Skip to content

Gemma 4 E2B/E4B iOS export - #301

Open
Lewis300 wants to merge 25 commits into
apple:mainfrom
Lewis300:lpanos/gemma-4-ios-export
Open

Lewis300 wants to merge 25 commits into
apple:mainfrom
Lewis300:lpanos/gemma-4-ios-export

Conversation

@Lewis300

@Lewis300 Lewis300 commented Sep 29, 2026 •

Copy link
Copy Markdown
Contributor

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 from
this export needs both to run.

cd models/gemma4
uv run export.py --model google/gemma-4-E2B-it

produces exports/gemma_4_e2b_it_4bit_palettized_static/: the .aimodel, an INT8
per-layer-embeddings sidecar (*_ple.safetensors), metadata.json and the tokenizer.

Model (models/ios/gemma4_text.py)

  • Gemma4ForCausalLMForiOS — the text decoder of the multimodal checkpoint: dual head
    dims (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_dict keeps only the
    model.language_model.* keys.
  • Two flat KV caches — a full-context global cache and a fixed-depth sliding-window ring
    (SLIDING_RING_SIZE, 576 for the 512 window), each holding only its layer type's storing
    layers.
  • RoPE arrives precomputed as rope_cos/rope_sin rows instead of position_ids: a 131k
    position overflows a 16-bit position input, and a 32-bit one feeding the RoPE gather fails
    to compile.
  • Per-layer embeddings are externalized. The multi-GB table is written as an INT8 sidecar
    and fed per step as ple_embeddings; ple_scale / ple_zero_point are frozen parameters
    the graph dequantizes with.
  • The final logit soft cap is applied in the runner, not the graph, and published as
    language.final_logit_softcapping.
  • Export contract hooks describe one ladder rung — 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)

  • Follows export/ios.py, except that each transformer (context bucket, query length)
    pair is traced by torch.export at fully static shapes and registered directly as
    extend_{ctx}_{q} / prompt_opt_{ctx}_{q}. BlockedSDPA unrolls its block loop, so the
    graph's op count depends on the context length. The gather is traced dynamic and
    specialized exactly as ios.py does it (a statically traced quantized fused dequant-gather
    doesn't convert).
  • Ladder: contexts 1024 / 8192 / 32768 / 131072, up to --max-context-length (a power
    of two; default and max 131072), which becomes the top bucket. extend at q=8,
    prompt_opt at q=64 — 11 functions for the full ladder.
  • Compression defaults to 4bit_palettized.yaml (4-bit per-grouped-channel k-means,
    8-bit per-tensor per-layer-embedding gate/projection); --compression none skips it.
  • Fails fast: an existing bundle without --overwrite, a --max-context-length that isn't
    a 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.
    --overwrite only removes the old asset just before the new one is saved.
  • Memory: the PLE sidecar is written straight after loading, quantized in row blocks, and
    the table is dropped before palettization.

Primitives and shared code

  • BlockedSDPA (primitives/ios/sdpa.py) — online-softmax attention over the flat global
    cache in block_size chunks, with block-scaled weights so no fp16 accumulator grows with
    the 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 whose
    hf_state_dict_prefix has already been stripped (layers.N. with no leading component).
  • test_mistral.py — seeds TestMistraliOSAttention.test_coreai. It drew unseeded
    weights 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.py rather than export/ios.py? Two things ios.py can't do:

    • One static graph per context bucket. BlockedSDPA unrolls its block loop, so the
      op count depends on the context length, and each rung has to be traced on its own.
      ios.py emits four fixed entrypoints into one AIProgram per call.
    • Dump the per-layer embeddings table. It is written as an INT8 sidecar next to the
      asset rather than exported into the graph, and declared in the bundle metadata.

    Otherwise export.py mirrors ios.py's _export_programs / _convert_to_coreai.

  • Private helpers. export.py imports pipeline._generate_output_name and
    llm.export._load_compression_config_object to get the same bundle naming and YAML
    validation 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 the
    ring 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 ring
    wrap 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 and
    reloaded, 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 — BlockedSDPA parity against SDPA and
    scaled_dot_product_attention (fp32/fp16, 65536 keys), and rank <= 4 for every op.

Testing

  • pytest python/tests (with USE_LOCAL_COREAI=1): all pass.
  • Export manifest (function names, input shapes, I/O names, hardware constraints) compared
    across the refactor on a synthetic model: identical.
  • export.py end to end on a small local checkpoint (--compression none,
    --max-context-length 2048).
  • E2B exported with the default palettization and run through llm-runner: generates.

@gokulkrishna98 gokulkrishna98 left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks.

@Lewis300
Lewis300 force-pushed the lpanos/gemma-4-ios-export branch from 7288761 to 84befe0 Compare September 29, 2026 21:12
Comment thread models/gemma4/README.md Outdated
Comment thread models/gemma4/README.md Outdated
Comment thread models/README.md Outdated
Comment thread models/gemma4/README.md Outdated
@Lewis300
Lewis300 force-pushed the lpanos/gemma-4-ios-export branch from d95a19c to 129805b Compare September 30, 2026 19:28
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.
Lewis300 and others added 9 commits September 30, 2026 15:48
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
Lewis300 force-pushed the lpanos/gemma-4-ios-export branch from 129805b to 9c2b3b6 Compare September 30, 2026 22:52
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants