Skip to content

Commit abedcd2

Browse files
feat(example): Add Spark-X2.5 model example with hybrid attention support (pytorch#22865)
## Summary Add `examples/models/spark_x2_5/` with support for the [Spark-X2.5-1.7B and Spark-X2.5-4B models](https://huggingface.co/collections/XHToken/spark-x25) 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) ```python 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 --------- Signed-off-by: dongjiang1989 <dongjiang1989@126.com> Co-authored-by: Mergen Nachin <mnachin@meta.com>
1 parent c58d35e commit abedcd2

16 files changed

Lines changed: 615 additions & 15 deletions

‎examples/models/llama/BUCK‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@ fbcode_target(_kind = runtime.python_library,
3232
":transformer_modules",
3333
"//caffe2:torch",
3434
"//executorch/examples/models/lfm2:lfm2",
35+
"//executorch/examples/models/spark_x2_5:spark_x2_5",
3536
],
3637
)
3738

‎examples/models/llama/attention.py‎

Lines changed: 37 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -418,7 +418,17 @@ def __init__(
418418
self.enable_dynamic_shape = args.enable_dynamic_shape
419419
self.scale_query_by = args.scale_query_by
420420
self.use_attn_o_gate = args.use_attn_o_gate
421+
self.headwise_attn_output_gate = args.headwise_attn_output_gate
422+
if self.use_attn_o_gate and self.headwise_attn_output_gate:
423+
raise ValueError(
424+
"use_attn_o_gate and headwise_attn_output_gate are mutually exclusive"
425+
)
421426
self.use_attn_o_norm = args.use_attn_o_norm
427+
self.is_sliding = (
428+
args.layer_types is not None
429+
and layer_id < len(args.layer_types)
430+
and args.layer_types[layer_id] == "sliding_attention"
431+
)
422432
q_out_dim = self.n_heads * self.head_dim * (2 if self.use_q_gate else 1)
423433

424434
# YOCO: Determine if this is a KV shared layer (receives shared KV from donor).
@@ -447,6 +457,10 @@ def __init__(
447457
device="cpu",
448458
)
449459
)
460+
# Sliding-window layers: additionally mask out positions outside the window.
461+
if self.is_sliding and args.sliding_window:
462+
window = args.sliding_window
463+
causal_mask = causal_mask.triu(diagonal=1 - window)
450464
self.register_buffer("mask", causal_mask, persistent=False)
451465

452466
if self.use_kv_cache:
@@ -481,6 +495,8 @@ def _init_norms(self, args: ModelArgs) -> None:
481495
self.o_norm = ScalelessRMSNorm(self.head_dim, eps=args.norm_eps)
482496
if self.use_attn_o_gate:
483497
self.og = nn.Linear(args.dim, self.n_heads * self.head_dim, bias=False)
498+
if self.headwise_attn_output_gate:
499+
self.og = nn.Linear(args.dim, self.n_local_heads, bias=False)
484500

485501
def _init_projections(self, args: ModelArgs, q_out_dim: int) -> None:
486502
"""Initialize Q/K/V/O projection layers."""
@@ -509,13 +525,22 @@ def _init_projections(self, args: ModelArgs, q_out_dim: int) -> None:
509525
def _init_kv_cache(self, args: ModelArgs) -> None:
510526
"""Initialize KV cache (only for non-shared layers)."""
511527
if self.has_kv_weights:
512-
self.kv_cache = KVCache(
513-
args.max_batch_size,
514-
args.max_context_len,
515-
self.n_kv_heads,
516-
self.head_dim,
517-
args.enable_dynamic_shape,
518-
)
528+
if self.is_sliding and args.sliding_window:
529+
self.kv_cache = RingKVCache(
530+
args.max_batch_size,
531+
args.sliding_window,
532+
self.n_kv_heads,
533+
self.head_dim,
534+
args.enable_dynamic_shape,
535+
)
536+
else:
537+
self.kv_cache = KVCache(
538+
args.max_batch_size,
539+
args.max_context_len,
540+
self.n_kv_heads,
541+
self.head_dim,
542+
args.enable_dynamic_shape,
543+
)
519544
else:
520545
self.kv_cache = None
521546

@@ -679,6 +704,11 @@ def _apply_output_transforms(
679704
og = self.og(x).view(bsz, seqlen, self.n_local_heads, self.head_dim)
680705
output_4d = torch.sigmoid(og) * output_4d
681706
output = output_4d.reshape(bsz, seqlen, -1)
707+
if self.headwise_attn_output_gate:
708+
output_4d = output.view(bsz, seqlen, self.n_local_heads, self.head_dim)
709+
og = self.og(x).unsqueeze(-1).to(output_4d.dtype)
710+
output_4d = torch.sigmoid(og) * output_4d
711+
output = output_4d.reshape(bsz, seqlen, -1)
682712
if gate is not None:
683713
output = output * torch.sigmoid(gate)
684714
return output

‎examples/models/llama/export_llama_lib.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -121,6 +121,8 @@
121121
"lfm2_1_2b", # hybrid
122122
"lfm2_5_350m", # hybrid
123123
"lfm2_5_1_2b", # hybrid
124+
"spark_x2_5_1_7b", # hybrid
125+
"spark_x2_5_4b", # hybrid
124126
]
125127
TORCHTUNE_DEFINED_MODELS = ["llama3_2_vision"]
126128
HUGGING_FACE_REPO_IDS = {
@@ -141,6 +143,8 @@
141143
"lfm2_1_2b": "LiquidAI/LFM2-1.2B",
142144
"lfm2_5_350m": "LiquidAI/LFM2.5-350M",
143145
"lfm2_5_1_2b": "LiquidAI/LFM2.5-1.2B-Instruct",
146+
"spark_x2_5_1_7b": "XHToken/Spark-X2.5-1.7B",
147+
"spark_x2_5_4b": "XHToken/Spark-X2.5-4B",
144148
}
145149

