Skip to content

Commit 5ca65be

Browse files
authored
Complete Qwen2.5-Omni Thinker export
Add a dedicated four-model task and align audio, vision, decoder, configuration, and weight routing with the current Hugging Face Thinker implementation. Cover the nested config and complete graph package with focused regression tests. Signed-off-by: GitHub <noreply@github.com>
1 parent a6303a7 commit 5ca65be

15 files changed

Lines changed: 581 additions & 75 deletions

src/mobius/_configs/_base.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -418,6 +418,7 @@ class ArchitectureConfig(BaseModelConfig):
418418
# Vision shared fields (accessed as top-level config.X by tasks)
419419
mm_tokens_per_image: int | None = None
420420
image_token_id: int | None = None
421+
video_token_id: int | None = None
421422
spatial_merge_size: int = 2
422423
temporal_patch_size: int = 2
423424
deepstack_visual_indexes: list[int] | None = None
@@ -649,6 +650,7 @@ def from_transformers(cls, config, parent_config=None) -> ArchitectureConfig:
649650
"bloom",
650651
"qwen2",
651652
"qwen2_5_vl_text",
653+
"qwen2_5_omni_text",
652654
"qwen2_moe",
653655
"qwen2_vl_text",
654656
),

src/mobius/_configs/_extractors.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -156,7 +156,7 @@ def extract_vision_config(config, parent_config, model_type: str) -> dict:
156156
Either step can populate ``fields`` (which become kwargs for
157157
:class:`VisionConfig`), or a per-model hook can return a fully-formed
158158
dict to short-circuit. The dispatcher also lifts a fixed set of
159-
"shared" vision fields (``image_token_id``, ``spatial_merge_size``,
159+
"shared" vision fields (``image_token_id``, ``video_token_id``, ``spatial_merge_size``,
160160
...) up to the top-level of the returned dict so callers can access
161161
them as ``config.image_token_id`` directly.
162162
"""
@@ -182,6 +182,7 @@ def extract_vision_config(config, parent_config, model_type: str) -> dict:
182182
for shared in (
183183
"mm_tokens_per_image",
184184
"image_token_id",
185+
"video_token_id",
185186
"spatial_merge_size",
186187
"temporal_patch_size",
187188
"deepstack_visual_indexes",

src/mobius/_configs/_sub_configs.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,7 @@ class VisionConfig:
4747
norm_eps: float = 1e-6
4848
mm_tokens_per_image: int | None = None
4949
image_token_id: int | None = None
50+
video_token_id: int | None = None
5051
# Pixtral / Mistral-3 vision fields
5152
model_type: str | None = None
5253
head_dim: int | None = None

src/mobius/_configs/per_model/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,5 +31,6 @@
3131
_phi4mm_audio,
3232
_phi4mm_vision,
3333
_qwen3_asr_audio,
34+
_qwen25_omni_vision,
3435
_sensevoice_audio,
3536
)
Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,39 @@
1+
# Copyright (c) Microsoft Corporation.
2+
# Licensed under the MIT License.
3+
4+
"""Qwen2.5-Omni vision extractor (vision config lives under thinker_config)."""
5+
6+
from __future__ import annotations
7+
8+
from mobius._configs._extractors import register_vision_hook
9+
10+
11+
@register_vision_hook("qwen2_5_omni_text")
12+
def _qwen25_omni_vision(config, parent_config, model_type: str, fields: dict):
13+
thinker = getattr(parent_config, "thinker_config", None)
14+
if thinker is None:
15+
return None
16+
if isinstance(thinker, dict):
17+
thinker = type("ThinkerConfig", (), thinker)()
18+
vision = getattr(thinker, "vision_config", None)
19+
if vision is None:
20+
return None
21+
if isinstance(vision, dict):
22+
vision = type("VisionConfig", (), vision)()
23+
24+
fields.update(
25+
hidden_size=getattr(vision, "hidden_size", None),
26+
intermediate_size=getattr(vision, "intermediate_size", None),
27+
num_hidden_layers=getattr(vision, "depth", None),
28+
num_attention_heads=getattr(vision, "num_heads", None),
29+
patch_size=getattr(vision, "patch_size", None),
30+
out_hidden_size=getattr(vision, "out_hidden_size", None),
31+
in_channels=getattr(vision, "in_channels", 3),
32+
spatial_merge_size=getattr(vision, "spatial_merge_size", 2),
33+
temporal_patch_size=getattr(vision, "temporal_patch_size", 2),
34+
fullatt_block_indexes=getattr(vision, "fullatt_block_indexes", None),
35+
window_size=getattr(vision, "window_size", 112),
36+
image_token_id=getattr(thinker, "image_token_id", None),
37+
)
38+
fields["video_token_id"] = getattr(thinker, "video_token_id", None)
39+
return None

src/mobius/_registry.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -626,11 +626,10 @@ def _detect_fallback_registration(hf_config) -> ModelRegistration | None:
626626
task="speech-to-text",
627627
config_class=WhisperConfig,
628628
),
629-
630629
# --- Omni ---
631630
"qwen2_5_omni": ModelRegistration(
632631
Qwen25OmniThinkerForConditionalGeneration,
633-
task="speech-language",
632+
task="qwen25-omni",
634633
),
635634
# --- Encoder-only ---
636635
"albert": ModelRegistration(BertModel, task="feature-extraction"),

src/mobius/components/__init__.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -54,6 +54,8 @@
5454
"PostNormDecoderLayer",
5555
"QuantizedEmbedding",
5656
"QuantizedLinear",
57+
"Qwen25OmniAudioAttention",
58+
"Qwen25OmniAudioEncoderLayer",
5759
"RMSNorm",
5860
"SelectiveScan",
5961
"SiLU",
@@ -221,6 +223,10 @@
221223
from mobius.components._qwen3_vl_vision import (
222224
Qwen3VLVisionRotaryEmbedding as Qwen3VLVisionRotaryEmbedding,
223225
)
226+
from mobius.components._qwen25_omni_audio import (
227+
Qwen25OmniAudioAttention,
228+
Qwen25OmniAudioEncoderLayer,
229+
)
224230
from mobius.components._qwen25_vl_vision import (
225231
Qwen2VLVisionBlock as Qwen2VLVisionBlock,
226232
)

src/mobius/components/_conv.py

Lines changed: 3 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -5,14 +5,9 @@
55

66
from __future__ import annotations
77

8-
from typing import TYPE_CHECKING
9-
108
import onnx_ir as ir
119
from onnxscript import OpBuilder, nn
1210

13-
if TYPE_CHECKING:
14-
pass
15-
1611

1712
class Conv2d(nn.Module):
1813
"""2D convolution with bias.
@@ -70,16 +65,14 @@ def __init__(
7065
groups: int = 1,
7166
):
7267
super().__init__()
73-
self.weight = nn.Parameter(
74-
(out_channels, in_channels // groups, kernel_size)
75-
)
76-
self.bias = nn.Parameter((out_channels))
68+
self.weight = nn.Parameter((out_channels, in_channels // groups, kernel_size))
69+
self.bias = nn.Parameter((out_channels,))
7770
self._kernel_size = kernel_size
7871
self._stride = stride
7972
self._padding = padding
8073
self._groups = groups
8174

82-
def forward(self, op: builder.OpBuilder, x: ir.Value):
75+
def forward(self, op: OpBuilder, x: ir.Value):
8376
p = self._padding
8477
return op.Conv(
8578
x,

src/mobius/components/_qwen25_omni_audio.py

Lines changed: 64 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -1,25 +1,23 @@
1-
"""Qwen25-Omni audio encoder components.
1+
# Copyright (c) Microsoft Corporation.
2+
# Licensed under the MIT License.
23

3-
Whisper-inspired audio encoder with 3x Conv1d,
4-
sinusoidal positional embeddings, and bidirectional transformer
5-
encoder layers with LayerNorm.
4+
"""Qwen2.5-Omni audio encoder components.
5+
6+
Packed bidirectional transformer layers with LayerNorm.
67
78
Reference: Transformers
89
https://github.com/huggingface/transformers/blob/main/src/transformers/models/qwen2_5_omni/modeling_qwen2_5_omni.py
910
"""
1011

1112
from __future__ import annotations
1213

13-
from typing import TYPE_CHECKING
14-
14+
import onnx_ir as ir
1515
from onnxscript import nn
1616
from onnxscript._internal import builder
1717

18+
from mobius._build_context import get_build_dtype
1819
from mobius.components._common import LayerNorm, Linear
1920

20-
if TYPE_CHECKING:
21-
import onnx_ir as ir
22-
2321

2422
class Qwen25OmniAudioAttention(nn.Module):
2523
"""Bidirectional multi-head attention for Qwen2_5Omni audio encoder.
@@ -37,7 +35,12 @@ def __init__(self, d_model: int, num_heads: int):
3735
self._num_heads = num_heads
3836
self._head_dim = d_model // num_heads
3937

40-
def forward(self, op: builder.OpBuilder, hidden_states: ir.Value):
38+
def forward(
39+
self,
40+
op: builder.OpBuilder,
41+
hidden_states: ir.Value,
42+
cu_seqlens: ir.Value,
43+
):
4144
"""Bidirectional self-attention.
4245
4346
Args:
@@ -46,19 +49,53 @@ def forward(self, op: builder.OpBuilder, hidden_states: ir.Value):
4649
Returns:
4750
output: (batch, seq_len, d_model)
4851
"""
49-
q = self.q_proj(op, hidden_states)
50-
k = self.k_proj(op, hidden_states)
51-
v = self.v_proj(op, hidden_states)
52+
seq_len = op.Shape(hidden_states, start=0, end=1)
53+
packed_shape = op.Concat(seq_len, [self._num_heads, self._head_dim], axis=0)
54+
q = op.Reshape(self.q_proj(op, hidden_states), packed_shape)
55+
k = op.Reshape(self.k_proj(op, hidden_states), packed_shape)
56+
v = op.Reshape(self.v_proj(op, hidden_states), packed_shape)
57+
58+
# Build the block-diagonal mask represented by HF's cu_seqlens.
59+
positions = op.Range(0, op.Squeeze(seq_len, [0]), 1)
60+
segment_ids = op.Sub(
61+
op.ReduceSum(
62+
op.Cast(
63+
op.GreaterOrEqual(
64+
op.Unsqueeze(positions, [1]),
65+
op.Unsqueeze(op.Cast(cu_seqlens, to=7), [0]),
66+
),
67+
to=7,
68+
),
69+
[1],
70+
keepdims=False,
71+
),
72+
1,
73+
)
74+
same_segment = op.Equal(
75+
op.Unsqueeze(segment_ids, [1]),
76+
op.Unsqueeze(segment_ids, [0]),
77+
)
78+
attention_bias = op.Where(
79+
same_segment,
80+
op.CastLike(0.0, q),
81+
op.CastLike(-1e9, q),
82+
)
83+
attention_bias = op.Unsqueeze(attention_bias, [0, 1])
5284

53-
# Use ONNX Attention op (bidirectional: no causal mask)
85+
q = op.Unsqueeze(op.Transpose(q, perm=[1, 0, 2]), [0])
86+
k = op.Unsqueeze(op.Transpose(k, perm=[1, 0, 2]), [0])
87+
v = op.Unsqueeze(op.Transpose(v, perm=[1, 0, 2]), [0])
5488
attn_output = op.Attention(
5589
q,
5690
k,
5791
v,
92+
attention_bias,
5893
q_num_heads=self._num_heads,
5994
kv_num_heads=self._num_heads,
6095
scale=float(self._head_dim**-0.5),
6196
)
97+
attn_output = op.Transpose(op.Squeeze(attn_output, [0]), perm=[1, 0, 2])
98+
attn_output = op.Reshape(attn_output, op.Concat(seq_len, [-1], axis=0))
6299
return self.out_proj(op, attn_output)
63100

64101

@@ -86,7 +123,12 @@ def __init__(
86123
self.fc2 = Linear(ffn_dim, d_model, bias=True)
87124
self.final_layer_norm = LayerNorm(d_model, eps=eps)
88125

89-
def forward(self, op: builder.OpBuilder, hidden_states: ir.Value):
126+
def forward(
127+
self,
128+
op: builder.OpBuilder,
129+
hidden_states: ir.Value,
130+
cu_seqlens: ir.Value,
131+
):
90132
"""Pre-norm encoder layer with bidirectional attention.
91133
92134
Args:
@@ -98,7 +140,7 @@ def forward(self, op: builder.OpBuilder, hidden_states: ir.Value):
98140
# Self-attention with pre-norm and residual
99141
residual = hidden_states
100142
hidden_states = self.self_attn_layer_norm(op, hidden_states)
101-
hidden_states = self.self_attn(op, hidden_states)
143+
hidden_states = self.self_attn(op, hidden_states, cu_seqlens)
102144
hidden_states = op.Add(residual, hidden_states)
103145

104146
# FFN with pre-norm, GELU, and residual
@@ -108,5 +150,11 @@ def forward(self, op: builder.OpBuilder, hidden_states: ir.Value):
108150
hidden_states = op.Gelu(hidden_states)
109151
hidden_states = self.fc2(op, hidden_states)
110152
hidden_states = op.Add(residual, hidden_states)
153+
if get_build_dtype() == ir.DataType.FLOAT16:
154+
hidden_states = op.Clip(
155+
hidden_states,
156+
op.CastLike(-64504.0, hidden_states),
157+
op.CastLike(64504.0, hidden_states),
158+
)
111159

112160
return hidden_states

src/mobius/models/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -93,6 +93,7 @@
9393
"Phi4MMMultiModalModel",
9494
"PhiCausalLMModel",
9595
"Qwen25VLCausalLMModel",
96+
"Qwen25OmniThinkerForConditionalGeneration",
9697
"Qwen25VLDecoderModel",
9798
"Qwen25VLEmbeddingModel",
9899
"Qwen25VLTextModel",
@@ -250,6 +251,7 @@
250251
Qwen3TTSCodecEncoderModel,
251252
Qwen3TTSTokenizerV2Model,
252253
)
254+
from mobius.models.qwen25_omni import Qwen25OmniThinkerForConditionalGeneration
253255
from mobius.models.qwen35 import (
254256
Qwen35CausalLMModel,
255257
Qwen35MoECausalLMModel,

0 commit comments

Comments
 (0)