Conversation
carinapeng
approved these changes
Sep 29, 2026
Lewis300
force-pushed
the
lpanos/gemma-4-ios-export
branch
from
September 29, 2026 21:12
7288761 to
84befe0
Compare
tjia1818
reviewed
Sep 29, 2026
tjia1818
reviewed
Sep 29, 2026
tjia1818
reviewed
Sep 29, 2026
stikves
approved these changes
Sep 30, 2026
tjia1818
reviewed
Sep 30, 2026
tjia1818
approved these changes
Sep 30, 2026
Lewis300
force-pushed
the
lpanos/gemma-4-ios-export
branch
from
September 30, 2026 19:28
d95a19c to
129805b
Compare
Adds the Python export for the Gemma 4 E2B/E4B per-layer-embedding variants on iOS. The Swift runner support is a separate change. * Gemma4ForCausalLMForiOS (models/ios/gemma4_text.py) for the E2B/E4B PLE variants: dual head dims, shared KV across same-type layers, v-norm, proportional RoPE on global layers. The final logit soft cap is left to the runner and published as language.final_logit_softcapping. * models/gemma4/export.py emits a per-context blocked ladder of statically-shaped programs, a flat global KV cache paired with a sliding-window ring, and an INT8 per-layer embeddings sidecar. * BlockedSDPA for chunked global attention; apply_rope takes a head_axis so callers holding (B, S, n_heads, head_dim) skip a transpose. * _is_layer_key_beyond tolerates keys whose hf_state_dict_prefix has already been stripped.
* Name bundles with the shared pipeline._generate_output_name instead of a local copy that dropped the compression suffix, which sent the palettized and uncompressed exports to the same bundle. The default palettized export is now gemma_4_e2b_it_gemma4_4bit_palettized_static; --compression none stays gemma_4_e2b_it_static. * Resolve the bundle paths before loading the model, so a missing --overwrite fails immediately instead of after the full ladder export. * Always write the PLE sidecar. The graph always takes ple_embeddings, so the hasattr-guarded skip produced bundles the runner cannot serve. Likewise the model now raises when the checkpoint has no per-layer embedding table or forward() gets no ple_embeddings, rather than silently skipping the per-layer inputs and producing wrong logits.
Only code added by the Gemma 4 port is touched. * export.py: drop --platform (one choice, never read) and --compute-precision/PRECISIONS (everything but float16 was rejected); drop context_ladder's unused block_size, the redundant vocab_size parameters, a model-type check that could not fail and the Python version guard pyproject already enforces. Read kv_block_size from the model instead of a duplicated default. Shorten the module docstring, which also claimed --compression selects a preset. * gemma4_text.py: drop the optional HF import and _HF_MODEL_CLASS (from_hf is overridden), a pass-through __init__, Attention's unused layer_idx and is_blocked, the duplicated Q reshape and the global_offset alias. QUERY_LENGTHS becomes MAX_QUERY_LEN, derived from the base class's IOS_STATIC_QUERY_LENS. Gemma4CombinedRoPE moves into the test module, its only user. * Tests: use the production sliding_ring_size, drop the redundant short-prompt test, and remove references to scripts that don't exist.
* _mutate_state_dict needed only the PLE scale but quantized the whole table to get it, then dump_ple_embedding quantized it again. Compute the scale directly from max|w| instead: scaling by a positive constant commutes with the max, so it matches quantize_per_tensor bit for bit (checked on fp16, bf16 and fp32) without an fp32 copy of the multi-GB table. * Drop the untied lm_head path. from_hf keeps only model.language_model.* keys, so a real lm_head.weight was discarded and replaced by the embedding table. E2B and E4B tie their embeddings; construction now raises for an untied config. * The E2B parity test skips only on OSError (download and gated-repo failures) instead of any exception, so a bug in from_hf fails the test rather than skipping it.
…ooks * ple_scale / ple_zp become frozen nn.Parameters on Gemma4Extend, like emb_scale / emb_zero_point, written into the state dict by _mutate_state_dict. They were plain tensor attributes, invisible to state_dict() and attached after load_state_dict through _ple_*_pending glue that from_hf and the test helper each repeated. The scale now follows the table's dtype, as emb_scale does; the shipping fp16 export is unchanged. * export.py loads its compression YAML with the shared llm.export._load_compression_config_object and reads config.head_dim directly, dropping its own copies of both. * --dev-extend-qlens / --dev-prompt-qlens reject query lengths that don't divide MAX_QUERY_LEN: the sliding ring and every context bucket are sized in multiples of it, so anything else could wrap a cache write. * export_static_shape_configs and export_hardware_constraints raise NotImplementedError like the other two single-graph hooks, instead of returning base-class values that are wrong for Gemma 4.
The E2B parity test nulled attention_logit_cap / attn_logit_softcapping on the HF config before comparing. Gemma4TextAttention never applies an attention logit cap (only Gemma4AudioAttention reads attention_logit_cap), so the loop had no effect on the text logits under comparison. Only the final logit soft cap needs matching, which the test already applies to our logits in place of the runner.
…per rung
export.py now follows coreai_models.export.ios. The model class owns the
graph contract for one (context bucket, query length) rung --
build_reference_inputs, build_dynamic_shapes, export_static_shape_configs
and export_hardware_constraints -- and export.py only loops over the
rungs and dumps the PLE sidecar.
Each rung is traced by torch.export at fully static shapes and
registered directly as extend_{ctx}_{q} / prompt_opt_{ctx}_{q} /
gather_embeddings_{q}. BlockedSDPA unrolls its block loop, so the graph
already differed per context; tracing the query length statically too
means there are no static shape configs left to apply, and
export_static_shape_configs returns empty.
The sliding ring depth is a class constant (SLIDING_RING_SIZE, from
SLIDING_WINDOW = 512) because export_hardware_constraints gets no
config; build_reference_inputs rejects a config whose sliding_window
differs.
This removes export.py's own reference-input builder, palettization
input builder and constraint/shape tables (the palettizer now calibrates
on the smallest rung's reference inputs). Checked against the previous
commit on a synthetic Gemma 4 with a recording converter: the same 11
functions, identical declared input shapes, I/O names and hardware
constraints for every function, and identical palettization inputs.
Drop Gemma4ForCausalLMForiOS.from_hf, which copied the base loader and differed in only three ways, each now expressed where the base expects it: * _HF_MODEL_CLASS is Gemma4ForConditionalGeneration. * _mutate_state_dict keeps only the model.language_model.* text-decoder keys (as model.*) when given a multimodal state dict, dropping the vision and audio towers and the top-level lm_head. * The override's dtype check exempted every non-float tensor; the base exempts *zero_point* names, so ple_zp becomes ple_zero_point, matching emb_zero_point. It also moves off the deprecated torch_dtype= argument. Checked on a tiny local multimodal checkpoint with vision and audio towers: the old override and the base loader produce identical state dicts with and without embedding quantization, full and truncated to 4 layers.
* The ring-wrap parity test runs over one, four and three ragged global attention blocks (kv_block_size None / 8 / 12), so the flash loop's cross-block rescale is checked at the model level, not only in the BlockedSDPA primitive tests. * test_gemma4_export.py exports a tiny Gemma 4 through models/gemma4/export.py with the real converter, saves and reloads the asset, and checks every function's name, inputs and states against what the runner looks up.
TestMistraliOSAttention.test_coreai draws its weights and inputs from the global RNG without seeding, so its result depended on which tests ran before it. After the Gemma 4 tests (which call torch.manual_seed) it drew inputs that missed its 1e-5 tolerance by ~2e-5. Seeded, it passes at every seed tried, alone and in the full test_ios_layers run.
4a21236 traced gather_embeddings at fully static shapes, one program
per query length. With quantized embeddings (the shipping default) the
gather lowers to the fused dequant-gather composite, and a statically
traced one fails the converter's optimization passes ("input types did
not match callee signature"); traced dynamic in the query length it
converts. So every quantized export on this branch failed to convert.
Neither the manifest comparison nor the conversion test caught it,
because both used an unquantized model.
The gather is now traced once, dynamic in the query length, and
specialized with a static shape config -- the base class's gather entry
for both build_dynamic_shapes and export_static_shape_configs, exactly
as ios.py exports it. export.py narrows the specializations to the query
lengths the ladder emits, so the functions are unchanged
(gather_embeddings_8 / _64). The transformer rungs stay fully static.
--dev-*-qlens must now be among IOS_STATIC_QUERY_LENS, the gather's
specializations.
test_gemma4_export.py now exports a quantized model; it fails with the
static gather and passes with this change. The export manifest again
matches the original port exactly.
Fixes: * A --max-context-length that isn't a power of two crashed after palettization: the ladder rounds the top bucket up, but each rung was traced with TraceSpec(max_context_length=<requested>), which rejects a longer cache. Rungs are now bounded by the top bucket, and the CLI rejects a context no longer than the prefill query length. * --overwrite deleted the existing asset before anything had run; the check still happens up front, but the old asset is only removed just before the new one is saved. * The PLE table (~4.7 GB fp16 for E2B) is not a module weight, so mmap never evicted it and dump_ple_embedding made a full fp32 copy. The sidecar is now written straight after loading, quantized a block of rows at a time, and the table is dropped before palettization. * Fail fast instead of at the end of the export: a sliding window other than the one the export is built for (checked before any weights are loaded), a checkpoint without per-layer embeddings, and an export with disable_embedding_quantization (the graph and runner take INT8 PLE). * The runner-facing metadata (sliding_window, rope) is required rather than silently omitted, with no 0.25 partial_rotary_factor fallback; only a missing generation_config is tolerated for eos ids. * --dev-ladder-only errors when none of its buckets are in the ladder instead of silently exporting the top bucket. * The E2B parity test downloads the real model twice and `slow` tests run by default (CI runs them), so it now needs RUN_GEMMA4_E2B_PARITY=1. Cleanups: * BlockedSDPA takes a required head_dim and names its transpose limit and fp16 -inf floor. * The compression recipe is renamed 4bit_palettized.yaml, so the default bundle is gemma_4_e2b_it_4bit_palettized_static rather than repeating "gemma4"; its regexes are escaped consistently and it has a header. * Revert the comment-only change to tests/_runner_infra/_deps.py. * Docstrings: one-rung meaning of the hooks' max_context_length, which 32-bit input fails to compile, stale flash-attention test wording, calibration inputs only driving the trace. Tests: non-power-of-two ladders trace every rung; CLI bounds; the metadata requires the runner's fields; the PLE sidecar's rows and scale equal quantize_per_tensor's and the graph's; unquantized PLE and non-PLE checkpoints are rejected; _is_layer_key_beyond with and without prefix. The export manifest is unchanged against the original port, a quantized tiny model converts, and export.py runs end to end on a local checkpoint.
It downloaded google/gemma-4-E2B-it and loaded a ~2B model twice, and nothing in the repo keeps `slow` tests out of a default run, so it either ran wherever the Hub was reachable or needed a one-off env var to gate it. The synthetic-config parity tests cover the same forward (ring wrap, multi-block attention, the runner-side soft cap) without the download.
The updated wheel requires a graph's hardware constraints to be set before its static shape config, and removes AIProgram.optimize(). Same change as export/ios.py in apple#291.
* Remove --dev-extend-qlens / --dev-prompt-qlens / --dev-ladder-only and the DevOverrides plumbing behind them. The ladder always uses the shipping query lengths (extend q=8, prompt_opt q=64) and context buckets. * --max-context-length must be a power of two: it becomes the ladder's top bucket, so any other value was silently rounded up.
The docstring and README said only that the PLE variants don't fit the generic pipeline. State the two reasons: each context bucket needs its own statically traced graph, and the per-layer embeddings table is dumped as a sidecar.
The export already passes build_aimodel_metadata(hf_model_id) to save_asset, but neither checkpoint had an entry, so exported assets carried no author, license or description and the export warned.
Co-authored-by: tjia1818 <35608981+tjia1818@users.noreply.github.com>
Co-authored-by: tjia1818 <35608981+tjia1818@users.noreply.github.com>
The runner now reads the Gemma-specific language settings from one `language.overrides` block rather than top-level `language` keys, so the export writes them there: `sliding_window`, `rope`, and `final_logit_softcapping`. The export test checks all three under `overrides`; the README and the gemma4_text docstring point at the new keys. Bundles exported before this change need re-exporting for the updated runner.
The comment on the shipping query lengths said prompts of 64 tokens or fewer go through prompt_opt. The runner prefills in q=64 prompt_opt chunks only while more than 64 tokens remain, and runs the rest as q=8 extend steps.
Lewis300
force-pushed
the
lpanos/gemma-4-ios-export
branch
from
September 30, 2026 22:52
129805b to
9c2b3b6
Compare
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
iOS: Gemma 4 E2B/E4B export
Adds an iOS export for the Gemma 4 per-layer-embedding variants (
google/gemma-4-E2B-it,google/gemma-4-E4B-it). Runner support is in the companion Swift PR (#302); a bundle fromthis export needs both to run.
produces
exports/gemma_4_e2b_it_4bit_palettized_static/: the.aimodel, an INT8per-layer-embeddings sidecar (
*_ple.safetensors),metadata.jsonand the tokenizer.Model (
models/ios/gemma4_text.py)Gemma4ForCausalLMForiOS— the text decoder of the multimodal checkpoint: dual headdims (sliding 256 / global 512), KV shared across same-type layers, v-norm, proportional
RoPE on the global layers, per-layer embeddings. Loads through the base
from_hf(
_HF_MODEL_CLASS = Gemma4ForConditionalGeneration);_mutate_state_dictkeeps only themodel.language_model.*keys.(
SLIDING_RING_SIZE, 576 for the 512 window), each holding only its layer type's storinglayers.
rope_cos/rope_sinrows instead ofposition_ids: a 131kposition overflows a 16-bit position input, and a 32-bit one feeding the RoPE gather fails
to compile.
and fed per step as
ple_embeddings;ple_scale/ple_zero_pointare frozen parametersthe graph dequantizes with.
language.final_logit_softcapping.build_reference_inputs/build_dynamic_shapes/export_static_shape_configs/export_hardware_constraints—with the rung's context bucket carried in
TraceSpec.cache_seq_len.Export (
models/gemma4/export.py)export/ios.py, except that each transformer (context bucket, query length)pair is traced by
torch.exportat fully static shapes and registered directly asextend_{ctx}_{q}/prompt_opt_{ctx}_{q}.BlockedSDPAunrolls its block loop, so thegraph's op count depends on the context length. The gather is traced dynamic and
specialized exactly as
ios.pydoes it (a statically traced quantized fused dequant-gatherdoesn't convert).
--max-context-length(a powerof two; default and max 131072), which becomes the top bucket.
extendat q=8,prompt_optat q=64 — 11 functions for the full ladder.4bit_palettized.yaml(4-bit per-grouped-channel k-means,8-bit per-tensor per-layer-embedding gate/projection);
--compression noneskips it.--overwrite, a--max-context-lengththat isn'ta power of two or is no longer than the prefill query length, or a sliding window other than
512 fails before any weights are loaded; a checkpoint without per-layer embeddings fails while loading, before conversion.
--overwriteonly removes the old asset just before the new one is saved.the table is dropped before palettization.
Primitives and shared code
BlockedSDPA(primitives/ios/sdpa.py) — online-softmax attention over the flat globalcache in
block_sizechunks, with block-scaled weights so no fp16 accumulator grows withthe key count.
apply_rope(..., head_axis=)(primitives/ios/rope.py) — lets callers holding Q/K as(B, S, n_heads, head_dim)skip a transpose._is_layer_key_beyond(models/base.py) — also matches keys whosehf_state_dict_prefixhas already been stripped (layers.N.with no leading component).test_mistral.py— seedsTestMistraliOSAttention.test_coreai. It drew unseededweights and inputs, so it passed or failed depending on which tests ran first; after the
Gemma tests it missed its 1e-5 tolerance by ~2e-5.
Design notes
Why a separate
export.pyrather thanexport/ios.py? Two thingsios.pycan't do:BlockedSDPAunrolls its block loop, so theop count depends on the context length, and each rung has to be traced on its own.
ios.pyemits four fixed entrypoints into oneAIProgramper call.asset rather than exported into the graph, and declared in the bundle metadata.
Otherwise
export.pymirrorsios.py's_export_programs/_convert_to_coreai.Private helpers.
export.pyimportspipeline._generate_output_nameandllm.export._load_compression_config_objectto get the same bundle naming and YAMLvalidation as the shared pipeline. Happy to make them public in a follow-up.
Sliding window as a class constant.
export_hardware_constraints(cls, max_context_length)receives no config, but the sliding cache's alignment depends on thering depth; the export checks the checkpoint's window against it up front.
Tied embeddings only; an untied config raises. E2B and E4B are tied.
Tests
test_gemma4.py— chunked-prefill parity with HF over a synthetic config, including ringwrap and one / several / ragged global-attention blocks; the runner-side soft cap; the
one-rung export contract; the PLE sidecar's rows and scale matching the graph's; rejection
of unquantized PLE and non-PLE checkpoints;
_is_layer_key_beyond.test_gemma4_export.py— a quantized tiny model through the real converter, saved andreloaded, with every function's name and I/O checked against what the runner looks up;
a shortened ladder traces every rung; CLI bounds; required metadata fields.
test_flash_attention.py—BlockedSDPAparity againstSDPAandscaled_dot_product_attention(fp32/fp16, 65536 keys), and rank <= 4 for every op.Testing
pytest python/tests(withUSE_LOCAL_COREAI=1): all pass.across the refactor on a synthetic model: identical.
export.pyend to end on a small local checkpoint (--compression none,--max-context-length 2048).llm-runner: generates.