146150

@@ -718,6 +722,8 @@ def export_llama( # noqa: C901
718722
from executorch.examples.models.smollm2 import convert_weights
719723
elif model_name.startswith("lfm2"):
720724
from executorch.examples.models.lfm2 import convert_weights
725+
elif model_name.startswith("spark_x2_5"):
726+
from executorch.examples.models.spark_x2_5 import convert_weights
721727
else:
722728
raise ValueError(
723729
f"Converting weights to meta format for {model_name} is not yet supported"

‎examples/models/llama/feed_forward.py‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,14 +6,15 @@
66

77

88
class FeedForward(nn.Module):
9-
def __init__(self, dim: int, hidden_dim: int):
9+
def __init__(self, dim: int, hidden_dim: int, act_fn=F.silu):
1010
super().__init__()
1111
self.w1 = nn.Linear(dim, hidden_dim, bias=False)
1212
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
1313
self.w3 = nn.Linear(dim, hidden_dim, bias=False)
14+
self.act_fn = act_fn
1415

1516
def forward(self, x):
16-
return self.w2(F.silu(self.w1(x)) * self.w3(x))
17+
return self.w2(self.act_fn(self.w1(x)) * self.w3(x))
1718

1819

1920
class LoRAFeedForward(nn.Module):

‎examples/models/llama/llama_transformer.py‎

Lines changed: 64 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -215,7 +215,11 @@ def __init__(
215215
):
216216
self.feed_forward = LoRAFeedForward(args.dim, args.hidden_dim, args)
217217
else:
218-
self.feed_forward = FeedForward(dim=args.dim, hidden_dim=args.hidden_dim)
218+
self.feed_forward = FeedForward(
219+
dim=args.dim,
220+
hidden_dim=args.hidden_dim,
221+
act_fn=args.act_fn.get_function(),
222+
)
219223

220224
if isinstance(self.attention, AttentionSkip):
221225
self.attention_norm = nn.Identity()
@@ -357,6 +361,7 @@ def __init__(self, params: ModelArgs, layers: nn.ModuleList, rope: Rope):
357361
self.output_prune_map = params.output_prune_map
358362
# YOCO (You Only Cache Once) KV sharing configuration.
359363
self.num_kv_shared_layers = params.num_kv_shared_layers
364+
self.layer_types = params.layer_types
360365

361366
def _forward_layers(
362367
self,
@@ -365,6 +370,7 @@ def _forward_layers(
365370
freqs_sin: torch.Tensor,
366371
attn_options_: Dict,
367372
seqlen: int,
373+
freqs_by_type: Optional[Dict[str, Tuple[torch.Tensor, torch.Tensor]]] = None,
368374
) -> Tuple[torch.Tensor, Optional[Any]]:
369375
"""Run transformer layers with YOCO KV sharing support."""
370376
attn_options_update = None
@@ -383,7 +389,14 @@ def _forward_layers(
383389
if donor_idx in shared_kv:
384390
attn_options_["shared_kv"] = shared_kv[donor_idx]
385391

386-
h, attn_options_update = layer(h, freqs_cos, freqs_sin, attn_options_)
392+
# Per-layer-type RoPE: select freqs based on layer type when available.
393+
l_cos, l_sin = freqs_cos, freqs_sin
394+
if freqs_by_type is not None and self.layer_types is not None:
395+
layer_type = self.layer_types[layer_idx]
396+
if layer_type in freqs_by_type:
397+
l_cos, l_sin = freqs_by_type[layer_type]
398+
399+
h, attn_options_update = layer(h, l_cos, l_sin, attn_options_)
387400

388401
if _is_kv_donor_layer(layer_idx, self.n_layers, self.num_kv_shared_layers):
389402
assert (
@@ -425,10 +438,23 @@ def forward(
425438
attn_options.get("input_pos"), seqlen
426439
)
427440

441+
# Compute per-layer-type freqs when per-layer RoPE is configured.
442+
freqs_by_type = None
443+
if hasattr(self, "ropes"):
444+
input_pos = attn_options.get("input_pos")
445+
freqs_by_type = {
446+
lt: r.get_freqs(input_pos, seqlen) for lt, r in self.ropes.items()
447+
}
448+
428449
attn_options_ = attn_options.copy() if attn_options is not None else {}
429450

430451
h, attn_options_update = self._forward_layers(
431-
h, freqs_cos, freqs_sin, attn_options_, seqlen
452+
h,
453+
freqs_cos,
454+
freqs_sin,
455+
attn_options_,
456+
seqlen,
457+
freqs_by_type=freqs_by_type,
432458
)
433459

434460
if not self.generate_full_logits:
@@ -467,11 +493,35 @@ def forward(
467493
return logits
468494

469495

496+
def _build_ropes(model_args: ModelArgs) -> Tuple[Rope, Dict[str, Rope]]:
497+
"""Build Rope instances, creating per-layer-type ropes when rope_parameters is set.
498+
499+
Returns (default_rope, ropes_by_type). ropes_by_type is empty when no
500+
per-layer-type configuration is provided.
501+
"""
502+
import copy as _copy
503+
504+
if not model_args.rope_parameters:
505+
return Rope(model_args), {}
506+
507+
ropes: Dict[str, Rope] = {}
508+
for layer_type, rope_params in model_args.rope_parameters.items():
509+
rope_args = _copy.copy(model_args)
510+
if "rope_theta" in rope_params:
511+
rope_args.rope_theta = rope_params["rope_theta"]
512+
rope_args.rope_freq_base = rope_params["rope_theta"]
513+
if "partial_rotary_factor" in rope_params:
514+
rope_args.partial_rotary_factor = rope_params["partial_rotary_factor"]
515+
ropes[layer_type] = Rope(rope_args)
516+
return next(iter(ropes.values())), ropes
517+
518+
470519
def construct_transformer(model_args: ModelArgs) -> Transformer:
471520
"""
472521
Construct a Transformer model from the given model arguments.
473522
"""
474-
rope = Rope(model_args)
523+
rope, ropes = _build_ropes(model_args)
524+
475525
if model_args.attention_type not in ATTENTION_REGISTRY:
476526
raise ValueError(
477527
f"Unknown attention type: {model_args.attention_type}. "
@@ -521,12 +571,20 @@ def construct_transformer(model_args: ModelArgs) -> Transformer:
521571
)
522572
layers.append(transformer_block)
523573
else:
574+
# Select per-layer-type RoPE when available.
575+
layer_rope = rope
576+
if ropes and model_args.layer_types:
577+
layer_type = model_args.layer_types[layer_id]
578+
layer_rope = ropes.get(layer_type, rope)
524579
attention = cls(
525-
model_args, layer_id, rope, **model_args.attention_kwargs
580+
model_args, layer_id, layer_rope, **model_args.attention_kwargs
526581
) # pyre-ignore[45]
527582
transformer_block = TransformerBlock(
528583
model_args, attention, layer_id=layer_id
529584
)
530585
layers.append(transformer_block)
531586

532-
return Transformer(model_args, layers, rope)
587+
transformer = Transformer(model_args, layers, rope)
588+
if ropes:
589+
transformer.ropes = torch.nn.ModuleDict(ropes)
590+
return transformer

‎examples/models/llama/model_args.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -126,6 +126,9 @@ class ModelArgs:
126126
local_rope_theta: Optional[float] = (
127127
None # For sliding window attention. e.g., gemma3-1b
128128
)
129+
rope_parameters: Optional[Dict[str, Dict[str, Any]]] = (
130+
None # Per-layer-type RoPE configs. e.g., {"full_attention": {"rope_theta": 5000000, "partial_rotary_factor": 0.25}}
131+
)
129132
rope_freq_base: float = 10000.0 # The base frequency for RoPE. Keep it for BC.
130133
use_scaled_rope: bool = False # Use scaled RoPE, introduced in llama3.1.
131134
rope_scale_factor: int = 8
@@ -184,6 +187,7 @@ class ModelArgs:
184187
normalize_tok_embeddings: bool = False
185188
scale_query_by: float = 1.0
186189
use_attn_o_gate: bool = False
190+
headwise_attn_output_gate: bool = False
187191
use_attn_o_norm: bool = False
188192
use_residual_gate: bool = False
189193
use_ffn_learnable_scales: bool = False

‎examples/models/spark_x2_5/BUCK‎

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,22 @@
1+
load("@fbcode_macros//build_defs:build_file_migration.bzl", "fbcode_target", "non_fbcode_target")
2+
# Any targets that should be shared between fbcode and xplat must be defined in
3+
# targets.bzl. This file can contain fbcode-only targets.
4+
5+
load("@fbsource//xplat/executorch/build:runtime_wrapper.bzl", "runtime")
6+
7+
oncall("executorch")
8+
9+
fbcode_target(_kind = runtime.python_library,
10+
name = "spark_x2_5",
11+
srcs = [
12+
"__init__.py",
13+
"convert_weights.py",
14+
],
15+
base_module = "executorch.examples.models.spark_x2_5",
16+
visibility = ["PUBLIC"],
17+
deps = [
18+
"//caffe2:torch",
19+
"//executorch/examples/models/llama:transformer_modules",
20+
"fbsource//third-party/pypi/safetensors:safetensors",
21+
],
22+
)

0 commit comments

Comments
 (0)