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
# ---------------------------------------------------------------------------