# Copyright (c) 2025, NVIDIA CORPORATION.  All rights reserved.
# Copyright (c) 2026, ZDTaichu-5.0-9B Contributors.  All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
# SPDX-License-Identifier: Apache-2.0

"""Standalone inference-only C-RADIO ViT vision tower.

This file intentionally contains the small subset of C-RADIO needed by the
ZDTaichu-5.0-9B checkpoint. It does not depend on the cradio_v4 package.
"""

from __future__ import annotations

import math
from contextlib import contextmanager
from types import MethodType
from typing import Callable, Iterable, List, NamedTuple, Optional, Tuple, Union

import torch
import torch.nn.functional as F
from torch import nn
from transformers import PreTrainedModel

try:
    from timm.models import VisionTransformer, checkpoint_seq
except ImportError as exc:  # pragma: no cover - import-time dependency guard
    raise ImportError("cradio_model.py requires timm to build the C-RADIO ViT tower") from exc

from .cradio_config import RADIOConfig

class Resolution(NamedTuple):
    height: int
    width: int


class RadioOutput(NamedTuple):
    summary: Optional[torch.Tensor]
    features: Optional[torch.Tensor]

    def to(self, *args, **kwargs) -> "RadioOutput":
        return RadioOutput(
            self.summary.to(*args, **kwargs) if self.summary is not None else None,
            self.features.to(*args, **kwargs) if self.features is not None else None,
        )


class InputConditioner(nn.Module):
    def __init__(
        self,
        input_scale: float,
        norm_mean: Union[Tuple[float, float, float], torch.Tensor],
        norm_std: Union[Tuple[float, float, float], torch.Tensor],
        dtype: Optional[torch.dtype] = None,
    ) -> None:
        super().__init__()
        self.dtype = dtype
        self.register_buffer("norm_mean", torch.as_tensor(norm_mean, dtype=torch.float32).view(-1, 1, 1) / input_scale)
        self.register_buffer("norm_std", torch.as_tensor(norm_std, dtype=torch.float32).view(-1, 1, 1) / input_scale)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        y = (x - self.norm_mean) / self.norm_std
        if self.dtype is not None:
            y = y.to(self.dtype)
        return y


def get_default_conditioner() -> InputConditioner:
    from timm.data.constants import OPENAI_CLIP_MEAN, OPENAI_CLIP_STD

    return InputConditioner(1.0, OPENAI_CLIP_MEAN, OPENAI_CLIP_STD)


class ClsToken(nn.Module):
    def __init__(
        self,
        ndim: int,
        num_tokens: int = 1,
        enabled: bool = True,
        register_multiple: Optional[int] = None,
        num_registers: Optional[int] = None,
    ) -> None:
        super().__init__()
        self.ndim = ndim
        self.enabled = enabled
        self.num_registers = 0
        self.num_tokens = num_tokens
        if enabled:
            if num_registers:
                self.num_registers = num_registers
            elif register_multiple:
                self.num_registers = register_multiple - (num_tokens % register_multiple)
            scale = ndim ** -0.5
            self.token = nn.Parameter(torch.randn(num_tokens + self.num_registers, ndim) * scale)
        else:
            self.token = None
        self.num_patches = self.num_tokens + self.num_registers

    def disable(self) -> None:
        self.token = None
        self.enabled = False

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        if self.token is None:
            return x
        token = self.token.unsqueeze(0).expand(x.shape[0], -1, -1)
        return torch.cat([token, x], dim=1)

    def no_weight_decay(self) -> List[str]:
        return ["token"]


class Im2Patches(nn.Module):
    def __init__(self, patch_size: int) -> None:
        super().__init__()
        self.patch_size = patch_size

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        if self.patch_size == 1:
            return x.flatten(2).transpose(1, 2)
        return F.unfold(x, kernel_size=self.patch_size, stride=self.patch_size).transpose(1, 2)


class ViTPatchLinear(nn.Linear):
    def __init__(self, patch_size: int, embed_dim: int, bias: bool = False, **factory) -> None:
        super().__init__(3 * (patch_size ** 2), embed_dim, bias=bias, **factory)
        self.patch_size = patch_size


