ZDTaichu5.0-9B-NVFP4 / cradio_model.py
TaichuAI's picture
Upload folder using huggingface_hub
f5f6734 verified
Raw History Blame Contribute Delete
26.5 kB
# 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",
]