diff --git a/README.md b/README.md index f9f2320b6..a4447b29b 100644 --- a/README.md +++ b/README.md @@ -29,7 +29,7 @@ multi-component export for pipelines. |---|---| | **Text Generation** | Llama 2/3/4, Mistral, Qwen 2/2.5/3/3.5/3.6, Phi-3/3.5, Gemma 1/2/3/4, Granite, GPT-2, OPT, OLMo, SmolLM3, and many more | | **Mixture of Experts** | PhiMoE, GPTOSS, Mixtral, OLMoE, DeepSeek-V2/V3, Qwen2-MoE, Qwen3-MoE, Qwen3-Next, GLM-4-MoE, Arctic, DBRX, Jamba | -| **Multimodal** | Gemma 3/4, Phi-4MM (vision + audio + LoRA), LLaVA, InternVL2, Qwen2.5-VL, Qwen3-VL, Qwen3.5/3.6-VL, Pixtral | +| **Multimodal** | Gemma 3/4, Phi-4MM (vision + audio + LoRA), Nemotron Parse, LLaVA, InternVL2, Qwen2.5-VL, Qwen3-VL, Qwen3.5/3.6-VL, Pixtral | | **Encoder-only** | BERT, RoBERTa, ALBERT, DeBERTa, DistilBERT, ELECTRA, XLNet | | **Encoder-Decoder** | BART, T5/mT5, Marian, M2M-100, Pegasus, BigBird-Pegasus | | **Speech-to-Text** | Whisper, FastConformer-RNNT, FunASR, Qwen3-ASR, SenseVoice | diff --git a/scripts/generate_golden.py b/scripts/generate_golden.py index 95dec62eb..1264880fb 100644 --- a/scripts/generate_golden.py +++ b/scripts/generate_golden.py @@ -511,6 +511,86 @@ def _generate_vision_language(case: TestCase, json_path: Path, device: str) -> N ) +def _generate_image_to_text(case: TestCase, json_path: Path, device: str) -> None: + """Generate Nemotron Parse-style image-to-text reference data.""" + import torch + from PIL import Image + from transformers import AutoModel, AutoProcessor + + from mobius._testing.golden import save_generation_json, save_golden_ref + + dtype_map = { + "float32": torch.float32, + "float16": torch.float16, + "bfloat16": torch.bfloat16, + } + model = AutoModel.from_pretrained( + case.model_id, + revision=case.revision, + torch_dtype=dtype_map[case.dtype], + trust_remote_code=case.trust_remote_code, + ).to(device) + model.eval() + processor = AutoProcessor.from_pretrained( + case.model_id, + revision=case.revision, + trust_remote_code=case.trust_remote_code, + ) + images = [ + Image.open(Path("testdata") / image_path).convert("RGB") for image_path in case.images + ] + processed = processor( + images=images, + text=case.decoder_prompt, + return_tensors="pt", + add_special_tokens=False, + ).to(device) + decoder_input_ids = processed["input_ids"] + + with torch.no_grad(): + outputs = model( + pixel_values=processed["pixel_values"], + decoder_input_ids=decoder_input_ids, + ) + last_logits = outputs.logits[0, -1].float().cpu().numpy() + golden = _extract_logits_golden(last_logits) + input_ids = decoder_input_ids.cpu().numpy() + + generated_ids = None + if "L5" in case.level: + with torch.no_grad(): + generated = model.generate( + pixel_values=processed["pixel_values"], + decoder_input_ids=decoder_input_ids, + max_new_tokens=case.generation_params.get("max_new_tokens", 20), + do_sample=False, + ) + generated_ids = generated[0, input_ids.shape[1] :].cpu().numpy() + + save_golden_ref( + json_path, + top1_id=golden["top1_id"], + top2_id=golden["top2_id"], + top10_ids=golden["top10_ids"], + top10_logits=golden["top10_logits"], + logits_summary=golden["logits_summary"], + input_ids=input_ids, + ) + + if generated_ids is not None: + gen_path = json_path.with_name(json_path.stem + "_generation.json") + save_generation_json( + gen_path, + model_id=case.model_id, + prompt=case.decoder_prompt, + generated_tokens=generated_ids.tolist(), + generated_text=processor.decode( + generated_ids.tolist(), + skip_special_tokens=False, + ), + ) + + def _generate_speech_to_text(case: TestCase, json_path: Path, device: str) -> None: """Generate golden data for a speech-to-text (Whisper) model.""" import librosa @@ -1646,6 +1726,7 @@ def _hook(_module, _args, kwargs, output): "feature-extraction": _generate_encoder, "seq2seq": _generate_seq2seq, "image-text-to-text": _generate_vision_language, + "image-to-text": _generate_image_to_text, "image-classification": _generate_image_classification, "speech-to-text": _generate_speech_to_text, "speech-language": _generate_speech_language, diff --git a/src/mobius/_configs/__init__.py b/src/mobius/_configs/__init__.py index afb241a16..50293ff9a 100644 --- a/src/mobius/_configs/__init__.py +++ b/src/mobius/_configs/__init__.py @@ -45,6 +45,7 @@ MMSConfig, NanoChatConfig, NemotronHConfig, + NemotronParseConfig, Qwen35MtpConfig, Sam2Config, SegformerConfig, @@ -107,6 +108,7 @@ "MllamaConfig", "MMSConfig", "NanoChatConfig", + "NemotronParseConfig", "NemotronHConfig", "QuantizationConfig", "Qwen35MtpConfig", diff --git a/src/mobius/_configs/_base.py b/src/mobius/_configs/_base.py index 475b69b9a..412fb40cc 100644 --- a/src/mobius/_configs/_base.py +++ b/src/mobius/_configs/_base.py @@ -1267,6 +1267,101 @@ class VisionLanguageConfig(CausalLMConfig): """ +@dataclasses.dataclass +class NemotronParseConfig(ArchitectureConfig): + """Configuration for NVIDIA Nemotron Parse image-to-text models.""" + + image_height: int = 2048 + image_width: int = 1664 + vision_max_grid_size: int = 128 + num_summary_tokens: int = 3 + decoder_start_token_id: int = 2 + scale_embedding: bool = True + add_final_layer_norm: bool = True + + @classmethod + def from_transformers(cls, config, parent_config=None) -> NemotronParseConfig: + del parent_config + import types + + def _namespace(value): + if isinstance(value, dict): + return types.SimpleNamespace( + **{key: _namespace(item) for key, item in value.items()} + ) + if isinstance(value, list): + return [_namespace(item) for item in value] + return value + + decoder = _namespace(getattr(config, "decoder", None)) + if decoder is None: + raise ValueError("Nemotron Parse config is missing its decoder sub-config") + base = ArchitectureConfig.from_transformers(decoder, parent_config=config) + fields = _shallow_fields(base) + num_attention_heads = int( + getattr(decoder, "decoder_attention_heads", None) + or getattr(decoder, "num_attention_heads", fields["num_attention_heads"]) + ) + hidden_size = int(fields["hidden_size"]) + if hidden_size % num_attention_heads: + raise ValueError( + "Nemotron Parse decoder hidden size must be divisible by its attention heads" + ) + fields.update( + num_attention_heads=num_attention_heads, + num_key_value_heads=num_attention_heads, + head_dim=hidden_size // num_attention_heads, + ) + + raw_image_size = getattr(config, "image_size", (2048, 1664)) + if isinstance(raw_image_size, int): + image_height = image_width = raw_image_size + else: + image_height, image_width = (int(raw_image_size[0]), int(raw_image_size[1])) + + encoder = _namespace(getattr(config, "encoder", None)) + patch_size = int(getattr(encoder, "patch_size", 16)) + max_resolution = int( + getattr(encoder, "max_resolution", max(image_height, image_width)) + ) + fields.update( + model_type="nemotron_parse", + bos_token_id=getattr(config, "bos_token_id", fields.get("bos_token_id")), + eos_token_id=getattr(config, "eos_token_id", fields.get("eos_token_id")), + pad_token_id=getattr(config, "pad_token_id", fields["pad_token_id"]), + tie_word_embeddings=getattr(config, "tie_word_embeddings", True), + max_position_embeddings=( + getattr(config, "max_sequence_length", None) + or fields["max_position_embeddings"] + ), + vision=VisionConfig( + hidden_size=1280, + intermediate_size=5120, + num_hidden_layers=32, + num_attention_heads=16, + image_size=max_resolution, + patch_size=patch_size, + norm_eps=1e-6, + model_type="radio_v2.5-h", + in_channels=3, + ), + ) + resolved_dtype = _resolve_dtype(config) + if resolved_dtype is not None: + fields["dtype"] = resolved_dtype + + return cls( + **fields, + image_height=image_height, + image_width=image_width, + vision_max_grid_size=max_resolution // patch_size, + num_summary_tokens=3, + decoder_start_token_id=int(getattr(config, "decoder_start_token_id", 2)), + scale_embedding=bool(getattr(decoder, "scale_embedding", True)), + add_final_layer_norm=bool(getattr(decoder, "add_final_layer_norm", True)), + ) + + # --------------------------------------------------------------------------- # Model-family subclasses — add model-specific fields # --------------------------------------------------------------------------- diff --git a/src/mobius/_configs/_base_test.py b/src/mobius/_configs/_base_test.py new file mode 100644 index 000000000..7aead7549 --- /dev/null +++ b/src/mobius/_configs/_base_test.py @@ -0,0 +1,42 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Tests for architecture-specific configuration extraction.""" + +from __future__ import annotations + +import types + +from mobius._configs import NemotronParseConfig + + +def test_nemotron_parse_maps_raw_mbart_decoder_attention_heads(): + """The Hub's non-trusted config exposes MBART's decoder-specific aliases.""" + config = types.SimpleNamespace( + model_type="nemotron_parse", + decoder={ + "model_type": "nemotron_parse_text", + "d_model": 1024, + "decoder_attention_heads": 16, + "decoder_ffn_dim": 4096, + "decoder_layers": 10, + "num_hidden_layers": 12, + "vocab_size": 72256, + "pad_token_id": 1, + }, + encoder={"patch_size": 16, "max_resolution": 2048}, + image_size=[2048, 1664], + max_sequence_length=9000, + bos_token_id=0, + eos_token_id=2, + pad_token_id=1, + tie_word_embeddings=True, + decoder_start_token_id=2, + ) + + extracted = NemotronParseConfig.from_transformers(config) + + assert extracted.num_attention_heads == 16 + assert extracted.num_key_value_heads == 16 + assert extracted.head_dim == 64 + assert extracted.num_decoder_layers == 10 diff --git a/src/mobius/_registry.py b/src/mobius/_registry.py index 1c611e97c..5ae820ec4 100644 --- a/src/mobius/_registry.py +++ b/src/mobius/_registry.py @@ -28,6 +28,7 @@ Gemma4AssistantConfig, Gemma4Config, MMSConfig, + NemotronParseConfig, WhisperConfig, ) from mobius.models import ( @@ -73,6 +74,7 @@ MoECausalLMModel, NanoChatCausalLMModel, NemotronCausalLMModel, + NemotronParseForConditionalGeneration, OLMo2CausalLMModel, OLMoCausalLMModel, Phi3CausalLMModel, @@ -564,6 +566,11 @@ def _detect_fallback_registration(hf_config) -> ModelRegistration | None: "bamba": ModelRegistration(BambaCausalLMModel), "jamba": ModelRegistration(JambaCausalLMModel), "nemotron_h": ModelRegistration(NemotronHCausalLMModel), + "nemotron_parse": ModelRegistration( + NemotronParseForConditionalGeneration, + task="vision-encoder-decoder", + config_class=NemotronParseConfig, + ), # --- Hybrid linear-attention --- "longcat_flash": ModelRegistration(LongcatFlashCausalLMModel), # --- Multimodal --- @@ -863,6 +870,7 @@ def _create_default_registry() -> ModelRegistry: "csm": "sesame/csm-1b", "evolla": "westlake-repl/Evolla-10B-hf", "nemotron_h": "nvidia/NVIDIA-Nemotron-3-Nano-4B-BF16", + "nemotron_parse": "nvidia/NVIDIA-Nemotron-Parse-2.0", "open-llama": "openlm-research/open_llama_3b", "persimmon": "adept/persimmon-8b-base", "shieldgemma2": "google/shieldgemma-2b", diff --git a/src/mobius/components/__init__.py b/src/mobius/components/__init__.py index 3baeb6e77..99f1d56c2 100644 --- a/src/mobius/components/__init__.py +++ b/src/mobius/components/__init__.py @@ -56,6 +56,7 @@ "PostNormDecoderLayer", "QuantizedEmbedding", "QuantizedLinear", + "RadioVisionModel", "RMSNorm", "SelectiveScan", "SiLU", @@ -251,6 +252,7 @@ from mobius.components._qwen25_vl_vision import ( Qwen25VLVisionRotaryEmbedding as Qwen25VLVisionRotaryEmbedding, ) +from mobius.components._radio_vision import RadioVisionModel from mobius.components._rms_norm import ( GatedRMSNorm, OffsetRMSNorm, diff --git a/src/mobius/components/_conv.py b/src/mobius/components/_conv.py index 4cbccd1f0..0256431f3 100644 --- a/src/mobius/components/_conv.py +++ b/src/mobius/components/_conv.py @@ -5,6 +5,7 @@ from __future__ import annotations +from collections.abc import Sequence from typing import TYPE_CHECKING import onnx_ir as ir @@ -14,6 +15,14 @@ pass +def _pair(value: int | Sequence[int], name: str) -> tuple[int, int]: + if isinstance(value, int): + return value, value + if len(value) != 2: + raise ValueError(f"{name} must contain exactly two values, got {value!r}") + return int(value[0]), int(value[1]) + + class Conv2d(nn.Module): """2D convolution with bias. @@ -25,30 +34,31 @@ def __init__( self, in_channels: int, out_channels: int, - kernel_size: int = 3, - stride: int = 1, - padding: int = 0, + kernel_size: int | tuple[int, int] = 3, + stride: int | tuple[int, int] = 1, + padding: int | tuple[int, int] = 0, groups: int = 1, ): super().__init__() - self.weight = nn.Parameter( - (out_channels, in_channels // groups, kernel_size, kernel_size) - ) + kernel_h, kernel_w = _pair(kernel_size, "kernel_size") + self.weight = nn.Parameter((out_channels, in_channels // groups, kernel_h, kernel_w)) self.bias = nn.Parameter((out_channels,)) - self._kernel_size = kernel_size + self._kernel_size = (kernel_h, kernel_w) self._stride = stride + self._strides = _pair(stride, "stride") self._padding = padding + self._pads = _pair(padding, "padding") self._groups = groups def forward(self, op: OpBuilder, x: ir.Value): - p = self._padding + pad_h, pad_w = self._pads return op.Conv( x, self.weight, self.bias, - kernel_shape=[self._kernel_size, self._kernel_size], - strides=[self._stride, self._stride], - pads=[p, p, p, p], + kernel_shape=list(self._kernel_size), + strides=list(self._strides), + pads=[pad_h, pad_w, pad_h, pad_w], group=self._groups, ) @@ -60,28 +70,29 @@ def __init__( self, in_channels: int, out_channels: int, - kernel_size: int = 3, - stride: int = 1, - padding: int = 0, + kernel_size: int | tuple[int, int] = 3, + stride: int | tuple[int, int] = 1, + padding: int | tuple[int, int] = 0, groups: int = 1, ): super().__init__() - self.weight = nn.Parameter( - (out_channels, in_channels // groups, kernel_size, kernel_size) - ) - self._kernel_size = kernel_size + kernel_h, kernel_w = _pair(kernel_size, "kernel_size") + self.weight = nn.Parameter((out_channels, in_channels // groups, kernel_h, kernel_w)) + self._kernel_size = (kernel_h, kernel_w) self._stride = stride + self._strides = _pair(stride, "stride") self._padding = padding + self._pads = _pair(padding, "padding") self._groups = groups def forward(self, op: OpBuilder, x: ir.Value): - p = self._padding + pad_h, pad_w = self._pads return op.Conv( x, self.weight, - kernel_shape=[self._kernel_size, self._kernel_size], - strides=[self._stride, self._stride], - pads=[p, p, p, p], + kernel_shape=list(self._kernel_size), + strides=list(self._strides), + pads=[pad_h, pad_w, pad_h, pad_w], group=self._groups, ) diff --git a/src/mobius/components/_conv_test.py b/src/mobius/components/_conv_test.py index c00054183..1b2318117 100644 --- a/src/mobius/components/_conv_test.py +++ b/src/mobius/components/_conv_test.py @@ -64,6 +64,18 @@ def test_forward_with_stride(self): conv(op, x) assert count_op_type(graph, "Conv") >= 1 + def test_asymmetric_kernel_stride_and_padding(self): + conv = Conv2d( + 3, + 16, + kernel_size=(1, 4), + stride=(1, 4), + padding=(0, 2), + ) + assert list(conv.weight.shape) == [16, 3, 1, 4] + assert conv._strides == (1, 4) + assert conv._pads == (0, 2) + class TestConv2dNoBias: """Tests for 2D convolution without bias.""" @@ -86,6 +98,11 @@ def test_forward_builds_graph(self): builder._adapt_outputs([result], "") assert count_op_type(graph, "Conv") >= 1 + def test_asymmetric_kernel_and_stride(self): + conv = Conv2dNoBias(3, 16, kernel_size=(1, 4), stride=(1, 4)) + assert list(conv.weight.shape) == [16, 3, 1, 4] + assert conv._strides == (1, 4) + class TestBatchNorm2d: """Tests for 2D batch normalization.""" diff --git a/src/mobius/components/_radio_vision.py b/src/mobius/components/_radio_vision.py new file mode 100644 index 000000000..b932a206c --- /dev/null +++ b/src/mobius/components/_radio_vision.py @@ -0,0 +1,226 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""C-RADIO ViT vision encoder components.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +from onnxscript import OpBuilder, nn + +from mobius.components import INT64_MAX, LayerNorm, Linear + +if TYPE_CHECKING: + import onnx_ir as ir + + +class _RadioClsToken(nn.Module): + def __init__(self, num_tokens: int, hidden_size: int): + super().__init__() + self.token = nn.Parameter([num_tokens, hidden_size]) + self._num_tokens = num_tokens + self._hidden_size = hidden_size + + def forward(self, op: OpBuilder, patches: ir.Value) -> ir.Value: + batch = op.Shape(patches, start=0, end=1) + shape = op.Concat( + batch, + op.Constant(value_ints=[self._num_tokens, self._hidden_size]), + axis=0, + ) + tokens = op.Expand(op.Unsqueeze(self.token, [0]), shape) + return op.Concat(tokens, patches, axis=1) + + +class _RadioPatchEmbedder(nn.Module): + """Checkpoint-aligned linear patch projector emitted as an ONNX Conv.""" + + def __init__(self, patch_size: int, hidden_size: int): + super().__init__() + self.weight = nn.Parameter([hidden_size, 3 * patch_size * patch_size]) + self._patch_size = patch_size + self._hidden_size = hidden_size + + def forward(self, op: OpBuilder, pixel_values: ir.Value) -> ir.Value: + patch_weight = op.Reshape( + self.weight, + [self._hidden_size, 3, self._patch_size, self._patch_size], + ) + return op.Conv( + pixel_values, + patch_weight, + kernel_shape=[self._patch_size, self._patch_size], + strides=[self._patch_size, self._patch_size], + ) + + +class RadioPatchGenerator(nn.Module): + """Linear patchification plus CPE positional embeddings and register tokens.""" + + def __init__( + self, + *, + image_height: int, + image_width: int, + patch_size: int, + max_grid_size: int, + hidden_size: int, + num_register_tokens: int = 8, + ): + super().__init__() + if image_height % patch_size or image_width % patch_size: + raise ValueError("RADIO image dimensions must be divisible by patch_size") + self.embedder = _RadioPatchEmbedder(patch_size, hidden_size) + self.pos_embed = nn.Parameter([1, max_grid_size * max_grid_size, hidden_size]) + self.cls_token = _RadioClsToken(num_register_tokens, hidden_size) + self._image_height = image_height + self._image_width = image_width + self._patch_size = patch_size + self._max_grid_size = max_grid_size + self._hidden_size = hidden_size + + grid_h = image_height // patch_size + grid_w = image_width // patch_size + if max(grid_h, grid_w) != max_grid_size: + raise ValueError( + "RADIO export currently requires a canvas whose longest patch-grid " + "dimension equals max_grid_size" + ) + + def forward(self, op: OpBuilder, pixel_values: ir.Value) -> ir.Value: + grid_h = self._image_height // self._patch_size + grid_w = self._image_width // self._patch_size + + # The checkpoint stores a linear patchification matrix; reshaping it + # inside the embedder gives the equivalent strided Conv2d. + patches = self.embedder(op, pixel_values) # (B, hidden, grid_h, grid_w) + patches = op.Reshape(patches, [0, self._hidden_size, -1]) + patches = op.Transpose(patches, perm=[0, 2, 1]) # (B, grid_h*grid_w, hidden) + + # At the checkpoint's 2048x1664 canvas the CPE algorithm is exactly a + # top-left crop from the learned 128x128 grid (no interpolation). + pos = op.Reshape( + self.pos_embed, + [1, self._max_grid_size, self._max_grid_size, self._hidden_size], + ) + pos = op.Slice(pos, [0, 0], [grid_h, grid_w], [1, 2]) + pos = op.Reshape(pos, [1, grid_h * grid_w, self._hidden_size]) + patches = op.Add(patches, pos) + + # Four teacher CLS tokens plus four padding registers are prepended. + return self.cls_token(op, patches) + + +class RadioAttention(nn.Module): + """Fused-QKV bidirectional attention used by timm ViT-Huge.""" + + def __init__(self, hidden_size: int, num_heads: int): + super().__init__() + self.qkv = Linear(hidden_size, 3 * hidden_size) + self.proj = Linear(hidden_size, hidden_size) + self._num_heads = num_heads + self._scale = float((hidden_size // num_heads) ** -0.5) + + def forward(self, op: OpBuilder, hidden_states: ir.Value) -> ir.Value: + qkv = self.qkv(op, hidden_states) + query, key, value = op.Split(qkv, axis=-1, num_outputs=3, _outputs=3) + hidden_states = op.Attention( + query, + key, + value, + q_num_heads=self._num_heads, + kv_num_heads=self._num_heads, + scale=self._scale, + ) + return self.proj(op, hidden_states) + + +class RadioMLP(nn.Module): + """Exact-GELU ViT feed-forward network with checkpoint-aligned names.""" + + def __init__(self, hidden_size: int, intermediate_size: int): + super().__init__() + self.fc1 = Linear(hidden_size, intermediate_size) + self.fc2 = Linear(intermediate_size, hidden_size) + + def forward(self, op: OpBuilder, hidden_states: ir.Value) -> ir.Value: + return self.fc2(op, op.Gelu(self.fc1(op, hidden_states))) + + +class RadioBlock(nn.Module): + """Pre-norm C-RADIO transformer block.""" + + def __init__( + self, + hidden_size: int, + intermediate_size: int, + num_heads: int, + norm_eps: float, + ): + super().__init__() + self.norm1 = LayerNorm(hidden_size, eps=norm_eps) + self.attn = RadioAttention(hidden_size, num_heads) + self.norm2 = LayerNorm(hidden_size, eps=norm_eps) + self.mlp = RadioMLP(hidden_size, intermediate_size) + + def forward(self, op: OpBuilder, hidden_states: ir.Value) -> ir.Value: + hidden_states = op.Add( + hidden_states, + self.attn(op, self.norm1(op, hidden_states)), + ) + return op.Add( + hidden_states, + self.mlp(op, self.norm2(op, hidden_states)), + ) + + +class RadioVisionModel(nn.Module): + """C-RADIOv2-H ViT backbone returning summary and spatial features.""" + + def __init__( + self, + *, + image_height: int, + image_width: int, + patch_size: int, + max_grid_size: int, + hidden_size: int, + intermediate_size: int, + num_layers: int, + num_heads: int, + norm_eps: float = 1e-6, + ): + super().__init__() + self.patch_generator = RadioPatchGenerator( + image_height=image_height, + image_width=image_width, + patch_size=patch_size, + max_grid_size=max_grid_size, + hidden_size=hidden_size, + ) + self.blocks = nn.ModuleList( + [ + RadioBlock( + hidden_size, + intermediate_size, + num_heads, + norm_eps, + ) + for _ in range(num_layers) + ] + ) + + def forward(self, op: OpBuilder, pixel_values: ir.Value) -> tuple[ir.Value, ir.Value]: + hidden_states = self.patch_generator(op, pixel_values) + for block in self.blocks: + hidden_states = block(op, hidden_states) + + # The first four outputs are teacher CLS tokens; summaries [0,1,2] + # are flattened, while all eight CLS/register tokens are discarded + # from the spatial feature sequence. + all_summary = op.Slice(hidden_states, [0], [4], [1]) + summary = op.Gather(all_summary, [0, 1, 2], axis=1) + summary = op.Reshape(summary, [0, -1]) + features = op.Slice(hidden_states, [8], [INT64_MAX], [1]) + return summary, features diff --git a/src/mobius/integrations/ort_genai/auto_export.py b/src/mobius/integrations/ort_genai/auto_export.py index 3b4786b5d..6fd31e945 100644 --- a/src/mobius/integrations/ort_genai/auto_export.py +++ b/src/mobius/integrations/ort_genai/auto_export.py @@ -1007,6 +1007,14 @@ def write_ort_genai_config( "This is set automatically when building with mobius.build(). " "Diffusion models (which have no config) are not supported." ) + if {"vision_encoder", "decoder"}.issubset(pkg) and "embedding" not in pkg: + model_type = getattr(config, "model_type", "unknown") + raise NotImplementedError( + "onnxruntime-genai does not support generic vision encoder-decoder " + f"packages such as {model_type!r}. Run the vision_encoder and decoder " + "ONNX sessions directly; emitting genai_config.json would create an " + "artifact that the runtime cannot load." + ) os.makedirs(directory, exist_ok=True) diff --git a/src/mobius/integrations/ort_genai/auto_export_test.py b/src/mobius/integrations/ort_genai/auto_export_test.py index 174db33dd..49bd9a8b4 100644 --- a/src/mobius/integrations/ort_genai/auto_export_test.py +++ b/src/mobius/integrations/ort_genai/auto_export_test.py @@ -665,6 +665,29 @@ def test_genai_config_json_is_written(self, tmp_path): assert "model" in data assert data["model"]["type"] == "qwen2" + def test_rejects_generic_vision_encoder_decoder_package(self, tmp_path): + import dataclasses + + from mobius._model_package import ModelPackage + + @dataclasses.dataclass + class FakeConfig: + model_type: str = "nemotron_parse" + + pkg = ModelPackage( + { + "vision_encoder": mock.MagicMock(), + "decoder": mock.MagicMock(), + }, + config=FakeConfig(), + ) + with pytest.raises( + NotImplementedError, + match="does not support generic vision encoder-decoder", + ): + write_ort_genai_config(pkg, str(tmp_path)) + assert not (tmp_path / "genai_config.json").exists() + def test_processor_config_written_with_vision(self, tmp_path): """image_processor.json is written when pkg.config.vision is set.""" import dataclasses diff --git a/src/mobius/models/__init__.py b/src/mobius/models/__init__.py index 9ca9e2ce3..f785f59c1 100644 --- a/src/mobius/models/__init__.py +++ b/src/mobius/models/__init__.py @@ -87,6 +87,7 @@ "NanoChatCausalLMModel", "NemotronCausalLMModel", "NemotronHCausalLMModel", + "NemotronParseForConditionalGeneration", "OLMo2CausalLMModel", "OLMoCausalLMModel", "OPTCausalLMModel", @@ -236,6 +237,7 @@ from mobius.models.nemo_rnnt import EncDecRNNTModel from mobius.models.nemotron import NemotronCausalLMModel from mobius.models.nemotron_h import NemotronHCausalLMModel +from mobius.models.nemotron_parse import NemotronParseForConditionalGeneration from mobius.models.olmo import OLMo2CausalLMModel, OLMoCausalLMModel from mobius.models.opt import OPTCausalLMModel from mobius.models.persimmon import PersimmonCausalLMModel diff --git a/src/mobius/models/nemotron_parse.py b/src/mobius/models/nemotron_parse.py new file mode 100644 index 000000000..b8a177678 --- /dev/null +++ b/src/mobius/models/nemotron_parse.py @@ -0,0 +1,299 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""NVIDIA Nemotron Parse vision-language encoder-decoder model. + +Replicates HuggingFace's ``NemotronParseForConditionalGeneration`` with a +C-RADIOv2-H vision encoder, convolutional feature neck, and a cross-attentive +mBART text decoder. The architecture is exported as separate +``vision_encoder`` and ``decoder`` ONNX models. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch +from onnxscript import OpBuilder, nn + +from mobius._configs import ArchitectureConfig +from mobius.components import ( + Conv2dNoBias, + Embedding, + EncoderDecoderAttention, + LayerNorm, + Linear, + RadioVisionModel, + create_padding_mask, +) + +if TYPE_CHECKING: + import onnx_ir as ir + + +class _RadioModel(nn.Module): + """Wrapper matching ``radio_model.model`` in the HuggingFace checkpoint.""" + + def __init__(self, config: ArchitectureConfig): + super().__init__() + self.model = RadioVisionModel( + image_height=config.image_height, + image_width=config.image_width, + patch_size=config.vision.patch_size, + max_grid_size=config.vision_max_grid_size, + hidden_size=config.vision.hidden_size, + intermediate_size=config.vision.intermediate_size, + num_layers=config.vision.num_hidden_layers, + num_heads=config.vision.num_attention_heads, + norm_eps=config.vision.norm_eps, + ) + + def forward(self, op: OpBuilder, pixel_values: ir.Value): + return self.model(op, pixel_values) + + +class _RadioEncoder(nn.Module): + """Wrapper matching the checkpoint's ``model_encoder`` module.""" + + def __init__(self, config: ArchitectureConfig): + super().__init__() + self.radio_model = _RadioModel(config) + + def forward(self, op: OpBuilder, pixel_values: ir.Value): + return self.radio_model(op, pixel_values) + + +class NemotronParseVisionEncoder(nn.Module): + """C-RADIO vision encoder and Nemotron Parse feature-compression neck.""" + + def __init__(self, config: ArchitectureConfig): + super().__init__() + self.config = config + self.model_encoder = _RadioEncoder(config) + self.conv1 = Linear(config.vision.hidden_size, config.hidden_size) + self.layer_norm1 = LayerNorm(config.hidden_size, eps=1e-6) + self.conv2 = Conv2dNoBias( + config.hidden_size, + config.hidden_size, + kernel_size=(1, 4), + stride=(1, 4), + ) + self.layer_norm2 = LayerNorm(config.hidden_size, eps=1e-6) + self.sum_proj = Linear( + config.num_summary_tokens * config.vision.hidden_size, + config.hidden_size, + ) + self.layer_norm3 = LayerNorm(config.hidden_size, eps=1e-6) + + def forward(self, op: OpBuilder, pixel_values: ir.Value): + # C-RADIO returns three flattened teacher summaries plus the spatial + # patch sequence: (B, 3*1280), (B, H/16*W/16, 1280). + summary, features = self.model_encoder(op, pixel_values) + + # Pointwise projection and normalization: (B, H/16*W/16, 1024). + features = self.conv1(op, features) + features = self.layer_norm1(op, features) + + # Restore the patch grid, compress every four horizontal patches, then + # flatten back to a sequence: (B, H/16*W/64, 1024). + batch = op.Shape(features, start=0, end=1) + grid_height = self.config.image_height // self.config.vision.patch_size + grid_width = self.config.image_width // self.config.vision.patch_size + grid_shape = op.Concat( + batch, + op.Constant(value_ints=[grid_height, grid_width, self.config.hidden_size]), + axis=0, + ) + features = op.Reshape(features, grid_shape) + features = op.Transpose(features, perm=[0, 3, 1, 2]) + features = self.conv2(op, features) + features = op.Transpose(features, perm=[0, 2, 3, 1]) + sequence_shape = op.Concat( + batch, + op.Constant( + value_ints=[ + grid_height * (grid_width // 4), + self.config.hidden_size, + ] + ), + axis=0, + ) + features = op.Reshape(features, sequence_shape) + features = self.layer_norm2(op, features) + + # Project the concatenated teacher summaries to one final visual token. + summary = self.sum_proj(op, summary) # (B, 1024) + summary = self.layer_norm3(op, summary) + summary = op.Unsqueeze(summary, [1]) # (B, 1, 1024) + return op.Concat(features, summary, axis=1) + + def preprocess_weights( + self, state_dict: dict[str, torch.Tensor] + ) -> dict[str, torch.Tensor]: + """Select vision weights and reshape the checkpoint's Conv1d kernel.""" + weights: dict[str, torch.Tensor] = {} + for key, value in state_dict.items(): + if not key.startswith("encoder."): + continue + key = key.removeprefix("encoder.") + if ".input_conditioner." in key or key.endswith(".summary_idxs"): + continue + if key == "conv1.weight": + value = value.squeeze(-1) + weights[key] = value + return weights + + +class _NemotronParseDecoderLayer(nn.Module): + """Pre-norm mBART decoder layer with self- and cross-attention.""" + + def __init__(self, config: ArchitectureConfig): + super().__init__() + self.self_attn = EncoderDecoderAttention(config, is_causal=True) + self.self_attn_layer_norm = LayerNorm(config.hidden_size, eps=1e-5) + self.encoder_attn = EncoderDecoderAttention(config) + self.encoder_attn_layer_norm = LayerNorm(config.hidden_size, eps=1e-5) + self.fc1 = Linear(config.hidden_size, config.intermediate_size) + self.fc2 = Linear(config.intermediate_size, config.hidden_size) + self.final_layer_norm = LayerNorm(config.hidden_size, eps=1e-5) + + def forward( + self, + op: OpBuilder, + hidden_states: ir.Value, + encoder_hidden_states: ir.Value, + self_attention_bias: ir.Value, + past_key_value: tuple[ir.Value, ir.Value] | None = None, + ): + # Pre-norm causal self-attention. + residual = hidden_states + hidden_states = self.self_attn_layer_norm(op, hidden_states) + hidden_states, self_kv = self.self_attn( + op, + hidden_states, + attention_bias=self_attention_bias, + past_key_value=past_key_value, + ) + hidden_states = op.Add(residual, hidden_states) + + # Pre-norm cross-attention over the compressed C-RADIO sequence. + residual = hidden_states + hidden_states = self.encoder_attn_layer_norm(op, hidden_states) + hidden_states, _ = self.encoder_attn( + op, hidden_states, key_value_states=encoder_hidden_states + ) + hidden_states = op.Add(residual, hidden_states) + + # Pre-norm exact-GELU feed-forward block. + residual = hidden_states + hidden_states = self.final_layer_norm(op, hidden_states) + hidden_states = self.fc1(op, hidden_states) + hidden_states = op.Gelu(hidden_states) + hidden_states = self.fc2(op, hidden_states) + return op.Add(residual, hidden_states), self_kv + + +class NemotronParseDecoder(nn.Module): + """Position-free scaled-embedding mBART decoder used by Nemotron Parse.""" + + def __init__(self, config: ArchitectureConfig): + super().__init__() + self.config = config + self.embed_tokens = Embedding( + config.vocab_size, config.hidden_size, config.pad_token_id + ) + self.layernorm_embedding = LayerNorm(config.hidden_size, eps=1e-5) + self.layers = nn.ModuleList( + [_NemotronParseDecoderLayer(config) for _ in range(config.num_decoder_layers)] + ) + self.layer_norm = LayerNorm(config.hidden_size, eps=1e-5) + self.lm_head = Linear(config.hidden_size, config.vocab_size, bias=False) + self.lm_head.weight = self.embed_tokens.weight + self._embedding_scale = float(config.hidden_size**0.5) + + def forward( + self, + op: OpBuilder, + input_ids: ir.Value, + attention_mask: ir.Value, + encoder_hidden_states: ir.Value, + past_key_values: list[tuple[ir.Value, ir.Value]] | None = None, + ): + hidden_states = self.embed_tokens(op, input_ids) + embedding_scale = op.CastLike( + op.Constant(value_float=self._embedding_scale), hidden_states + ) + hidden_states = op.Mul( + hidden_states, + embedding_scale, + ) + hidden_states = self.layernorm_embedding(op, hidden_states) + self_attention_bias = create_padding_mask(op, input_ids, attention_mask) + + present_self_kvs = [] + layer_past = past_key_values or [None] * len(self.layers) + for layer, past_key_value in zip(self.layers, layer_past): + hidden_states, self_kv = layer( + op, + hidden_states, + encoder_hidden_states, + self_attention_bias, + past_key_value=past_key_value, + ) + present_self_kvs.append(self_kv) + + hidden_states = self.layer_norm(op, hidden_states) + logits = self.lm_head(op, hidden_states) + return logits, present_self_kvs + + def preprocess_weights( + self, state_dict: dict[str, torch.Tensor] + ) -> dict[str, torch.Tensor]: + """Select decoder weights and materialize the tied language-model head.""" + weights = { + key.removeprefix("decoder."): value + for key, value in state_dict.items() + if key.startswith("decoder.") + and ".extra_heads." not in key + and ".extra_proj." not in key + } + embed_weight = weights.get("embed_tokens.weight") + if embed_weight is not None: + weights["lm_head.weight"] = embed_weight + return weights + + +class NemotronParseForConditionalGeneration(nn.Module): + """NVIDIA Nemotron Parse OCR/document parser with C-RADIO and mBART.""" + + default_task = "vision-encoder-decoder" + category: str = "Multimodal" + + def __init__(self, config: ArchitectureConfig): + super().__init__() + self.config = config + self.vision_encoder = NemotronParseVisionEncoder(config) + self.decoder = NemotronParseDecoder(config) + + def preprocess_weights( + self, state_dict: dict[str, torch.Tensor] + ) -> dict[str, torch.Tensor]: + """Map the HuggingFace checkpoint to the two exported ONNX models.""" + weights = { + key: value + for key, value in state_dict.items() + if key.startswith("vision_encoder.") + } + weights.update( + { + f"vision_encoder.{key}": value + for key, value in self.vision_encoder.preprocess_weights(state_dict).items() + } + ) + weights.update( + { + f"decoder.{key}": value + for key, value in self.decoder.preprocess_weights(state_dict).items() + } + ) + return weights diff --git a/src/mobius/tasks/__init__.py b/src/mobius/tasks/__init__.py index d7f947ffc..df7a8f89b 100644 --- a/src/mobius/tasks/__init__.py +++ b/src/mobius/tasks/__init__.py @@ -67,6 +67,7 @@ "VAETask", "VideoDenoisingTask", "VisionLanguageTask", + "VisionEncoderDecoderTask", "WorldModelTask", "build_decoder_from_embeds", "build_embedding_from_features", @@ -120,6 +121,7 @@ from mobius.tasks._tts import TTSTask from mobius.tasks._vae import VAETask from mobius.tasks._video_denoising import VideoDenoisingTask +from mobius.tasks._vision_encoder_decoder import VisionEncoderDecoderTask from mobius.tasks._vision_language import Qwen3VLVisionLanguageTask from mobius.tasks._vision_language_3model import ( Cosmos3EdgeVLTask, @@ -160,6 +162,7 @@ "vae": VAETask, "qwen-image-vae": QwenImageVAETask, "vision-language": VisionLanguageTask, + "vision-encoder-decoder": VisionEncoderDecoderTask, "cosmos3-edge-vl": Cosmos3EdgeVLTask, "pixtral-vl": PixtralVLTask, "mllama-vision-language": MllamaVisionLanguageTask, diff --git a/src/mobius/tasks/_vision_encoder_decoder.py b/src/mobius/tasks/_vision_encoder_decoder.py new file mode 100644 index 000000000..c9ccab646 --- /dev/null +++ b/src/mobius/tasks/_vision_encoder_decoder.py @@ -0,0 +1,141 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Vision encoder-decoder task for image-to-text generation.""" + +from __future__ import annotations + +from typing import ClassVar + +import onnx_ir as ir +from onnxscript import nn + +from mobius._configs import ArchitectureConfig +from mobius._model_package import ModelPackage +from mobius.tasks._base import ComponentSpec, _make_graph, _make_model +from mobius.tasks._seq2seq import Seq2SeqTask + + +class VisionEncoderDecoderTask(Seq2SeqTask): + """Split an image encoder and cross-attentive text decoder into two models.""" + + model_roles: ClassVar[dict[str, str]] = { + "vision_encoder": "vision", + "decoder": "decoder", + } + components: ClassVar[ComponentSpec] = ComponentSpec( + vision_encoder="vision_encoder", + decoder="decoder", + ) + + def build( + self, + module: nn.Module, + config: ArchitectureConfig, + ) -> ModelPackage: + self._validate_components(module) + return ModelPackage( + { + "vision_encoder": self._build_vision_encoder_graph(module, config), + "decoder": self._build_decoder_graph(module, config), + }, + config=config, + ) + + def _build_vision_encoder_graph( + self, + module: nn.Module, + config: ArchitectureConfig, + ) -> ir.Model: + batch = ir.SymbolicDim("batch") + image_height = getattr(config, "image_height", None) or config.image_size + image_width = getattr(config, "image_width", None) or config.image_size + + graph, builder = _make_graph(name="vision_encoder") + pixel_values = builder.input( + "pixel_values", + dtype=ir.DataType.FLOAT, + shape=[batch, 3, image_height, image_width], + ) + # Image processors universally emit float32. Cast at the graph boundary + # so reduced-precision exports remain directly consumable. + model_pixel_values = builder.op.Cast(pixel_values, to=config.dtype) + encoder_hidden_states = module.vision_encoder( + builder.op, + pixel_values=model_pixel_values, + ) + builder.add_output(encoder_hidden_states, "last_hidden_state") + return _make_model(graph) + + def _build_decoder_graph( + self, + module: nn.Module, + config: ArchitectureConfig, + ) -> ir.Model: + """Build the image-to-text decoder with self-attention caches only. + + The compressed visual sequence remains constant during decoding. The + decoder projects it for cross-attention on each step, so exporting + unused cross-cache inputs would force callers to allocate invalid dummy + tensors and would expose outputs that cannot be consumed. + """ + batch = ir.SymbolicDim("batch") + dec_seq_len = ir.SymbolicDim("decoder_sequence_len") + enc_seq_len = ir.SymbolicDim("encoder_sequence_len") + past_seq_len = ir.SymbolicDim("past_sequence_len") + total_seq_len = ir.SymbolicDim("total_sequence_len") + + graph, builder = _make_graph(name="decoder") + input_ids = builder.input( + "input_ids", + dtype=ir.DataType.INT64, + shape=[batch, dec_seq_len], + ) + attention_mask = builder.input( + "attention_mask", + dtype=ir.DataType.INT64, + shape=[batch, total_seq_len], + ) + encoder_hidden_states = builder.input( + "encoder_hidden_states", + dtype=config.dtype, + shape=[batch, enc_seq_len, config.hidden_size], + ) + + past_self_kvs: list[tuple[ir.Value, ir.Value]] = [] + num_decoder_layers = getattr(config, "num_decoder_layers", config.num_hidden_layers) + for i in range(num_decoder_layers): + past_key = builder.input( + f"past_key_values.{i}.self.key", + dtype=config.dtype, + shape=[ + batch, + config.num_key_value_heads, + past_seq_len, + config.head_dim, + ], + ) + past_value = builder.input( + f"past_key_values.{i}.self.value", + dtype=config.dtype, + shape=[ + batch, + config.num_key_value_heads, + past_seq_len, + config.head_dim, + ], + ) + past_self_kvs.append((past_key, past_value)) + + logits, present_self_kvs = module.decoder( + builder.op, + input_ids=input_ids, + attention_mask=attention_mask, + encoder_hidden_states=encoder_hidden_states, + past_key_values=past_self_kvs, + ) + builder.add_output(logits, "logits") + for i, (key, value) in enumerate(present_self_kvs): + builder.add_output(key, f"present.{i}.self.key") + builder.add_output(value, f"present.{i}.self.value") + return _make_model(graph) diff --git a/testdata/cases/vision-language/nemotron-parse-2.yaml b/testdata/cases/vision-language/nemotron-parse-2.yaml new file mode 100644 index 000000000..150ab6bcb --- /dev/null +++ b/testdata/cases/vision-language/nemotron-parse-2.yaml @@ -0,0 +1,20 @@ +model_id: "nvidia/NVIDIA-Nemotron-Parse-2.0" +model_type: "nemotron_parse" +revision: "635b84d9b09bb9526b9a684d0b2c953d3cc3df05" +task_type: "image-to-text" +dtype: "bfloat16" +level: "L4+L5" +trust_remote_code: true + +inputs: + images: + - "nemotron-parse-document.png" + decoder_prompt: "" + +generation: + max_new_tokens: 24 + do_sample: false + eos_token_id: 2 + +ci_skip_reason: "The pinned 903M checkpoint and 2048x1664 C-RADIO input require a CUDA runner with at least 8 GB VRAM." +notes: "Real nonzero OCR document; BF16 CUDA L4 prefill and 24-token deterministic L5 generation." diff --git a/testdata/golden/vision-language/nemotron-parse-2.json b/testdata/golden/vision-language/nemotron-parse-2.json new file mode 100644 index 000000000..73286491d --- /dev/null +++ b/testdata/golden/vision-language/nemotron-parse-2.json @@ -0,0 +1,42 @@ +{ + "top1_id": 50251, + "top2_id": 50250, + "top10_ids": [ + 50251, + 50250, + 50252, + 50249, + 50247, + 50248, + 50253, + 50254, + 50211, + 50246 + ], + "top10_logits": [ + "0x1.2600000000000p+4", + "0x1.1600000000000p+4", + "0x1.b400000000000p+3", + "0x1.ae00000000000p+3", + "0x1.8200000000000p+3", + "0x1.7a00000000000p+3", + "0x1.7600000000000p+3", + "0x1.5e00000000000p+3", + "0x1.5200000000000p+3", + "0x1.4a00000000000p+3" + ], + "logits_summary": [ + "0x1.2600000000000p+4", + "-0x1.5800000000000p+4", + "-0x1.9b9863c8e660ap+2", + "0x1.984e048ac7d92p+1" + ], + "input_ids": [ + 2, + 0, + 50004, + 50008, + 50001, + 50010 + ] +} diff --git a/testdata/golden/vision-language/nemotron-parse-2_generation.json b/testdata/golden/vision-language/nemotron-parse-2_generation.json new file mode 100644 index 000000000..a9d48967d --- /dev/null +++ b/testdata/golden/vision-language/nemotron-parse-2_generation.json @@ -0,0 +1,31 @@ +{ + "model_id": "nvidia/NVIDIA-Nemotron-Parse-2.0", + "prompt": "", + "generated_tokens": [ + 50251, + 51312, + 25, + 7671, + 6134, + 3718, + 41839, + 31347, + 1778, + 5152, + 50630, + 51346, + 52316, + 221, + 221, + 50251, + 51361, + 68, + 390, + 26070, + 36360, + 243, + 40, + 2 + ], + "generated_text": "# MOBIUS OCR VALIDATION\n\nNemotron Parse 2" +} diff --git a/testdata/nemotron-parse-document.png b/testdata/nemotron-parse-document.png new file mode 100644 index 000000000..f71fba90c Binary files /dev/null and b/testdata/nemotron-parse-document.png differ diff --git a/tests/_test_configs.py b/tests/_test_configs.py index 19ea16689..c498af706 100644 --- a/tests/_test_configs.py +++ b/tests/_test_configs.py @@ -36,6 +36,7 @@ MllamaConfig, NanoChatConfig, NemotronHConfig, + NemotronParseConfig, Sam2Config, SegformerConfig, VisionConfig, @@ -2111,6 +2112,30 @@ def _base_config(config_cls=None, **overrides) -> ArchitectureConfig: # The test parametrization in build_graph_test.py uses specialised test # methods that invoke the correct task and assert the right output models. VL_CONFIGS: list[tuple[str, dict, bool]] = [ + # --- Nemotron Parse (C-RADIO + feature neck + cross-attentive decoder) --- + ( + "nemotron_parse", + { + "_config_cls": NemotronParseConfig, + "hidden_act": "gelu", + "num_decoder_layers": 1, + "num_key_value_heads": TINY_HEADS, + "image_height": 32, + "image_width": 64, + "vision_max_grid_size": 4, + "num_summary_tokens": 3, + "vision": VisionConfig( + hidden_size=32, + intermediate_size=64, + num_hidden_layers=1, + num_attention_heads=4, + image_size=64, + patch_size=16, + norm_eps=1e-6, + ), + }, + True, + ), # --- LLaVA family (vision-language, 3-model split) --- ("llava", {"vision": _TINY_VISION, "image_token_id": 32000}, True), ( diff --git a/tests/build_graph_test.py b/tests/build_graph_test.py index 8c986cdca..bd0788ece 100644 --- a/tests/build_graph_test.py +++ b/tests/build_graph_test.py @@ -5621,6 +5621,7 @@ def test_gemma4_static_cache_input_ordering(self): # VL models that produce a single "model" key instead of 3-model split _VL_SINGLE_MODEL_TASKS = {"qwen3-vl-vision-language"} +_VL_TWO_MODEL_TASKS = {"vision-encoder-decoder"} @pytest.mark.parametrize("model_type,config_overrides", _VL_MODEL_PARAMS) @@ -5642,6 +5643,14 @@ def test_package_builds(self, model_type: str, config_overrides: dict): assert model.graph is not None output_names = {o.name for o in model.graph.outputs} assert "logits" in output_names + elif task_name in _VL_TWO_MODEL_TASKS: + assert set(pkg) == {"decoder", "vision_encoder"} + decoder = pkg["decoder"] + assert "encoder_hidden_states" in {i.name for i in decoder.graph.inputs} + assert "logits" in {o.name for o in decoder.graph.outputs} + vision = pkg["vision_encoder"] + pixel_values = next(i for i in vision.graph.inputs if i.name == "pixel_values") + assert pixel_values.dtype == ir.DataType.FLOAT else: assert "decoder" in pkg, f"{model_type} should produce 'decoder'" assert "vision_encoder" in pkg, f"{model_type} should produce 'vision_encoder'" @@ -5652,7 +5661,8 @@ def test_package_builds(self, model_type: str, config_overrides: dict): assert "logits" in {o.name for o in decoder.graph.outputs} vision = pkg["vision_encoder"] - assert "pixel_values" in {i.name for i in vision.graph.inputs} + pixel_values = next(i for i in vision.graph.inputs if i.name == "pixel_values") + assert pixel_values.dtype == ir.DataType.FLOAT def test_has_initializers(self, model_type: str, config_overrides: dict): """Verify all sub-models have non-empty initializers.""" diff --git a/tests/e2e_golden_test.py b/tests/e2e_golden_test.py index 37ea7ca2b..22ca8fc4f 100644 --- a/tests/e2e_golden_test.py +++ b/tests/e2e_golden_test.py @@ -509,6 +509,53 @@ def _run_seq2seq_prefill( return outputs +def _prepare_image_to_text_inputs( + case: GoldenTestCase, +) -> tuple[np.ndarray, np.ndarray]: + """Apply the checkpoint's processor to one real image and decoder prompt.""" + import transformers + from PIL import Image + + processor = transformers.AutoProcessor.from_pretrained( + case.model_id, + revision=case.revision, + trust_remote_code=case.trust_remote_code, + ) + images = [Image.open(_TESTDATA_DIR / path).convert("RGB") for path in case.images] + processed = processor( + images=images, + text=case.decoder_prompt, + return_tensors="pt", + add_special_tokens=False, + ) + pixel_values = processed["pixel_values"].float().cpu().numpy() + return pixel_values, processed["input_ids"].cpu().numpy().astype(np.int64) + + +def _run_image_to_text_prefill( + pkg: ModelPackage, + case: GoldenTestCase, + config: object, +) -> dict[str, np.ndarray]: + """Run a real image through the vision encoder and decoder prefill.""" + pixel_values, input_ids = _prepare_image_to_text_inputs(case) + device_kwargs = _get_test_device_kwargs() + enc_session = OnnxModelSession(pkg["vision_encoder"], **device_kwargs) + dec_session = OnnxModelSession(pkg["decoder"], **device_kwargs) + try: + encoder_hidden = enc_session.run({"pixel_values": pixel_values})["last_hidden_state"] + feeds: dict[str, np.ndarray] = { + "input_ids": input_ids, + "attention_mask": np.ones_like(input_ids, dtype=np.int64), + "encoder_hidden_states": encoder_hidden, + } + feeds.update(_make_empty_kv_cache(dec_session, config)) + return dec_session.run(feeds) + finally: + enc_session.close() + dec_session.close() + + def _prepare_prefill_feeds( golden: GoldenRef, config: object, @@ -1860,6 +1907,8 @@ def test_prefill_argmax_matches_golden(self, case: GoldenTestCase) -> None: outputs = _run_speech_language_prefill(pkg, case, golden, config) elif case.task_type == "image-text-to-text": outputs = _run_vision_language_prefill(pkg, case, config) + elif case.task_type == "image-to-text": + outputs = _run_image_to_text_prefill(pkg, case, config) elif case.task_type == "phi4mm-multimodal": outputs = _run_phi4mm_multimodal_prefill(pkg, case, golden, config) elif case.task_type == "gemma4-assistant": @@ -1947,6 +1996,7 @@ def test_prefill_argmax_matches_golden(self, case: GoldenTestCase) -> None: { "text-generation", "image-text-to-text", + "image-to-text", "seq2seq", "speech-to-text", "speech-language", @@ -2164,6 +2214,54 @@ def _run_seq2seq_generation( return all_ids[0] +def _run_image_to_text_generation( + pkg: ModelPackage, + case: GoldenTestCase, + config: object, +) -> np.ndarray: + """Run greedy image-to-text generation through the two ONNX sessions.""" + pixel_values, input_ids = _prepare_image_to_text_inputs(case) + device_kwargs = _get_test_device_kwargs() + enc_session = OnnxModelSession(pkg["vision_encoder"], **device_kwargs) + dec_session = OnnxModelSession(pkg["decoder"], **device_kwargs) + try: + encoder_hidden = enc_session.run({"pixel_values": pixel_values})["last_hidden_state"] + past_cache = _make_empty_kv_cache(dec_session, config) + generated: list[np.ndarray] = [] + current_ids = input_ids + attention_mask = np.ones_like(input_ids, dtype=np.int64) + max_new_tokens = case.generation_params.get("max_new_tokens", 20) + eos_token_id = case.generation_params.get("eos_token_id") + for _ in range(max_new_tokens): + feeds: dict[str, np.ndarray] = { + "input_ids": current_ids, + "attention_mask": attention_mask, + "encoder_hidden_states": encoder_hidden, + **past_cache, + } + outputs = dec_session.run(feeds) + next_token = np.argmax( + outputs["logits"][:, -1], + axis=-1, + keepdims=True, + ).astype(np.int64) + generated.append(next_token) + for name in past_cache: + present_name = name.replace("past_key_values.", "present.") + past_cache[name] = outputs[present_name] + if eos_token_id is not None and np.all(next_token == eos_token_id): + break + current_ids = next_token + attention_mask = np.concatenate( + [attention_mask, np.ones_like(next_token, dtype=np.int64)], + axis=1, + ) + return np.concatenate(generated, axis=1)[0] + finally: + enc_session.close() + dec_session.close() + + def _run_speech_to_text_generation( pkg: ModelPackage, case: GoldenTestCase, @@ -2565,6 +2663,8 @@ def test_generation_matches_golden(self, case: GoldenTestCase) -> None: max_new_tokens=case.generation_params.get("max_new_tokens", 30), eos_token_id=case.generation_params.get("eos_token_id"), ) + elif case.task_type == "image-to-text": + new_tokens = _run_image_to_text_generation(pkg, case, config) elif case.task_type == "gemma4-assistant": new_tokens = _run_gemma4_assistant_generation(pkg, case) elif case.task_type == "seq2seq": diff --git a/tests/integration_test.py b/tests/integration_test.py index 99218321c..b22bbd0f8 100644 --- a/tests/integration_test.py +++ b/tests/integration_test.py @@ -21,6 +21,7 @@ from __future__ import annotations +import gc import os import numpy as np @@ -46,6 +47,112 @@ ) +@pytest.mark.integration +@pytest.mark.integration_slow +def test_nemotron_parse_real_weight_cuda_parity(): + """Compare real BF16 C-RADIO features and decoder logits on a document image.""" + if _get_test_device() != "cuda" or not torch.cuda.is_available(): + pytest.skip("Nemotron Parse real-weight parity requires CUDA") + + import ml_dtypes + from transformers import AutoModel, AutoProcessor + + model_id = "nvidia/NVIDIA-Nemotron-Parse-2.0" + revision = "635b84d9b09bb9526b9a684d0b2c953d3cc3df05" + prompt = "" + image_path = os.path.join( + os.path.dirname(__file__), + "..", + "testdata", + "nemotron-parse-document.png", + ) + processor = AutoProcessor.from_pretrained( + model_id, + revision=revision, + trust_remote_code=True, + ) + processed = processor( + images=[Image.open(image_path).convert("RGB")], + text=prompt, + return_tensors="pt", + add_special_tokens=False, + ) + hf_model = AutoModel.from_pretrained( + model_id, + revision=revision, + dtype=torch.bfloat16, + trust_remote_code=True, + ).to("cuda") + hf_model.eval() + pixel_values = processed["pixel_values"].to("cuda") + decoder_input_ids = processed["input_ids"].to("cuda") + with torch.no_grad(): + encoder_outputs = hf_model.encoder(pixel_values=pixel_values) + hf_encoder = encoder_outputs[0].float().cpu().numpy() + hf_logits = ( + hf_model( + encoder_outputs=encoder_outputs, + decoder_input_ids=decoder_input_ids, + ) + .logits[:, -1] + .float() + .cpu() + .numpy() + ) + del hf_model, encoder_outputs + gc.collect() + torch.cuda.empty_cache() + + pkg = build( + model_id, + dtype="bf16", + load_weights=True, + trust_remote_code=True, + execution_provider="cuda", + ) + vision_session = _make_session(pkg["vision_encoder"]) + decoder_session = _make_session(pkg["decoder"]) + try: + onnx_pixel_values = processed["pixel_values"].float().cpu().numpy() + onnx_encoder = vision_session.run({"pixel_values": onnx_pixel_values})[ + "last_hidden_state" + ] + empty_cache = { + name: np.zeros( + (1, pkg.config.num_key_value_heads, 0, pkg.config.head_dim), + dtype=ml_dtypes.bfloat16, + ) + for name in decoder_session.input_names + if name.startswith("past_key_values.") + } + onnx_logits = decoder_session.run( + { + "input_ids": processed["input_ids"].numpy().astype(np.int64), + "attention_mask": np.ones_like(processed["input_ids"].numpy(), dtype=np.int64), + "encoder_hidden_states": onnx_encoder, + **empty_cache, + } + )["logits"][:, -1] + finally: + vision_session.close() + decoder_session.close() + + onnx_encoder_f32 = onnx_encoder.astype(np.float32) + encoder_cosine = np.dot(onnx_encoder_f32.ravel(), hf_encoder.ravel()) / ( + np.linalg.norm(onnx_encoder_f32) * np.linalg.norm(hf_encoder) + ) + onnx_logits_f32 = onnx_logits.astype(np.float32) + logits_cosine = np.dot(onnx_logits_f32.ravel(), hf_logits.ravel()) / ( + np.linalg.norm(onnx_logits_f32) * np.linalg.norm(hf_logits) + ) + assert encoder_cosine > 0.99 + assert logits_cosine > 0.995 + np.testing.assert_array_equal( + np.argmax(onnx_logits_f32, axis=-1), + np.argmax(hf_logits, axis=-1), + ) + + def _get_test_device() -> str: """Return 'cuda' if MOBIUS_TEST_DEVICE=cuda, else 'cpu'.""" return os.environ.get("MOBIUS_TEST_DEVICE", "cpu").strip().lower() diff --git a/tests/synthetic_parity_test.py b/tests/synthetic_parity_test.py index 306406edc..ca1e37b41 100644 --- a/tests/synthetic_parity_test.py +++ b/tests/synthetic_parity_test.py @@ -707,6 +707,312 @@ def _fill_random_weights(model: ir.Model, rng: np.random.Generator) -> None: init.const_value = ir.Tensor(data) +def _nemotron_parse_torch_attention( + hidden_states: torch.Tensor, + weights: dict[str, torch.Tensor], + prefix: str, + *, + num_heads: int, + key_value_states: torch.Tensor | None = None, + causal: bool = False, +) -> torch.Tensor: + """Evaluate one Nemotron Parse attention block with exported weights.""" + source = hidden_states if key_value_states is None else key_value_states + query = torch.nn.functional.linear( + hidden_states, + weights[f"{prefix}.q_proj.weight"], + weights[f"{prefix}.q_proj.bias"], + ) + key = torch.nn.functional.linear( + source, + weights[f"{prefix}.k_proj.weight"], + weights[f"{prefix}.k_proj.bias"], + ) + value = torch.nn.functional.linear( + source, + weights[f"{prefix}.v_proj.weight"], + weights[f"{prefix}.v_proj.bias"], + ) + batch, query_len, hidden_size = query.shape + key_len = key.shape[1] + head_dim = hidden_size // num_heads + query = query.reshape(batch, query_len, num_heads, head_dim).transpose(1, 2) + key = key.reshape(batch, key_len, num_heads, head_dim).transpose(1, 2) + value = value.reshape(batch, key_len, num_heads, head_dim).transpose(1, 2) + scores = query @ key.transpose(-1, -2) * head_dim**-0.5 + if causal: + mask = torch.ones(query_len, key_len, dtype=torch.bool).tril() + scores = scores.masked_fill(~mask, float("-inf")) + output = torch.softmax(scores, dim=-1) @ value + output = output.transpose(1, 2).reshape(batch, query_len, hidden_size) + return torch.nn.functional.linear( + output, + weights[f"{prefix}.out_proj.weight"], + weights[f"{prefix}.out_proj.bias"], + ) + + +def _nemotron_parse_torch_vision( + pixel_values: torch.Tensor, + weights: dict[str, torch.Tensor], +) -> torch.Tensor: + """Evaluate the tiny one-layer C-RADIO encoder and compression neck.""" + prefix = "vision_encoder.model_encoder.radio_model.model" + patch_weight = weights[f"{prefix}.patch_generator.embedder.weight"].reshape(32, 3, 16, 16) + hidden = torch.nn.functional.conv2d(pixel_values, patch_weight, stride=16) + hidden = hidden.flatten(2).transpose(1, 2) + pos = weights[f"{prefix}.patch_generator.pos_embed"].reshape(1, 4, 4, 32) + hidden = hidden + pos[:, :2].reshape(1, 8, 32) + cls = weights[f"{prefix}.patch_generator.cls_token.token"].unsqueeze(0) + hidden = torch.cat((cls, hidden), dim=1) + + block = f"{prefix}.blocks.0" + norm = torch.nn.functional.layer_norm( + hidden, + (32,), + weights[f"{block}.norm1.weight"], + weights[f"{block}.norm1.bias"], + 1e-6, + ) + qkv = torch.nn.functional.linear( + norm, + weights[f"{block}.attn.qkv.weight"], + weights[f"{block}.attn.qkv.bias"], + ) + query, key, value = qkv.chunk(3, dim=-1) + batch, sequence, _ = query.shape + query = query.reshape(batch, sequence, 4, 8).transpose(1, 2) + key = key.reshape(batch, sequence, 4, 8).transpose(1, 2) + value = value.reshape(batch, sequence, 4, 8).transpose(1, 2) + attn = torch.softmax(query @ key.transpose(-1, -2) * 8**-0.5, dim=-1) @ value + attn = attn.transpose(1, 2).reshape(batch, sequence, 32) + attn = torch.nn.functional.linear( + attn, + weights[f"{block}.attn.proj.weight"], + weights[f"{block}.attn.proj.bias"], + ) + hidden = hidden + attn + norm = torch.nn.functional.layer_norm( + hidden, + (32,), + weights[f"{block}.norm2.weight"], + weights[f"{block}.norm2.bias"], + 1e-6, + ) + mlp = torch.nn.functional.gelu( + torch.nn.functional.linear( + norm, + weights[f"{block}.mlp.fc1.weight"], + weights[f"{block}.mlp.fc1.bias"], + ) + ) + hidden = hidden + torch.nn.functional.linear( + mlp, + weights[f"{block}.mlp.fc2.weight"], + weights[f"{block}.mlp.fc2.bias"], + ) + + summary = hidden[:, [0, 1, 2]].reshape(batch, -1) + features = hidden[:, 8:] + features = torch.nn.functional.linear( + features, + weights["vision_encoder.conv1.weight"], + weights["vision_encoder.conv1.bias"], + ) + features = torch.nn.functional.layer_norm( + features, + (64,), + weights["vision_encoder.layer_norm1.weight"], + weights["vision_encoder.layer_norm1.bias"], + 1e-6, + ) + features = features.reshape(batch, 2, 4, 64).permute(0, 3, 1, 2) + features = torch.nn.functional.conv2d( + features, + weights["vision_encoder.conv2.weight"], + stride=(1, 4), + ) + features = features.permute(0, 2, 3, 1).reshape(batch, 2, 64) + features = torch.nn.functional.layer_norm( + features, + (64,), + weights["vision_encoder.layer_norm2.weight"], + weights["vision_encoder.layer_norm2.bias"], + 1e-6, + ) + summary = torch.nn.functional.linear( + summary, + weights["vision_encoder.sum_proj.weight"], + weights["vision_encoder.sum_proj.bias"], + ) + summary = torch.nn.functional.layer_norm( + summary, + (64,), + weights["vision_encoder.layer_norm3.weight"], + weights["vision_encoder.layer_norm3.bias"], + 1e-6, + ) + return torch.cat((features, summary[:, None]), dim=1) + + +def _nemotron_parse_torch_decoder( + input_ids: torch.Tensor, + encoder_hidden_states: torch.Tensor, + weights: dict[str, torch.Tensor], +) -> torch.Tensor: + """Evaluate the tiny one-layer pre-norm mBART decoder.""" + hidden = torch.nn.functional.embedding(input_ids, weights["decoder.embed_tokens.weight"]) + hidden = hidden * 8.0 + hidden = torch.nn.functional.layer_norm( + hidden, + (64,), + weights["decoder.layernorm_embedding.weight"], + weights["decoder.layernorm_embedding.bias"], + ) + layer = "decoder.layers.0" + residual = hidden + norm = torch.nn.functional.layer_norm( + hidden, + (64,), + weights[f"{layer}.self_attn_layer_norm.weight"], + weights[f"{layer}.self_attn_layer_norm.bias"], + ) + hidden = residual + _nemotron_parse_torch_attention( + norm, + weights, + f"{layer}.self_attn", + num_heads=4, + causal=True, + ) + residual = hidden + norm = torch.nn.functional.layer_norm( + hidden, + (64,), + weights[f"{layer}.encoder_attn_layer_norm.weight"], + weights[f"{layer}.encoder_attn_layer_norm.bias"], + ) + hidden = residual + _nemotron_parse_torch_attention( + norm, + weights, + f"{layer}.encoder_attn", + num_heads=4, + key_value_states=encoder_hidden_states, + ) + residual = hidden + hidden = torch.nn.functional.layer_norm( + hidden, + (64,), + weights[f"{layer}.final_layer_norm.weight"], + weights[f"{layer}.final_layer_norm.bias"], + ) + hidden = torch.nn.functional.gelu( + torch.nn.functional.linear( + hidden, + weights[f"{layer}.fc1.weight"], + weights[f"{layer}.fc1.bias"], + ) + ) + hidden = residual + torch.nn.functional.linear( + hidden, + weights[f"{layer}.fc2.weight"], + weights[f"{layer}.fc2.bias"], + ) + hidden = torch.nn.functional.layer_norm( + hidden, + (64,), + weights["decoder.layer_norm.weight"], + weights["decoder.layer_norm.bias"], + ) + return hidden @ weights["decoder.embed_tokens.weight"].T + + +def test_nemotron_parse_synthetic_parity(): + """L3 parity for the full tiny vision encoder and cross-attentive decoder.""" + from _test_configs import VL_CONFIGS + + from mobius._testing.ort_inference import OnnxModelSession + + overrides = next( + overrides for model_type, overrides, _ in VL_CONFIGS if model_type == "nemotron_parse" + ) + config = _base_config(**overrides) + _, pkg = _build_onnx_model("nemotron_parse", config) + decoder_layer_norms = [ + node for node in pkg["decoder"].graph if node.op_type == "LayerNormalization" + ] + assert len(decoder_layer_norms) == 3 * config.num_decoder_layers + 2 + assert all( + node.attributes["epsilon"].value == pytest.approx(1e-5) for node in decoder_layer_norms + ) + rng = np.random.default_rng(42) + for model in pkg.values(): + _fill_random_weights(model, rng) + + weights: dict[str, torch.Tensor] = {} + for model in pkg.values(): + for name, initializer in model.graph.initializers.items(): + if initializer.const_value is not None and not name.startswith("const_"): + weights[name] = torch.from_numpy(initializer.const_value.numpy()) + + pixel_values = rng.standard_normal((1, 3, 32, 64)).astype(np.float32) + input_ids = np.array([[2, 7, 11]], dtype=np.int64) + torch_encoder = _nemotron_parse_torch_vision(torch.from_numpy(pixel_values), weights) + torch_logits = _nemotron_parse_torch_decoder( + torch.from_numpy(input_ids), + torch_encoder, + weights, + ) + + vision_session = OnnxModelSession(pkg["vision_encoder"]) + decoder_session = OnnxModelSession(pkg["decoder"]) + try: + onnx_encoder = vision_session.run({"pixel_values": pixel_values})["last_hidden_state"] + empty_cache = { + name: np.zeros((1, 4, 0, 16), dtype=np.float32) + for name in decoder_session.input_names + if name.startswith("past_key_values.") + } + onnx_logits = decoder_session.run( + { + "input_ids": input_ids, + "attention_mask": np.ones_like(input_ids, dtype=np.int64), + "encoder_hidden_states": onnx_encoder, + **empty_cache, + } + )["logits"] + + unpadded_ids = np.array([[2, 7]], dtype=np.int64) + unpadded_logits = decoder_session.run( + { + "input_ids": unpadded_ids, + "attention_mask": np.ones_like(unpadded_ids, dtype=np.int64), + "encoder_hidden_states": onnx_encoder, + **empty_cache, + } + )["logits"] + padded_ids = np.array([[1, 1, 2, 7]], dtype=np.int64) + padded_logits = decoder_session.run( + { + "input_ids": padded_ids, + "attention_mask": np.array([[0, 0, 1, 1]], dtype=np.int64), + "encoder_hidden_states": onnx_encoder, + **empty_cache, + } + )["logits"] + finally: + vision_session.close() + decoder_session.close() + + np.testing.assert_allclose(onnx_encoder, torch_encoder.numpy(), rtol=1e-3, atol=1e-3) + np.testing.assert_allclose(onnx_logits, torch_logits.numpy(), rtol=1e-3, atol=1e-3) + np.testing.assert_allclose( + unpadded_logits[:, -1], + padded_logits[:, -1], + rtol=1e-5, + atol=1e-5, + ) + + # --------------------------------------------------------------------------- # Parametrize over all causal LM configs # --------------------------------------------------------------------------- diff --git a/tests/weight_alignment_test.py b/tests/weight_alignment_test.py index 39ae516a2..4e412dea6 100644 --- a/tests/weight_alignment_test.py +++ b/tests/weight_alignment_test.py @@ -33,6 +33,7 @@ ENCODER_CONFIGS, SEQ2SEQ_CONFIGS, VISION_CONFIGS, + VL_CONFIGS, _base_config, ) @@ -221,6 +222,23 @@ def test_identity_state_dict_roundtrip(self, model_type: str, config_overrides: _assert_identity_roundtrip(model_type, config_overrides) +@pytest.mark.parametrize( + "model_type,config_overrides", + [ + pytest.param(model_type, overrides, id=model_type) + for model_type, overrides, _ in VL_CONFIGS + if model_type == "nemotron_parse" + ], +) +class TestNemotronParseWeightAlignment: + """Verify the two-model package preserves every aligned parameter.""" + + def test_identity_state_dict_roundtrip( + self, model_type: str, config_overrides: dict + ) -> None: + _assert_identity_roundtrip(model_type, config_overrides) + + # --------------------------------------------------------------------------- # Detection model weight alignment # ---------------------------------------------------------------------------