class ViTPatchGenerator(nn.Module):
    def __init__(
        self,
        patch_size: int,
        embed_dim: int,
        input_dims: Union[int, Tuple[int, int]],
        abs_pos: bool = True,
        normalize_patches: bool = False,
        cls_token: bool = False,
        max_input_dims: Optional[Union[int, Tuple[int, int]]] = None,
        pos_dropout: float = 0.0,
        return_pos_enc: bool = False,
        num_cls_tokens: int = 1,
        register_multiple: Optional[int] = None,
        num_registers: Optional[int] = None,
        patch_bias: bool = False,
        device=None,
        dtype=None,
    ) -> None:
        super().__init__()
        if isinstance(input_dims, int):
            input_dims = (input_dims, input_dims)
        if max_input_dims is None:
            max_input_dims = input_dims
        if isinstance(max_input_dims, int):
            max_input_dims = (max_input_dims, max_input_dims)

        max_input_dims = tuple(int(math.ceil(d / patch_size) * patch_size) for d in max_input_dims)
        factory = dict(device=device, dtype=dtype)

        self.cpe_mode = max_input_dims != input_dims
        self.pos_dropout = pos_dropout
        self.return_pos_enc = return_pos_enc
        self.patch_size = patch_size
        self.abs_pos = abs_pos
        self.embed_dim = embed_dim
        self.num_rows = max_input_dims[0] // patch_size
        self.num_cols = max_input_dims[1] // patch_size
        self.input_dims = tuple(d // patch_size for d in input_dims)
        self.num_patches = self.num_rows * self.num_cols
        self.max_input_dims = max_input_dims
        self.im_to_patches = Im2Patches(patch_size)
        self.embedder = ViTPatchLinear(patch_size, embed_dim, bias=patch_bias, **factory)
        if abs_pos:
            scale = embed_dim ** -0.5
            self.pos_embed = nn.Parameter(torch.randn(1, self.num_patches, embed_dim, **factory) * scale)
        self.cls_token = ClsToken(
            embed_dim,
            num_tokens=num_cls_tokens,
            enabled=cls_token,
            register_multiple=register_multiple,
            num_registers=num_registers,
        )
        self.patch_normalizer = nn.LayerNorm(embed_dim) if normalize_patches else nn.Identity()
        self.num_video_frames = None

    @property
    def apply_cls_token(self) -> bool:
        return self.cls_token.enabled

    @property
    def num_cls_tokens(self) -> int:
        return self.cls_token.num_tokens

    @property
    def num_cls_patches(self) -> int:
        return self.cls_token.num_patches

    @property
    def num_registers(self) -> int:
        return self.cls_token.num_registers

    @property
    def num_skip(self) -> int:
        return self.num_cls_tokens + self.num_registers

    def no_weight_decay(self) -> List[str]:
        return ["pos_embed"]

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        patches = self.embedder(self.im_to_patches(x))
        patches, pos_enc = self.apply_pos_enc(patches, input_size=x.shape[2:])
        patches = self.cls_token(patches)
        patches = self.patch_normalizer(patches)
        if self.return_pos_enc:
            return patches, pos_enc
        return patches

    def apply_pos_enc(
        self,
        patches: torch.Tensor,
        patch_idxs: Optional[torch.Tensor] = None,
        input_size: Optional[Tuple[int, int]] = None,
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        if not self.abs_pos:
            return patches, torch.empty(0, device=patches.device, dtype=patches.dtype)
        pos_enc = self.get_pos_enc(patches.shape[0], patch_idxs, input_size)
        if self.training and self.pos_dropout > 0:
            keeps = torch.rand(patches.shape[0], 1, 1, dtype=pos_enc.dtype, device=pos_enc.device) > self.pos_dropout
            pos_enc_drop = torch.where(keeps, pos_enc, 0)
        else:
            pos_enc_drop = pos_enc
        return patches + pos_enc_drop, pos_enc

    def get_pos_enc(
        self,
        batch_size: int,
        patch_idxs: Optional[torch.Tensor] = None,
        input_size: Optional[Tuple[int, int]] = None,
    ) -> torch.Tensor:
        input_dims = self.input_dims if input_size is None else tuple(d // self.patch_size for d in input_size)
        pos_embed = self._get_pos_embeddings(batch_size, input_dims)
        if patch_idxs is None:
            return pos_embed
        exp_patch_idxs = patch_idxs.unsqueeze(-1).expand(-1, -1, pos_embed.shape[-1])
        return torch.gather(pos_embed.expand(patch_idxs.shape[0], -1, -1), dim=1, index=exp_patch_idxs)

    def _get_pos_embeddings(self, batch_size: int, input_dims: Tuple[int, int]) -> torch.Tensor:
        if (self.num_rows, self.num_cols) == input_dims:
            return self.pos_embed

        pos_embed = self.pos_embed.reshape(1, self.num_rows, self.num_cols, -1).permute(0, 3, 1, 2)

        def window_select(pe: torch.Tensor) -> torch.Tensor:
            if input_dims[0] < pe.shape[-2]:
                pe = pe[..., :input_dims[0], :]
            if input_dims[1] < pe.shape[-1]:
                pe = pe[..., :, :input_dims[1]]
            return pe

        if self.cpe_mode:
            if self.training:
                if self.num_video_frames is not None:
                    if batch_size % self.num_video_frames != 0:
                        raise ValueError(
                            f"Batch size {batch_size} must be divisible by num_video_frames "
                            f"{self.num_video_frames} for CPE mode."
                        )
                    batch_size //= self.num_video_frames

                min_scale = math.sqrt(0.1)
                scale = torch.rand(batch_size, 1, 1, device=pos_embed.device) * (1 - min_scale) + min_scale
                aspect_min = math.log(3 / 4)
                aspect = torch.exp(torch.rand(batch_size, 1, 1, device=pos_embed.device) * (-2 * aspect_min) + aspect_min)
                scale_xy = torch.stack([scale * aspect, scale / aspect], dim=-1).clamp_(0, 1)
                pos_xy = torch.rand(batch_size, 1, 1, 2, device=pos_embed.device) * (1 - scale_xy)
                lin_x = torch.linspace(0, 1, steps=input_dims[1], device=pos_embed.device)[None, None].expand(batch_size, input_dims[0], -1)
                lin_y = torch.linspace(0, 1, steps=input_dims[0], device=pos_embed.device)[None, :, None].expand(batch_size, -1, input_dims[1])
                grid_xy = torch.stack([lin_x, lin_y], dim=-1) * scale_xy + pos_xy
                grid_xy.mul_(2).sub_(1)
                pos_embed = F.grid_sample(
                    pos_embed.float().expand(batch_size, -1, -1, -1),
                    grid=grid_xy,
                    mode="bilinear",
                    padding_mode="zeros",
                    align_corners=True,
                ).to(pos_embed.dtype)
                if self.num_video_frames is not None:
                    pos_embed = torch.repeat_interleave(pos_embed, self.num_video_frames, dim=0)
            else:
                max_dim = max(input_dims)
                pos_embed = F.interpolate(pos_embed.float(), size=(max_dim, max_dim), align_corners=False, mode="bilinear").to(pos_embed.dtype)
                pos_embed = window_select(pos_embed)
        else:
            pos_embed = window_select(pos_embed)

        if pos_embed.shape[-2:] != input_dims:
            pos_embed = F.interpolate(pos_embed.float(), size=input_dims, align_corners=False, mode="bilinear").to(pos_embed.dtype)
        return pos_embed.flatten(2).permute(0, 2, 1)


def _forward_cpe(self: VisionTransformer, x: torch.Tensor) -> torch.Tensor:
    x = self.patch_generator(x)
    if getattr(self, "grad_checkpointing", False) and not torch.jit.is_scripting():
        x = checkpoint_seq(self.blocks, x)
    else:
        x = self.blocks(x)
    x = self.norm(x)
    return x


@contextmanager
def _video_mode(self: VisionTransformer, t: int):
    original_num_frames = self.patch_generator.num_video_frames
    self.patch_generator.num_video_frames = t
    try:
        yield
    finally:
        self.patch_generator.num_video_frames = original_num_frames


def enable_cpe(
    model: VisionTransformer,
    max_img_size: Union[int, Tuple[int, int]] = 1024,
    num_cls_tokens: int = 1,
    pos_dropout: float = 0.1,
    register_multiple: Optional[int] = None,
    num_registers: Optional[int] = None,
) -> None:
    if not isinstance(model, VisionTransformer):
        raise ValueError(f"CPE only supports timm VisionTransformer models, got {type(model)}")

    patch_size = model.patch_embed.patch_size[0]
    embed_dim = model.embed_dim
    input_dims = model.patch_embed.img_size
    normalize_patches = not isinstance(model.patch_embed.norm, nn.Identity)
    cls_token = model.cls_token is not None
    if isinstance(max_img_size, int):
        max_img_size = int(round(max_img_size / patch_size) * patch_size)
    else:
        max_img_size = tuple(int(round(d / patch_size) * patch_size) for d in max_img_size)

    model.patch_generator = ViTPatchGenerator(
        patch_size=patch_size,
        embed_dim=embed_dim,
        input_dims=input_dims,
        normalize_patches=normalize_patches,
        cls_token=cls_token,
        max_input_dims=max_img_size,
        pos_dropout=pos_dropout,
        num_cls_tokens=num_cls_tokens,
        register_multiple=register_multiple,
        num_registers=num_registers,
    )
    model.patch_embed = None
    model.cls_token = None
    model.pos_embed = None
    model.pos_drop = None
    model.patch_size = patch_size
    model.num_cls_tokens = num_cls_tokens
    model.num_registers = model.patch_generator.num_registers
    model.forward_features = MethodType(_forward_cpe, model)
    model.cpe_video_mode = MethodType(_video_mode, model)


class FeatureNormalizer(nn.Module):
    def __init__(self, embed_dim: int, dtype: torch.dtype = torch.float32) -> None:
        super().__init__()
        self.register_buffer("mean", torch.zeros(embed_dim, dtype=dtype))
        self.register_buffer("tx", torch.eye(embed_dim, dtype=dtype))

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        if x.ndim <= 3:
            return (x - self.mean) @ self.tx.T
        if x.ndim == 4:
            kernel = self.tx.reshape(*self.tx.shape, 1, 1)
            return F.conv2d(x - self.mean.reshape(1, -1, 1, 1), weight=kernel, bias=None, stride=1, padding=0)
        raise ValueError(f"Unsupported input dimension: {x.ndim}, shape: {x.shape}")


class InnerRADIOModel(nn.Module):
    def __init__(
        self,
        model: nn.Module,
        input_conditioner: nn.Module,
        patch_size: int,
        max_resolution: int,
        preferred_resolution: Resolution,
        summary_idxs: Optional[torch.Tensor] = None,
        feature_normalizer: Optional[nn.Module] = None,
        window_size: Optional[int] = None,
    ) -> None:
        super().__init__()
        self.model = model
        self.input_conditioner = input_conditioner
        if summary_idxs is not None:
            self.register_buffer("summary_idxs", summary_idxs)
        else:
            self.summary_idxs = None
        self._preferred_resolution = preferred_resolution
        self._patch_size = patch_size
        self._max_resolution = max_resolution
        self._window_size = window_size
        self.feature_normalizer = feature_normalizer if feature_normalizer is not None else nn.Identity()

    @property
    def num_summary_tokens(self) -> int:
        patch_gen = getattr(self.model, "patch_generator", None)
        if patch_gen is not None:
            return patch_gen.num_skip
        if getattr(self.model, "global_pool", None) == "avg":
            return 0
        return 1

    @property
    def num_cls_tokens(self) -> int:
        patch_gen = getattr(self.model, "patch_generator", None)
        if patch_gen is not None:
            return patch_gen.num_cls_tokens
        if getattr(self.model, "global_pool", None) == "avg":
            return 0
        return 1

    @property
    def patch_size(self) -> int:
        if self._patch_size is not None:
            return self._patch_size
        if hasattr(self.model, "patch_size"):
            return self.model.patch_size
        patch_gen = getattr(self.model, "patch_generator", None)
        if patch_gen is not None:
            return patch_gen.patch_size
        raise AttributeError("Unable to infer patch_size from RADIO vision model")

    @property
    def max_resolution(self) -> int:
        return self._max_resolution

    @property
    def preferred_resolution(self) -> Resolution:
        return self._preferred_resolution

    @property
    def window_size(self) -> Optional[int]:
        return self._window_size

    @property
    def min_resolution_step(self) -> int:
        res = self.patch_size
        if self.window_size is not None:
            res *= self.window_size
        return res

    @property
    def blocks(self) -> Iterable[nn.Module]:
        return getattr(self.model, "blocks", None)

    @property
    def embed_dim(self) -> int:
        return self.model.embed_dim

    @property
    def summary_dim(self) -> int:
        embed_dim = self.embed_dim
        if self.summary_idxs is not None:
            embed_dim *= self.summary_idxs.shape[0]
        return embed_dim

    def make_preprocessor_external(self) -> Callable[[torch.Tensor], torch.Tensor]:
        ret = self.input_conditioner
        self.input_conditioner = nn.Identity()
        return ret

    def get_nearest_supported_resolution(self, height: int, width: int) -> Resolution:
        height = int(round(height / self.min_resolution_step) * self.min_resolution_step)
        width = int(round(width / self.min_resolution_step) * self.min_resolution_step)
        return Resolution(max(height, self.min_resolution_step), max(width, self.min_resolution_step))

    def switch_to_deploy(self) -> None:
        fn = getattr(self.model, "switch_to_deploy", None)
        if fn is not None:
            fn()

    def cpe_video_mode(self, t: int):
        return self.model.cpe_video_mode(t)

    def forward(self, x: torch.Tensor, feature_fmt: str = "NLC") -> RadioOutput:
        res_step = self.min_resolution_step
        if res_step is not None and (x.shape[-2] % res_step != 0 or x.shape[-1] % res_step != 0):
            raise ValueError(
                "The input resolution must be a multiple of self.min_resolution_step. "
                f"Input: {x.shape[-2:]}, Nearest: {self.get_nearest_supported_resolution(*x.shape[-2:])}"
            )
        x = self.input_conditioner(x)
        y = self.model.forward_features(x)
        return self._extract_final(x, y, feature_fmt=feature_fmt)

    def _extract_final(self, x: torch.Tensor, y: torch.Tensor, feature_fmt: str = "NLC") -> RadioOutput:
        patch_gen = getattr(self.model, "patch_generator", None)
        if patch_gen is not None:
            all_summary = y[:, : patch_gen.num_cls_tokens]
            bb_summary = all_summary[:, self.summary_idxs] if self.summary_idxs is not None else all_summary
            all_feat = y[:, patch_gen.num_skip :]
        elif getattr(self.model, "global_pool", None) == "avg":
            all_summary = y[:, self.model.num_prefix_tokens :].mean(dim=1)
            bb_summary = all_summary
            all_feat = y
        else:
            all_summary = y[:, 0]
            bb_summary = all_summary
            all_feat = y[:, 1:]

        all_feat = self.feature_normalizer(all_feat)
        if feature_fmt == "NCHW":
            fmt_feat = all_feat.reshape(
                all_feat.shape[0],
                x.shape[-2] // self.patch_size,
                x.shape[-1] // self.patch_size,
                all_feat.shape[2],
            ).permute(0, 3, 1, 2)
        elif feature_fmt == "NLC":
            fmt_feat = all_feat
        else:
            raise ValueError(f"Unsupported feature_fmt: {feature_fmt}. Must be one of ['NLC', 'NCHW']")
        return RadioOutput(bb_summary.flatten(1), fmt_feat)


def _as_namespace(value):
    if value is None:
        return type("RADIOArgs", (), {})()
    if isinstance(value, dict):
        ns = type("RADIOArgs", (), {})()
        for k, v in value.items():
            setattr(ns, k, v)
        return ns
    return value


def _dtype_from_config(config: RADIOConfig) -> torch.dtype:
    dtype_name = getattr(config, "dtype", None) or getattr(config, "amp_dtype", None)
    if isinstance(dtype_name, torch.dtype):
        return dtype_name
    if isinstance(dtype_name, str) and hasattr(torch, dtype_name):
        return getattr(torch, dtype_name)
    return torch.float32


def create_vit_from_config(config: RADIOConfig) -> VisionTransformer:
    args = _as_namespace(getattr(config, "args", {}))
    model_name = getattr(args, "model", None) or "vit_huge_patch16_224"
    if model_name != "vit_huge_patch16_224":
        raise ValueError(
            "This standalone cradio_model.py keeps only the ZDTaichu ViT-H/16 structure. "
            f"Unsupported RADIO args.model={model_name!r}."
        )

    model = VisionTransformer(
        img_size=224,
        patch_size=16,
        embed_dim=1280,
        depth=32,
        num_heads=16,
        mlp_ratio=4.0,
        qkv_bias=True,
        num_classes=0,
        global_pool="",
    )

    # The ZDTaichu checkpoint was exported after RADIO removed the final ViT norm/head
    # and replaced patch embedding, cls token, and absolute pos embedding with CPE.
    if hasattr(model, "norm") and not getattr(args, "model_norm", False):
        model.norm = nn.Identity()
    model.head = nn.Identity()

    cpe_max_size = getattr(args, "cpe_max_size", None) or getattr(config, "max_resolution", None)
    if cpe_max_size is not None:
        teachers = getattr(args, "teachers", []) or []
        teacher_names = {t.get("name") for t in teachers if isinstance(t, dict) and t.get("name")}
        num_cls_tokens = len(teacher_names) if getattr(args, "cls_token_per_teacher", False) and teacher_names else 1
        enable_cpe(
            model,
            cpe_max_size,
            num_cls_tokens=num_cls_tokens,
            register_multiple=getattr(args, "register_multiple", None),
            num_registers=getattr(args, "cpe_num_registers", None),
        )
    return model


class RADIOModel(PreTrainedModel):
    """Inference-only HuggingFace wrapper for the ZDTaichu C-RADIO ViT tower."""

    config_class = RADIOConfig
    base_model_prefix = "radio_model"
    main_input_name = "pixel_values"
    supports_gradient_checkpointing = False

    def __init__(self, config: RADIOConfig) -> None:
        super().__init__(config)
        args = _as_namespace(getattr(config, "args", {}))
        dtype = _dtype_from_config(config)
        vit = create_vit_from_config(config)

        summary_idxs = None
        if getattr(args, "cls_token_per_teacher", False):
            teachers = getattr(args, "teachers", []) or []
            if teachers:
                summary_idxs = torch.tensor(
                    [i for i, t in enumerate(teachers) if not isinstance(t, dict) or t.get("use_summary", True)],
                    dtype=torch.int64,
                )

        feature_normalizer = None
        fn_cfg = getattr(config, "feature_normalizer_config", None)
        if fn_cfg is not None:
            embed_dim = fn_cfg.get("embed_dim", vit.embed_dim) if isinstance(fn_cfg, dict) else vit.embed_dim
            feature_normalizer = FeatureNormalizer(embed_dim, dtype=torch.float32)

        pref = getattr(config, "preferred_resolution", (512, 512))
        self.radio_model = InnerRADIOModel(
            model=vit,
            input_conditioner=get_default_conditioner(),
            patch_size=getattr(config, "patch_size", 16),
            max_resolution=getattr(config, "max_resolution", 2048),
            preferred_resolution=Resolution(int(pref[0]), int(pref[1])),
            summary_idxs=summary_idxs,
            feature_normalizer=feature_normalizer,
            window_size=getattr(config, "vitdet_window_size", None),
        )
        if dtype is not torch.float32:
            self.radio_model = self.radio_model.to(dtype=dtype)

    @property
    def adaptors(self):
        return nn.ModuleDict()

    @property
    def model(self) -> nn.Module:
        return self.radio_model.model

    @property
    def input_conditioner(self) -> nn.Module:
        return self.radio_model.input_conditioner

    @property
    def num_summary_tokens(self) -> int:
        return self.radio_model.num_summary_tokens

    @property
    def patch_size(self) -> int:
        return self.radio_model.patch_size

    @property
    def max_resolution(self) -> int:
        return self.radio_model.max_resolution

    @property
    def preferred_resolution(self) -> Resolution:
        return self.radio_model.preferred_resolution

    @property
    def window_size(self) -> Optional[int]:
        return self.radio_model.window_size

    @property
    def min_resolution_step(self) -> int:
        return self.radio_model.min_resolution_step

    def make_preprocessor_external(self) -> Callable[[torch.Tensor], torch.Tensor]:
        return self.radio_model.make_preprocessor_external()

    def get_nearest_supported_resolution(self, height: int, width: int) -> Resolution:
        return self.radio_model.get_nearest_supported_resolution(height, width)

    def switch_to_deploy(self) -> None:
        self.radio_model.switch_to_deploy()

    def forward(self, pixel_values: torch.Tensor, feature_fmt: str = "NLC", **kwargs) -> RadioOutput:
        return self.radio_model(pixel_values, feature_fmt=feature_fmt)


__all__ = [
    "RADIOModel",
    "RADIOConfig",
    "RadioOutput",
    "Resolution",
    "InputConditioner",
    "ViTPatchGenerator",
]
