Skip to content

Commit 839ccf3

Browse files
committed
Isolate Gemma4 per-layer embedding quantization
Keep embedding_bits and split per-layer embedding quantization separate from the shared text QuantizationConfig so decoder Linear, LM head, and token embedding quantization are not changed incidentally. Add dedicated Gemma4 config fields for per-layer embedding quantization and a regression test that split defaults do not mutate an existing quantization config.\n\nCo-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: justinchuby <justinchuby@users.noreply.github.com>
1 parent a1d6f94 commit 839ccf3

5 files changed

Lines changed: 124 additions & 82 deletions

File tree

src/mobius/_builder.py

Lines changed: 10 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -592,29 +592,17 @@ def build(
592592
# ``embedding_bits`` is only for Gemma4's per-layer embedding table. Do not
593593
# attach a QuantizationConfig to ordinary text models, because that changes
594594
# their regular token embedding/Linear modules.
595-
if getattr(config, "hidden_size_per_layer_input", 0) and getattr(
596-
config, "vocab_size_per_layer_input", 0
595+
if (
596+
hasattr(config, "per_layer_embedding_bits")
597+
and getattr(config, "hidden_size_per_layer_input", 0)
598+
and getattr(config, "vocab_size_per_layer_input", 0)
597599
):
598-
from mobius._configs import QuantizationConfig
599-
600-
if config.quantization is None:
601-
config = dataclasses.replace(
602-
config,
603-
quantization=QuantizationConfig(
604-
bits=embedding_bits,
605-
group_size=32,
606-
quant_method="none",
607-
sym=False,
608-
quantize_embeddings=True,
609-
),
610-
)
611-
else:
612-
qc = dataclasses.replace(
613-
config.quantization,
614-
bits=embedding_bits,
615-
quantize_embeddings=True,
616-
)
617-
config = dataclasses.replace(config, quantization=qc)
600+
config = dataclasses.replace(
601+
config,
602+
per_layer_embedding_bits=embedding_bits,
603+
per_layer_embedding_group_size=32,
604+
per_layer_embedding_sym=False,
605+
)
618606

619607
if task is None:
620608
task = _default_task_for_model(model_type)

src/mobius/_configs/_base.py

Lines changed: 12 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -324,7 +324,9 @@ def _extract_vision_config(config, parent_config, model_type: str) -> dict:
324324
hooks live under :mod:`mobius._configs.per_model` and are
325325
registered with :mod:`mobius._configs._extractors` at import time.
326326
"""
327-
from mobius._configs import per_model # noqa: F401 - imported for registration side effect
327+
from mobius._configs import (
328+
per_model, # ruff: ignore[unused-import] - imported for registration side effect
329+
)
328330
from mobius._configs._extractors import extract_vision_config as _dispatch
329331

330332
return _dispatch(config, parent_config, model_type)
@@ -337,7 +339,9 @@ def _extract_audio_config(config, parent_config, model_type: str) -> dict:
337339
hooks live under :mod:`mobius._configs.per_model` and are
338340
registered with :mod:`mobius._configs._extractors` at import time.
339341
"""
340-
from mobius._configs import per_model # noqa: F401 - imported for registration side effect
342+
from mobius._configs import (
343+
per_model, # ruff: ignore[unused-import] - imported for registration side effect
344+
)
341345
from mobius._configs._extractors import extract_audio_config as _dispatch
342346

343347
return _dispatch(config, parent_config, model_type)
@@ -1593,6 +1597,12 @@ class Gemma4Config(VisionLanguageConfig):
15931597
# is too small for the fused [V, L*D] per-layer embedding table, to split
15941598
# it into L separate [V, D] tables that each fit within the EP's buffer limit.
15951599
split_per_layer_embedding: bool = False
1600+
# Dedicated quantization settings for the split per-layer embedding table.
1601+
# Kept separate from ``quantization`` so embedding_bits does not affect text
1602+
# token embeddings, LM head, or decoder Linear projection quantization.
1603+
per_layer_embedding_bits: int | None = None
1604+
per_layer_embedding_group_size: int = 32
1605+
per_layer_embedding_sym: bool = False
15961606

15971607
@classmethod
15981608
def from_transformers(cls, config, parent_config=None) -> Gemma4Config:

src/mobius/models/gemma4.py

Lines changed: 38 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -230,6 +230,26 @@ def _per_layer_embedding_block_size(embedding_dim: int, group_size: int) -> int:
230230
return group_size
231231

232232

233+
def _per_layer_embedding_quantization(config: Gemma4Config) -> tuple[int, int, bool] | None:
234+
"""Return split per-layer embedding quantization as (bits, group_size, symmetric)."""
235+
bits = getattr(config, "per_layer_embedding_bits", None)
236+
if bits is None:
237+
quantization_config = getattr(config, "quantization", None)
238+
if quantization_config is None or not getattr(
239+
quantization_config, "quantize_embeddings", False
240+
):
241+
return None
242+
bits = quantization_config.bits
243+
group_size = quantization_config.group_size
244+
symmetric = quantization_config.sym
245+
else:
246+
group_size = getattr(config, "per_layer_embedding_group_size", 32)
247+
symmetric = getattr(config, "per_layer_embedding_sym", False)
248+
if bits not in (4, 8):
249+
raise ValueError(f"quantize_embeddings requires bits=4 or bits=8, got {bits}")
250+
return bits, group_size, symmetric
251+
252+
233253
def _dtype_safe_compress(
234254
op: OpBuilder, data: ir.Value, condition: ir.Value, *, axis: int
235255
) -> ir.Value:
@@ -1830,25 +1850,20 @@ def __init__(self, config: Gemma4Config):
18301850
# 256 MiB limit; ~128 MiB each vs ~4.7 GB fused).
18311851
# Only the table actually called in forward() is realized as an
18321852
# ONNX initializer, so the unused one adds no graph weight.
1833-
qc = getattr(config, "quantization", None)
1834-
if (
1835-
qc is not None
1836-
and getattr(qc, "quantize_embeddings", False)
1837-
and config.split_per_layer_embedding
1838-
):
1839-
block_size = _per_layer_embedding_block_size(
1840-
self._per_layer_dim, qc.group_size
1841-
)
1853+
per_layer_quant = _per_layer_embedding_quantization(config)
1854+
if per_layer_quant is not None and config.split_per_layer_embedding:
1855+
bits, group_size, symmetric = per_layer_quant
1856+
block_size = _per_layer_embedding_block_size(self._per_layer_dim, group_size)
18421857
self.embed_tokens_per_layer_split = nn.ModuleList(
18431858
[
18441859
QuantizedScaledWordEmbedding(
18451860
vocab_per_layer,
18461861
self._per_layer_dim,
18471862
config.pad_token_id,
18481863
embed_scale=float(self._per_layer_dim**0.5),
1849-
bits=qc.bits,
1864+
bits=bits,
18501865
block_size=block_size,
1851-
has_zero_point=not qc.sym,
1866+
has_zero_point=not symmetric,
18521867
)
18531868
for _ in range(self._num_layers)
18541869
]
@@ -2402,24 +2417,18 @@ def preprocess_weights(
24022417
f"got {fused.shape[1]}"
24032418
)
24042419
chunks = fused.chunk(num_layers, dim=1)
2405-
qc = getattr(self.config, "quantization", None)
2406-
quantize_per_layer = qc is not None and getattr(
2407-
qc, "quantize_embeddings", False
2408-
)
2409-
if quantize_per_layer:
2410-
if qc.bits not in (4, 8):
2411-
raise ValueError(
2412-
f"quantize_embeddings requires bits=4 or bits=8, got {qc.bits}"
2413-
)
2420+
per_layer_quant = _per_layer_embedding_quantization(self.config)
2421+
if per_layer_quant is not None:
2422+
bits, group_size, symmetric = per_layer_quant
24142423
from mobius._weight_utils import quantize_embedding_rtn
24152424

2416-
block_size = _per_layer_embedding_block_size(per_layer_dim, qc.group_size)
2425+
block_size = _per_layer_embedding_block_size(per_layer_dim, group_size)
24172426
for i, chunk in enumerate(chunks):
24182427
qweight, scales, zero_points = quantize_embedding_rtn(
24192428
chunk.contiguous(),
2420-
bits=qc.bits,
2429+
bits=bits,
24212430
block_size=block_size,
2422-
symmetric=qc.sym,
2431+
symmetric=symmetric,
24232432
)
24242433
state_dict[f"model.embed_tokens_per_layer_split.{i}.qweight"] = qweight
24252434
state_dict[f"model.embed_tokens_per_layer_split.{i}.scales"] = scales
@@ -3229,24 +3238,18 @@ def preprocess_weights(
32293238
f"got {fused.shape[1]}"
32303239
)
32313240
chunks = fused.chunk(num_layers, dim=1)
3232-
qc = getattr(self.config, "quantization", None)
3233-
quantize_per_layer = qc is not None and getattr(
3234-
qc, "quantize_embeddings", False
3235-
)
3236-
if quantize_per_layer:
3237-
if qc.bits not in (4, 8):
3238-
raise ValueError(
3239-
f"quantize_embeddings requires bits=4 or bits=8, got {qc.bits}"
3240-
)
3241+
per_layer_quant = _per_layer_embedding_quantization(self.config)
3242+
if per_layer_quant is not None:
3243+
bits, group_size, symmetric = per_layer_quant
32413244
from mobius._weight_utils import quantize_embedding_rtn
32423245

3243-
block_size = _per_layer_embedding_block_size(per_layer_dim, qc.group_size)
3246+
block_size = _per_layer_embedding_block_size(per_layer_dim, group_size)
32443247
for i, chunk in enumerate(chunks):
32453248
qweight, scales, zero_points = quantize_embedding_rtn(
32463249
chunk.contiguous(),
3247-
bits=qc.bits,
3250+
bits=bits,
32483251
block_size=block_size,
3249-
symmetric=qc.sym,
3252+
symmetric=symmetric,
32503253
)
32513254
renamed[f"decoder.model.embed_tokens_per_layer_split.{i}.qweight"] = (
32523255
qweight

src/mobius/models/gemma4_test.py

Lines changed: 44 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88
import onnx_ir as ir
99
import torch
1010

11-
from mobius._configs import AudioConfig, Gemma4Config
11+
from mobius._configs import AudioConfig, Gemma4Config, QuantizationConfig
1212
from mobius.models.gemma4 import Gemma4CausalLMModel, Gemma4EmbeddingModel, Gemma4Model
1313

1414

@@ -120,6 +120,49 @@ def test_per_expert_scale_not_folded(self):
120120
assert torch.allclose(result[key], torch.ones(4))
121121

122122

123+
class TestGemma4PerLayerEmbeddingQuantization:
124+
"""Per-layer embedding quantization stays independent from text quantization."""
125+
126+
def test_split_defaults_do_not_mutate_existing_quantization_config(self, monkeypatch):
127+
import dataclasses
128+
129+
import mobius.tasks._gemma4 as gemma4_task_module
130+
from mobius.tasks._gemma4 import Gemma4Task
131+
132+
class _TinyCaps:
133+
max_buffer_size = 1
134+
135+
monkeypatch.setattr(gemma4_task_module, "ep_capabilities", lambda: _TinyCaps())
136+
137+
config = _tiny_gemma4_config(
138+
enable_moe_block=False,
139+
hidden_size_per_layer_input=32,
140+
vocab_size_per_layer_input=64,
141+
)
142+
config = dataclasses.replace(
143+
config,
144+
quantization=QuantizationConfig(
145+
bits=8,
146+
group_size=64,
147+
quant_method="none",
148+
sym=True,
149+
quantize_embeddings=False,
150+
),
151+
)
152+
153+
pkg = Gemma4Task().build(Gemma4Model(config), config)
154+
155+
assert config.quantization is not None
156+
assert config.quantization.bits == 8
157+
assert config.quantization.group_size == 64
158+
assert config.quantization.sym is True
159+
assert config.quantization.quantize_embeddings is False
160+
assert config.per_layer_embedding_bits == 4
161+
assert config.per_layer_embedding_group_size == 32
162+
assert config.per_layer_embedding_sym is False
163+
assert any(node.op_type == "GatherBlockQuantized" for node in pkg["decoder"].graph)
164+
165+
123166
class TestGemma4EmbeddingModel:
124167
def test_reuses_token_id_fields_for_masking(self):
125168
config = _tiny_gemma4_config(

src/mobius/tasks/_gemma4.py

Lines changed: 20 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@
2525
from onnxscript import GraphBuilder, nn
2626

2727
from mobius._build_context import ep_capabilities
28-
from mobius._configs import Gemma4Config, QuantizationConfig
28+
from mobius._configs import Gemma4Config
2929
from mobius._model_package import ModelPackage
3030
from mobius._pipeline_contract import (
3131
declare_component_presence,
@@ -444,28 +444,24 @@ def build(
444444
# in mobius.build() / the MobiusBuilder Olive pass config.
445445
if config.split_per_layer_embedding:
446446
qc = config.quantization
447-
if qc is not None and qc.quantize_embeddings:
448-
# embedding_bits was set by the caller — use its bits value
447+
if getattr(config, "per_layer_embedding_bits", None) is not None:
448+
if config.per_layer_embedding_bits not in (4, 8):
449+
raise ValueError(
450+
"quantize_embeddings requires bits=4 or bits=8, "
451+
f"got {config.per_layer_embedding_bits}"
452+
)
453+
elif qc is not None and qc.quantize_embeddings:
449454
if qc.bits not in (4, 8):
450455
raise ValueError(
451456
f"quantize_embeddings requires bits=4 or bits=8, got {qc.bits}"
452457
)
453-
elif qc is not None:
454-
# Existing quantization config without embedding quantization —
455-
# enable it with INT4 (the existing bits may be for MatMul
456-
# quantization which is unrelated, so override for embeddings)
457-
config.quantization.quantize_embeddings = True
458-
config.quantization.bits = 4
459-
config.quantization.group_size = 32
460-
config.quantization.sym = False
458+
config.per_layer_embedding_bits = qc.bits
459+
config.per_layer_embedding_group_size = qc.group_size
460+
config.per_layer_embedding_sym = qc.sym
461461
else:
462-
config.quantization = QuantizationConfig(
463-
bits=4,
464-
group_size=32,
465-
quant_method="mobius",
466-
sym=False,
467-
quantize_embeddings=True,
468-
)
462+
config.per_layer_embedding_bits = 4
463+
config.per_layer_embedding_group_size = 32
464+
config.per_layer_embedding_sym = False
469465
# The module was already constructed before build() runs, so the
470466
# per-layer embeddings are plain Embedding (Gather). Replace them
471467
# with QuantizedScaledWordEmbedding (GatherBlockQuantized).
@@ -474,23 +470,25 @@ def build(
474470
_per_layer_embedding_block_size,
475471
)
476472

477-
qc = config.quantization
473+
bits = config.per_layer_embedding_bits
474+
group_size = config.per_layer_embedding_group_size
475+
symmetric = config.per_layer_embedding_sym
478476
per_layer_dim = getattr(config, "hidden_size_per_layer_input", 0)
479477
vocab_per_layer = getattr(config, "vocab_size_per_layer_input", 0)
480478
import numpy as np
481479

482480
embed_scale = float(np.float16(per_layer_dim**0.5))
483-
block_size = _per_layer_embedding_block_size(per_layer_dim, qc.group_size)
481+
block_size = _per_layer_embedding_block_size(per_layer_dim, group_size)
484482
module.decoder.model.embed_tokens_per_layer_split = nn.ModuleList(
485483
[
486484
QuantizedScaledWordEmbedding(
487485
vocab_per_layer,
488486
per_layer_dim,
489487
config.pad_token_id,
490488
embed_scale=embed_scale,
491-
bits=qc.bits,
489+
bits=bits,
492490
block_size=block_size,
493-
has_zero_point=not qc.sym,
491+
has_zero_point=not symmetric,
494492
)
495493
for _ in range(config.num_hidden_layers)
496494
]

0 commit comments

Comments
 (0)