Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 2 additions & 4 deletions src/mobius/components/_audio.py
Original file line number Diff line number Diff line change
Expand Up @@ -409,10 +409,8 @@ def __init__(
self.layer_norm = LayerNorm(d_model)

def forward(self, op: builder.OpBuilder, x: ir.Value, relative_attention_bias: ir.Value):
half = op.Constant(value_float=0.5)

# Macaron feed-forward in
x = op.Add(x, op.Mul(self.feed_forward_in(op, x), half))
x = op.Add(x, op.Mul(self.feed_forward_in(op, x), 0.5))

# Multi-head attention with pre-norm
norm_x = self.layer_norm_att(op, x)
Expand All @@ -422,7 +420,7 @@ def forward(self, op: builder.OpBuilder, x: ir.Value, relative_attention_bias: i
x = op.Add(x, self.conv(op, x))

# Macaron feed-forward out
x = op.Add(x, op.Mul(self.feed_forward_out(op, x), half))
x = op.Add(x, op.Mul(self.feed_forward_out(op, x), 0.5))

return self.layer_norm(op, x)

Expand Down
219 changes: 219 additions & 0 deletions src/mobius/components/_castlike_dtype_test.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,219 @@
# Copyright (c) ONNX Project Contributors
# SPDX-License-Identifier: Apache-2.0

"""Tests that Python float literals auto-cast to match operand dtypes.

onnxscript auto-casts Python scalars (int/float/bool) to match the dtype of
the other operand in a binary op. These tests verify that components using
Python literals produce output in the correct dtype — no spurious FLOAT32
constants widening BF16/FP16 computations.
"""

from __future__ import annotations

import onnx_ir as ir
import pytest

from mobius._builder import _cast_module_dtype
from mobius._configs import Gemma3nConfig
from mobius._testing import count_op_type, create_test_builder, create_test_input
from mobius.components._activations import quick_gelu
from mobius.components._audio import ConformerEncoderLayer
from mobius.components._diffusion import AdaLayerNormOutput
from mobius.components._moe import SigmoidTopKGate, SparseMixerGate
from mobius.components._rms_norm import OffsetRMSNorm
from mobius.models.gemma3n import Gemma3nAltUp


def _get_output_dtype(graph: ir.Graph) -> ir.DataType | None:
"""Return the dtype of the first graph output, or None."""
if graph.outputs:
return graph.outputs[0].dtype
return None


class TestPythonLiteralAutocast:
"""Verify that Python float literals auto-cast to match tensor operand dtypes."""

@pytest.mark.parametrize("dtype", [ir.DataType.FLOAT16, ir.DataType.BFLOAT16])
def test_offset_rms_norm_constant_autocasts(self, dtype: ir.DataType):
"""OffsetRMSNorm: `1.0` literal in op.Add auto-casts to weight dtype.

After _cast_module_dtype, the weight param is BF16/FP16. The Python
literal 1.0 in op.Add(self.weight, 1.0) must auto-cast to match.
"""
norm = OffsetRMSNorm(hidden_size=64, eps=1e-6)
_cast_module_dtype(norm, dtype)
builder_, op, graph = create_test_builder()
x = create_test_input(builder_, "x", [2, 3, 64], dtype)
result = norm(op, x)
graph.outputs.append(result)
# No CastLike — auto-cast handles it; output stays in the expected dtype
assert count_op_type(graph, "CastLike") == 0
assert _get_output_dtype(graph) == dtype

@pytest.mark.parametrize("dtype", [ir.DataType.FLOAT16, ir.DataType.BFLOAT16])
def test_quick_gelu_constant_autocasts(self, dtype: ir.DataType):
"""quick_gelu: `1.702` literal in op.Mul auto-casts to input dtype."""
builder_, op, graph = create_test_builder()
x = create_test_input(builder_, "x", [2, 3, 64], dtype)
result = quick_gelu(op, x)
graph.outputs.append(result)
assert count_op_type(graph, "CastLike") == 0
assert _get_output_dtype(graph) == dtype

@pytest.mark.parametrize("dtype", [ir.DataType.FLOAT16, ir.DataType.BFLOAT16])
def test_ada_layer_norm_output_constant_autocasts(self, dtype: ir.DataType):
"""AdaLayerNormOutput: `1.0` literal auto-casts to scale tensor dtype."""
mod = AdaLayerNormOutput(hidden_size=64, eps=1e-6)
_cast_module_dtype(mod, dtype)
builder_, op, graph = create_test_builder()
hidden = create_test_input(builder_, "hidden", [1, 4, 64], dtype)
temb = create_test_input(builder_, "temb", [1, 64], dtype)
result = mod(op, hidden, temb)
graph.outputs.append(result)
assert count_op_type(graph, "CastLike") == 0
assert _get_output_dtype(graph) == dtype

def test_float32_inputs_produce_float32_output(self):
"""Float32 inputs — no special casting needed, output stays float32."""
norm = OffsetRMSNorm(hidden_size=64, eps=1e-6)
builder_, op, graph = create_test_builder()
x = create_test_input(builder_, "x", [2, 3, 64], ir.DataType.FLOAT)
result = norm(op, x)
graph.outputs.append(result)
assert count_op_type(graph, "CastLike") == 0
assert _get_output_dtype(graph) == ir.DataType.FLOAT


class TestSigmoidTopKGate:
"""SigmoidTopKGate: verify 1e-9 epsilon and routed_scaling_factor auto-cast."""

@pytest.mark.parametrize("dtype", [ir.DataType.FLOAT16, ir.DataType.BFLOAT16])
def test_routing_weights_dtype(self, dtype: ir.DataType):
"""Routing weights stay in input dtype — no FP32 widening from 1e-9 literal."""
gate = SigmoidTopKGate(
hidden_size=32,
num_experts=4,
top_k=2,
norm_topk_prob=True,
)
_cast_module_dtype(gate, dtype)
builder_, op, graph = create_test_builder()
x = create_test_input(builder_, "x", [1, 3, 32], dtype)
routing_weights, selected_experts = gate(op, x)
graph.outputs.extend([routing_weights, selected_experts])
# 1e-9 in op.Add(weight_sum, 1e-9) must auto-cast to routing dtype
assert count_op_type(graph, "CastLike") == 0
assert routing_weights.dtype == dtype

@pytest.mark.parametrize("dtype", [ir.DataType.FLOAT16, ir.DataType.BFLOAT16])
def test_routed_scaling_factor_autocasts(self, dtype: ir.DataType):
"""routed_scaling_factor Python float literal auto-casts to routing dtype."""
gate = SigmoidTopKGate(
hidden_size=32,
num_experts=4,
top_k=2,
norm_topk_prob=False,
routed_scaling_factor=2.5,
)
_cast_module_dtype(gate, dtype)
builder_, op, graph = create_test_builder()
x = create_test_input(builder_, "x", [1, 3, 32], dtype)
routing_weights, _ = gate(op, x)
graph.outputs.append(routing_weights)
# routed_scaling_factor=2.5 in op.Mul must auto-cast, not widen to FP32
assert count_op_type(graph, "CastLike") == 0
assert routing_weights.dtype == dtype


class TestSparseMixerGate:
"""SparseMixerGate: verify CastLike(-1e30) preserves input dtype."""

@pytest.mark.parametrize("dtype", [ir.DataType.FLOAT16, ir.DataType.BFLOAT16])
def test_castlike_neg_inf_uses_input_dtype(self, dtype: ir.DataType):
"""op.CastLike(-1e30, scores) must cast the constant to the scores dtype.

Without CastLike the -1e30 literal would be FLOAT32, causing type
mismatches in the op.Where and op.Expand downstream ops.
"""
gate = SparseMixerGate(hidden_size=32, num_experts=4, top_k=2)
_cast_module_dtype(gate, dtype)
builder_, op, graph = create_test_builder()
x = create_test_input(builder_, "x", [1, 3, 32], dtype)
routing_weights, selected_experts = gate(op, x)
graph.outputs.extend([routing_weights, selected_experts])
# CastLike nodes should be present (the pattern is intentional for -1e30)
assert count_op_type(graph, "CastLike") > 0
# Final routing weights must remain in the input dtype
assert routing_weights.dtype == dtype


class TestConformerEncoderLayer:
"""ConformerEncoderLayer: verify 0.5 Macaron weight auto-casts."""

@pytest.mark.parametrize("dtype", [ir.DataType.FLOAT16, ir.DataType.BFLOAT16])
def test_macaron_half_weight_autocasts(self, dtype: ir.DataType):
"""0.5 literal in op.Mul(feed_forward(x), 0.5) auto-casts to input dtype.

The Macaron structure applies half-weight feed-forward modules:
``x += 0.5 * feed_forward_in(x)`` and ``x += 0.5 * feed_forward_out(x)``.
Both 0.5 literals must auto-cast to the hidden state dtype.
"""
layer = ConformerEncoderLayer(d_model=32, num_heads=4, d_inner=64, kernel_size=3)
_cast_module_dtype(layer, dtype)
builder_, op, graph = create_test_builder()
x = create_test_input(builder_, "x", [1, 5, 32], dtype)
# relative_attention_bias: [num_heads, q_len, kv_len]
bias = create_test_input(builder_, "bias", [4, 5, 5], dtype)
result = layer(op, x, bias)
graph.outputs.append(result)
# No CastLike needed — Python float 0.5 auto-casts
assert count_op_type(graph, "CastLike") == 0
assert _get_output_dtype(graph) == dtype


class TestGemma3nAltUp:
"""Gemma3nAltUp: verify router_input_scale Python float auto-casts."""

def _make_config(self, hidden_size: int = 32) -> Gemma3nConfig:
from mobius._configs import ArchitectureConfig

base = ArchitectureConfig(
hidden_size=hidden_size,
intermediate_size=64,
num_attention_heads=4,
num_key_value_heads=2,
num_hidden_layers=1,
vocab_size=256,
)
return Gemma3nConfig(
**{k: getattr(base, k) for k in base.__dataclass_fields__ if hasattr(base, k)},
altup_num_inputs=2,
altup_active_idx=0,
altup_correct_scale=True,
laurel_rank=8,
hidden_size_per_layer_input=16,
vocab_size_per_layer_input=256,
)

@pytest.mark.parametrize("dtype", [ir.DataType.FLOAT16, ir.DataType.BFLOAT16])
def test_router_input_scale_autocasts(self, dtype: ir.DataType):
"""router_input_scale = hidden_size**-1.0 auto-casts in op.Mul.

AltUp._compute_router_modalities multiplies a normalized hidden state
by self.router_input_scale (a Python float). This must not widen
BF16/FP16 computations to FP32.
"""
config = self._make_config(hidden_size=32)
altup = Gemma3nAltUp(config)
_cast_module_dtype(altup, dtype)
builder_, op, graph = create_test_builder()
# altup_num_inputs=2 — provide two hidden state tensors
hs0 = create_test_input(builder_, "hs0", [1, 3, 32], dtype)
hs1 = create_test_input(builder_, "hs1", [1, 3, 32], dtype)
predicted = altup.predict(op, [hs0, hs1])
graph.outputs.extend(predicted)
# router_input_scale (float) in op.Mul must auto-cast
assert count_op_type(graph, "CastLike") == 0
assert predicted[0].dtype == dtype
3 changes: 1 addition & 2 deletions src/mobius/components/_diffusion.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,9 +83,8 @@ def forward(self, op: builder.OpBuilder, hidden_states: ir.Value, timestep_emb:
emb = self._silu(op, timestep_emb)
emb = self.linear(op, emb)
shift, scale = op.Split(emb, num_outputs=2, axis=-1, _outputs=2)
one = op.Constant(value_float=1.0)
hidden_states = self.norm(op, hidden_states)
hidden_states = op.Mul(hidden_states, op.Add(one, op.Unsqueeze(scale, [1])))
hidden_states = op.Mul(hidden_states, op.Add(1.0, op.Unsqueeze(scale, [1])))
hidden_states = op.Add(hidden_states, op.Unsqueeze(shift, [1]))
return hidden_states

Expand Down
17 changes: 8 additions & 9 deletions src/mobius/components/_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -105,13 +105,9 @@ def forward(self, op: builder.OpBuilder, hidden_states: ir.Value):
if self.norm_topk_prob:
# Renormalize selected weights to sum to 1 (prevents vanishing gradients)
weight_sum = op.ReduceSum(routing_weights, [-1], keepdims=True)
eps = op.CastLike(op.Constant(value_float=1e-9), routing_weights)
routing_weights = op.Div(routing_weights, op.Add(weight_sum, eps))
routing_weights = op.Div(routing_weights, op.Add(weight_sum, 1e-9))
if self.routed_scaling_factor != 1.0: # noqa: RUF069
scale = op.CastLike(
op.Constant(value_float=self.routed_scaling_factor), routing_weights
)
routing_weights = op.Mul(routing_weights, scale)
routing_weights = op.Mul(routing_weights, self.routed_scaling_factor)
return routing_weights, selected_experts


Expand Down Expand Up @@ -145,9 +141,12 @@ def _threshold_mask_and_select(self, op, scores, jitter_eps):
factor = op.Max(abs_scores, max_score)
diff = op.Sub(max_score, scores)
ratio = op.Div(diff, factor)
threshold = op.Constant(value_float=2.0 * jitter_eps)
threshold = 2.0 * jitter_eps
mask = op.Greater(ratio, threshold)
neg_inf = op.Constant(value_float=-1e30)
# op.CastLike with Python literal: reuses a single constant, avoids cache-key
# collision that would occur if -1e30 were used as a plain literal in both
# op.Where (auto-cast to typed constant) and op.Expand (unbound → FLOAT).
neg_inf = op.CastLike(-1e30, scores)
masked_scores = op.Where(mask, neg_inf, scores)
weights = op.Softmax(masked_scores, axis=-1)
k_one = op.Constant(value_ints=[1])
Expand All @@ -169,7 +168,7 @@ def forward(self, op: builder.OpBuilder, hidden_states: ir.Value):
)
all_weights.append(weight_k)
all_experts.append(expert_k)
neg_inf = op.Constant(value_float=-1e30)
neg_inf = op.CastLike(-1e30, current_scores)
current_scores = op.ScatterElements(
current_scores,
expert_k,
Expand Down
5 changes: 2 additions & 3 deletions src/mobius/models/gemma3n.py
Original file line number Diff line number Diff line change
Expand Up @@ -196,8 +196,7 @@ def __init__(self, config: Gemma3nConfig):
def _compute_router_modalities(self, op: builder.OpBuilder, x):
"""Compute router modalities: tanh(router(norm(x) * scale))."""
router_input = self.router_norm(op, x)
scale = op.Constant(value_float=self.router_input_scale)
router_input = op.Mul(router_input, scale)
router_input = op.Mul(router_input, self.router_input_scale)
routed = self.modality_router(op, router_input)
return op.Tanh(routed)

Expand Down Expand Up @@ -335,7 +334,7 @@ def forward(
# Residual + Laurel (with sqrt(2) normalization)
attn_gated = op.Add(active, attn_output)
attn_laurel = op.Add(attn_gated, laurel_output)
attn_laurel = op.Div(attn_laurel, op.Constant(value_float=float(math.sqrt(2))))
attn_laurel = op.Div(attn_laurel, float(math.sqrt(2)))

# MLP
mlp_input = self.pre_feedforward_layernorm(op, attn_laurel)
Expand Down
Loading