Skip to content

Commit 2101f88

Browse files
authored
Merge branch 'main' into draft/int3-qmoe-export
2 parents 78d5fb6 + 88f7d04 commit 2101f88

3 files changed

Lines changed: 20 additions & 12 deletions

File tree

‎src/mobius/integrations/ort_genai/auto_export.py‎

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1356,6 +1356,7 @@ def _write_audio_processor_config(
13561356
# 128-dim log-mel frame. Reproduced natively by the ort-extensions
13571357
# ``Gemma4Audio`` op with ``type="raw_frames"`` (pad to a whole number of
13581358
# frames, reshape to (num_tokens, 640)).
1359+
# Disable the decoder's Whisper-compatible 30-second truncation.
13591360
samples_per_token = getattr(audio, "hidden_size", None) or 640
13601361
processor = {
13611362
"feature_extraction": {
@@ -1364,6 +1365,7 @@ def _write_audio_processor_config(
13641365
"operation": {
13651366
"name": "audio_decoder",
13661367
"type": "AudioDecoder",
1368+
"attrs": {"max_samples": 0},
13671369
}
13681370
},
13691371
{
@@ -1384,7 +1386,7 @@ def _write_audio_processor_config(
13841386
proc_filename = "audio_feature_extraction.json"
13851387
elif model_type in _GEMMA4_MODEL_TYPES:
13861388
# Gemma4 USM-style 128-dim log-mel spectrogram via the ort-extensions
1387-
# ``Gemma4Audio`` op with ``type="log_mel"``.
1389+
# ``Gemma4LogMel`` op.
13881390
# OrtxCreateSpeechFeatureExtractor requires the feature_extraction.sequence format.
13891391
processor = {
13901392
"feature_extraction": {
@@ -1397,10 +1399,9 @@ def _write_audio_processor_config(
13971399
},
13981400
{
13991401
"operation": {
1400-
"name": "gemma4_audio",
1401-
"type": "Gemma4Audio",
1402+
"name": "gemma4_log_mel",
1403+
"type": "Gemma4LogMel",
14021404
"attrs": {
1403-
"type": "log_mel",
14041405
"feature_size": 128,
14051406
"sampling_rate": 16000,
14061407
"frame_length_ms": 20.0,

‎src/mobius/integrations/ort_genai/auto_export_test.py‎

Lines changed: 10 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1003,6 +1003,7 @@ def test_audio_gemma4_unified_writes_raw_frames(self, tmp_path):
10031003

10041004
seq = data["feature_extraction"]["sequence"]
10051005
assert seq[0]["operation"]["type"] == "AudioDecoder"
1006+
assert seq[0]["operation"]["attrs"] == {"max_samples": 0}
10061007
op = seq[1]["operation"]
10071008
assert op["name"] == "gemma4_audio"
10081009
assert op["type"] == "Gemma4Audio"
@@ -1029,9 +1030,8 @@ def test_audio_gemma4_writes_feature_extraction_json(self, tmp_path):
10291030
seq = data["feature_extraction"]["sequence"]
10301031
assert len(seq) == 2
10311032
assert seq[0]["operation"]["type"] == "AudioDecoder"
1032-
assert seq[1]["operation"]["type"] == "Gemma4Audio"
1033+
assert seq[1]["operation"]["type"] == "Gemma4LogMel"
10331034
attrs = seq[1]["operation"]["attrs"]
1034-
assert attrs["type"] == "log_mel"
10351035
assert attrs["feature_size"] == 128
10361036
assert attrs["sampling_rate"] == 16000
10371037
assert attrs["frame_length_ms"] == 20.0 # noqa: RUF069
@@ -1967,11 +1967,11 @@ class FakeConfig:
19671967
# First op: AudioDecoder
19681968
op0 = seq[0]["operation"]
19691969
assert op0["type"] == "AudioDecoder"
1970+
assert "max_samples" not in op0.get("attrs", {})
19701971

1971-
# Second op: Gemma4Audio (type=log_mel) with expected attrs
1972+
# Second op: Gemma4LogMel with expected attrs
19721973
op1 = seq[1]["operation"]
1973-
assert op1["type"] == "Gemma4Audio"
1974-
assert op1["attrs"]["type"] == "log_mel"
1974+
assert op1["type"] == "Gemma4LogMel"
19751975
assert op1["attrs"]["feature_size"] == 128
19761976
assert op1["attrs"]["sampling_rate"] == 16000
19771977
assert op1["attrs"]["mel_floor"] == 0.001 # noqa: RUF069
@@ -4198,6 +4198,11 @@ def test_gemma4_unified_native_processor_contract(self, tmp_path):
41984198

41994199
with open(result["audio_processor"], encoding="utf-8") as f:
42004200
audio_processor = json.load(f)
4201+
assert audio_processor["feature_extraction"]["sequence"][0]["operation"] == {
4202+
"name": "audio_decoder",
4203+
"type": "AudioDecoder",
4204+
"attrs": {"max_samples": 0},
4205+
}
42014206
audio_op = audio_processor["feature_extraction"]["sequence"][1]["operation"]
42024207
assert audio_op == {
42034208
"name": "gemma4_audio",

‎src/mobius/models/gemma4.py‎

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,8 @@
1313
- Dual head_dim: local sliding-window layers use ``config.head_dim``; global
1414
full-attention layers use ``config.global_head_dim``.
1515
- Dual RoPE: different ``rope_theta`` and ``partial_rotary_factor`` per layer type.
16-
- Per-layer input gating (disabled when ``hidden_size_per_layer_input == 0``).
16+
- Per-layer input gating with decoder-matched quantized projections
17+
(disabled when ``hidden_size_per_layer_input == 0``).
1718
- Vision encoder: pre-patchified input ``[B, N, 3*P^2]`` with 2D position lookup,
1819
bidirectional attention, 4-norm structure, and scale-then-project pooling.
1920
- Vision projector: scale-free RMSNorm -> Linear (matches ``embed_vision`` weights).
@@ -1577,10 +1578,11 @@ def __init__(self, config: Gemma4Config, layer_idx: int):
15771578

15781579
self._per_layer_dim = config.hidden_size_per_layer_input
15791580
if self._per_layer_dim > 0:
1580-
self.per_layer_input_gate = Linear(
1581+
linear_class = _text_linear_class(config) or Linear
1582+
self.per_layer_input_gate = linear_class(
15811583
config.hidden_size, self._per_layer_dim, bias=False
15821584
)
1583-
self.per_layer_projection = Linear(
1585+
self.per_layer_projection = linear_class(
15841586
self._per_layer_dim, config.hidden_size, bias=False
15851587
)
15861588
self.post_per_layer_input_norm = RMSNorm(

0 commit comments

Comments
 (0)