Skip to content

feat(example): Add Spark-X2.5 model example with hybrid attention support - #22865

Merged
nil-is-all merged 4 commits into
pytorch:mainfrom
XHToken:add-spark-x2.5
Sep 28, 2026
Merged

nil-is-all merged 4 commits into
pytorch:mainfrom
XHToken:add-spark-x2.5

Conversation

@dongjiang1989

@dongjiang1989 dongjiang1989 commented Sep 16, 2026 •

Copy link
Copy Markdown
Contributor

Summary

Add examples/models/spark_x2_5/ with support for the Spark-X2.5-1.7B and Spark-X2.5-4B models from XHToken. These are compact, general-purpose language models with a hybrid attention architecture (3 sliding-window attention layers + 1 full-attention layer, repeating) that natively supports context windows up to 1M tokens.

Architecture highlights

  • Hybrid attention: pattern of 3 sliding-window + 1 full-attention layer (28 layers for 1.7B, 36 for 4B)
  • Per-layer-type RoPE: full-attention layers use rope_theta=5M, partial_rotary_factor=0.25; sliding-attention layers use rope_theta=10K, partial_rotary_factor=1.0
  • Headwise attention output gate: per-head sigmoid gate broadcast over head dim (a lighter variant of use_attn_o_gate)
  • Sliding window KV cache (RingKVCache, window=512 tokens)
  • GELU activation in MLP (vs default SiLU)
  • Tied word embeddings

Files added (in examples/models/spark_x2_5/)

  • convert_weights.py — HF safetensors → Meta format conversion with fused QKV split, sharded checkpoint support, and tied embeddings handling
  • config/ — JSON configs for 1.7B and 4B variants; XNNPack (fp32, q8da4w), CoreML, and MLX export configs
  • test_spark_x2_5.py — config validation and model registration tests
  • BUCK, README.md — build target and usage documentation

Core changes to existing files

  • examples/models/llama/model_args.py — add headwise_attn_output_gate and rope_parameters fields
  • examples/models/llama/feed_forward.py — FeedForward now accepts act_fn parameter (was hardcoded to SiLU)
  • examples/models/llama/attention.py — headwise attention output gate (dim→n_heads, broadcast over head_dim), per-layer is_sliding detection, RingKVCache for sliding-window layers, mutual-exclusion validation for gate flags
  • examples/models/llama/llama_transformer.py — per-layer-type RoPE via _build_ropes() helper, freqs_by_type dispatch in _forward_layers(), act_fn plumbed to FeedForward
  • examples/models/llama/export_llama_lib.py — register spark_x2_5_1_7b and spark_x2_5_4b in EXECUTORCH_DEFINED_MODELS, HUGGING_FACE_REPO_IDS, and weight-conversion dispatch
  • extension/llm/export/config/llm_config.py — add spark_x2_5_1_7b and spark_x2_5_4b to ModelType enum

Example export

python -m extension.llm.export.export_llm \
  --config examples/models/spark_x2_5/config/spark_x2_5_xnnpack_q8da4w.yaml \
  +base.model_class="spark_x2_5_1_7b" \
  +base.params="examples/models/spark_x2_5/config/spark_x2_5_1_7b_config.json" \
  +export.output_name="spark_x2_5_1_7b_8da4w.pte"

Test plan

Unit tests

pytest examples/models/spark_x2_5/test_spark_x2_5.py -v
# 3 passed

Model construction + forward pass (both variants)

from executorch.examples.models.llama.model_args import ModelArgs
from executorch.examples.models.llama.llama_transformer import construct_transformer

# 1.7B: 28 layers, dim=2048
# 4B:   36 layers, dim=2560
# Both: prefill + multi-step decode with KV cache succeed
# Layer types: sliding→RingKVCache, full→KVCache
# Per-layer RoPE: 2 distinct RoPE instances (full_attention, sliding_attention)

Sharded checkpoint conversion

# Spark-X2.5-1.7B uses 2 shards, 4B uses 5 shards
# Verified: 283 keys converted correctly, QKV split matches, tied embeddings preserved

Lint

flake8 examples/models/spark_x2_5/ examples/models/llama/model_args.py \
    examples/models/llama/feed_forward.py examples/models/llama/attention.py \
    examples/models/llama/llama_transformer.py examples/models/llama/export_llama_lib.py \
    extension/llm/export/config/llm_config.py --max-line-length=120
# No warnings

This PR was authored with AI assistance (Claude Code).

cc @mergennachin @iseeyuan @lucylq @helunwencser @tarun292 @kimishpatel @jackzhxng @larryliu0820 @cccclai @digantdesai

@pytorch-bot

pytorch-bot Bot commented Sep 16, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/executorch/22865

Note: Links to docs will display an error until the docs builds have been completed.

❌ 4 Pending, 1 Unclassified Failure

As of commit cd0396c with merge base b253af8 (image):

UNCLASSIFIED FAILURE - DrCI could not classify the following job because the workflow did not run on the merge base. The failure may be pre-existing on trunk or introduced by this PR:

  • MLX / test-mlx / test-mlx (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
    ##[error]fatal: unable to access 'https://github.com/nlohmann/json.git/': Failed to connect to github.com port 443 after 75007 ms: Couldn't connect to server

This comment was automatically generated by Dr. CI and updates every 15 minutes.

@meta-cla

meta-cla Bot commented Sep 16, 2026

Copy link
Copy Markdown

Hi @dongjiang1989!

Thank you for your pull request and welcome to our community.

Action Required

In order to merge any pull request (code, docs, etc.), we require contributors to sign our Contributor License Agreement, and we don't seem to have one on file for you.

Process

In order for us to review and merge your suggested changes, please sign at https://code.facebook.com/cla. If you are contributing on behalf of someone else (eg your employer), the individual CLA may not be sufficient and your employer may need to sign the corporate CLA.

Once the CLA is signed, our tooling will perform checks and validations. Afterwards, the pull request will be tagged with CLA signed. The tagging process may take up to 1 hour after signing. Please give it that time before contacting us about it.

If you have received this in error or have any questions, please contact us at cla@meta.com. Thanks!

@linux-foundation-easycla

linux-foundation-easycla Bot commented Sep 16, 2026 •

Copy link
Copy Markdown

CLA Signed
The committers listed above are authorized under a signed CLA.

  • ✅ login: dongjiang1989 / name: dongjiang1989 (c825ccf)

@github-actions

Copy link
Copy Markdown

This PR needs a release notes: label

If your change should be included in the release notes (i.e. would users of this library care about this change?), please use a label starting with release notes:. This helps us keep track and include your important work in the next release notes.

To add a label, you can comment to pytorchbot, for example
@pytorchbot label "release notes: none"

For more information, see
https://github.com/pytorch/pytorch/wiki/PyTorch-AutoLabel-Bot#why-categorize-for-release-notes-and-how-does-it-work.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. label Sep 16, 2026
@dongjiang1989 dongjiang1989 changed the title Add Spark-X2.5 model example with hybrid attention support feat(example): Add Spark-X2.5 model example with hybrid attention support Sep 16, 2026
@nil-is-all nil-is-all added module: examples Issues related to demos under examples/ module: llm Issues related to LLM examples and apps, and to the extensions/llm/ code labels Sep 16, 2026
@dongjiang1989

Copy link
Copy Markdown
Contributor Author

Fixed lintrunner error

@dongjiang1989

Copy link
Copy Markdown
Contributor Author

Thx @nil-is-all Please re-check it

@nil-is-all

Copy link
Copy Markdown
Contributor

Thx @nil-is-all Please re-check it

Thanks for the ping, re-running CI

@dongjiang1989

dongjiang1989 commented Sep 18, 2026 •

Copy link
Copy Markdown
Contributor Author

Thx @nil-is-all Please re-check it

Thanks for the ping, re-running CI

All CI pass @nil-is-all

@mergennachin

Copy link
Copy Markdown
Contributor

I checked bac008b and found several issues that need attention:

  1. MLX export drops the output gate and sliding-window behavior. The MLX preset runs transform_attention_mha_to_mlx, but MLXAttentionMHA.from_attention_mha neither copies/applies og nor preserves the ring cache. With a small model and zero gate weights, the converted attention output is exactly twice the expected output, since the sigmoid gate should multiply it by 0.5. This needs support in the MLX transformation before enabling the preset.

  2. The CoreML transformation creates incompatible cache and mask shapes. With the CoreML preset, replace_kv_cache_with_coreml_kv_cache replaces RingKVCache with KVCacheCoreML, losing ring handling. At the default context length of 128, the sliding layers retain a 1,024-slot cache but use a 128-column mask. A single-token forward after transformation fails with a 1024 versus 128 dimension mismatch. The conversion needs to preserve sliding-window semantics.

  3. The documented 2,048-token export exceeds the ring cache's prefill limit. The extended-context example sets max_seq_length=2048, but the 512-token window creates a 1,024-slot ring cache. A 1,025-token prefill hits the cache update assertion, and dynamic export with the larger sequence bound fails its shape constraints. Prefill needs to be constrained/chunked, or the cache capacity needs to account for the requested prefill size independently of the attention window.

  4. The uncached path does not apply the sliding-window mask. The new is_sliding handling changes cache construction, but with use_kv_cache=False, attention still uses the full triangular mask. In a 600-token test, outputs match a correctly windowed reference through the first 512 tokens, then diverge. Sliding layers need a windowed mask in the uncached path too.

  5. The example prompts don't match Spark's chat template. The runner examples use <|user|>, <|end|>, and <|assistant|>, while the published template uses <|User|>, <|Bot|>, and Spark's sentence delimiters. Both runners consume the supplied prompt directly. Please generate the example prompt from the model's chat template.

The three new tests and 35 shared transformer tests pass. The checks above used reduced models, source transformations, and torch.export; I haven't validated a full .pte export or device inference because compiled bindings aren't available in the isolated checkout. Tests for transformed attention parity and prompts that cross the sliding-window boundary would help cover these cases.

@dongjiang1989

dongjiang1989 commented Sep 20, 2026 •

Copy link
Copy Markdown
Contributor Author

Thanks @mergennachin for the thorough review. All 5 issues have been fixed:

  1. MLX export — Removed the MLX preset (spark_x2_5_mlx_4w.yaml). The MLX source transformation doesn't yet preserve og or the ring cache, so the preset produced incorrect results.

  2. CoreML export — Removed the CoreML preset (spark_x2_5_coreml_fp32.yaml). replace_kv_cache_with_coreml_kv_cache replaces RingKVCache with KVCacheCoreML, losing the ring handling.

  3. Extended-context limit — Updated the README example from max_seq_length=2048 to max_seq_length=1024 with a note explaining the 2×sliding_window ring buffer bound.

  4. Uncached sliding-window mask — Fixed in attention.py: when is_sliding=True and sliding_window is set, the causal mask now additionally masks out positions outside the window. Verified with a 600-token forward pass (sliding layers correctly attend to only 512 tokens per row; full-attention layers attend to all tokens).

  5. Chat template — Updated README runner examples to use Spark-X2.5's actual template (<|start▁of▁sentence|><|System|>...<|end▁of▁sentence|><|start▁of▁sentence|><|User|>...<|end▁of▁sentence|><|start▁of▁sentence|><|Bot|></think>).

Tests updated: replaced the removed MLX test with XNNPack q8da4w validation. All 3 tests pass.

Regarding future support: MLX and CoreML presets can be re-added once their source transformations are extended to handle the headwise attention output gate and ring-buffer KV cache. Happy to help with those when the time comes.

@dongjiang1989

Copy link
Copy Markdown
Contributor Author

cc @nil-is-all @mergennachin PTAL, Thx

Comment thread examples/models/llama/attention.py Outdated
@dongjiang1989
dongjiang1989 force-pushed the add-spark-x2.5 branch 2 times, most recently from 457397f to c58424f Compare September 22, 2026 08:54
@dongjiang1989

Copy link
Copy Markdown
Contributor Author

cc @nil-is-all @mergennachin Please recheck it, thanks

@executorch-triage executorch-triage Bot added the community: contribution PRs coming from community (excluding hardware partners) label Sep 22, 2026
@dongjiang1989

Copy link
Copy Markdown
Contributor Author

@nil-is-all This CI fail

Error: libomp: A `brew install libomp` process has already locked /opt/homebrew/Cellar/cmake.

Add examples/models/spark_x2_5/ with support for the Spark-X2.5-1.7B and
Spark-X2.5-4B models from XHToken. These are hybrid-attention LLMs (3:1
sliding-window to full-attention layers) with native 1M-token context,
per-layer-type RoPE configs, headwise sigmoid attention gates, and GELU
activation.

New module:
- convert_weights.py: HF safetensors to Meta format conversion, including
  fused QKV projection splitting, sharded checkpoint support, and tied
  embeddings handling.
- config/: JSON configs for 1.7B (28 layers) and 4B (36 layers) variants,
  plus XNNPack (fp32/q8da4w), CoreML, and MLX export configs.
- test_spark_x2_5.py: config validation and model registration tests.
- BUCK, README.md: build and usage documentation.

Core changes to support Spark-X2.5 architecture:
- model_args.py: Add headwise_attn_output_gate and rope_parameters fields.
- feed_forward.py: FeedForward now accepts act_fn parameter (was hardcoded
  to SiLU).
- attention.py: Add headwise attention output gate (dim->n_heads broadcast
  over head_dim), per-layer is_sliding detection, RingKVCache for
  sliding-window layers.
- llama_transformer.py: Per-layer-type RoPE construction via _build_ropes(),
  freqs_by_type dispatch in _forward_layers(), act_fn plumbed to FeedForward.
- export_llama_lib.py, llm_config.py: Register spark_x2_5_1_7b and
  spark_x2_5_4b model types.

This commit was authored with AI assistance (Claude Code).

Signed-off-by: dongjiang1989 <dongjiang1989@126.com>
Apply black formatting to fix UFMT lint failures in CI:
- convert_weights.py: single-line FileNotFoundError
- test_spark_x2_5.py: break long assert line
- llama_transformer.py: break long FeedForward and _forward_layers calls

This commit was authored with AI assistance (Claude Code).

Signed-off-by: dongjiang1989 <dongjiang1989@126.com>
dongjiang1989 and others added 2 commits September 24, 2026 12:40
…pdate docs

1. Uncached path: apply sliding-window mask for sliding_attention layers
   (previously used full triangular mask, causing divergence past 512 tokens).
2. Remove MLX and CoreML export presets: their source transformations do not
   yet preserve the headwise attention output gate or ring-buffer KV cache.
3. README: correct extended-context example (max_seq_length=1024, bounded
   by the 2*sliding_window ring cache), and use Spark-X2.5's actual chat
   template (<|start▁of▁sentence|><|User|>... instead of <|user|>...).
4. Update tests: replace removed MLX test with XNNPack q8da4w validation.

This commit was authored with AI assistance (Claude Code).

Signed-off-by: dongjiang1989 <dongjiang1989@126.com>
Co-authored-by: Mergen Nachin <mnachin@meta.com>
@meta-codesync

meta-codesync Bot commented Sep 25, 2026

Copy link
Copy Markdown
Contributor

@nil-is-all has imported this pull request. If you are a Meta employee, you can view this in D121856543.

@nil-is-all
nil-is-all self-requested a review September 28, 2026 15:38

@nil-is-all nil-is-all left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Ran internal CI, all passed. Good to merge

@nil-is-all
nil-is-all merged commit abedcd2 into pytorch:main Sep 28, 2026
250 of 251 checks passed
@dongjiang1989
dongjiang1989 deleted the add-spark-x2.5 branch September 29, 2026 02:02
@dongjiang1989

Copy link
Copy Markdown
Contributor Author

Ran internal CI, all passed. Good to merge

@nil-is-all Thx

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. community: contribution PRs coming from community (excluding hardware partners) module: examples Issues related to demos under examples/ module: llm Issues related to LLM examples and apps, and to the extensions/llm/ code

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants