Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 8 additions & 2 deletions examples/models/llama/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -401,6 +401,10 @@ def __init__(

self.layer_id = layer_id
self.rope = rope
# NoPE (no positional encoding): layers listed in args.no_rope_layers skip
# RoPE entirely (e.g. Llama4 global-attention layers). Resolved at
# construction time so torch.export folds the branch structurally.
self.use_rope: bool = layer_id not in (args.no_rope_layers or [])

causal_mask = torch.tril(
torch.ones(
Expand Down Expand Up @@ -498,7 +502,8 @@ def _prepare_qkv_shared(
q = q * self.scale_query_by

# Apply RoPE to Q only (K already has RoPE from donor layer)
q, _ = self.rope.forward(q, q, freqs_cos, freqs_sin)
if self.use_rope:
q, _ = self.rope.forward(q, q, freqs_cos, freqs_sin)
q = q.transpose(1, 2)

if self.use_qk_norm and not self.qk_norm_before_rope:
Expand Down Expand Up @@ -532,7 +537,8 @@ def _prepare_qkv(
q = q * self.scale_query_by
k = self.k_norm_fn(k)

q, k = self.rope.forward(q, k, freqs_cos, freqs_sin)
if self.use_rope:
q, k = self.rope.forward(q, k, freqs_cos, freqs_sin)

q = q.transpose(1, 2) # (bs, n_local_heads, seqlen, head_dim)
k = k.transpose(1, 2)
Expand Down
3 changes: 3 additions & 0 deletions examples/models/llama/model_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -119,6 +119,9 @@ class ModelArgs:
no_rope_layer_interval: Optional[int] = (
None # Interval at which to skip RoPE. From Rope to Nope and Back Again: A New Hybrid Attention Strategy (https://huggingface.co/papers/2501.18795).
)
no_rope_layers: Optional[list[int]] = (
None # Explicit layer indices that skip RoPE (NoPE). Takes precedence over no_rope_layer_interval; used for Llama4 global (NoPE) layers.
)
partial_rotary_factor: float = 1.0
rope_theta: Optional[float] = (
None # The official name to override self.rope_freq_base.
Expand Down
16 changes: 13 additions & 3 deletions examples/models/llama/static_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -874,6 +874,9 @@ def __init__(
self._init_wo(config)
self.rope = _Rope(rope.params)
self.layer_id = layer_id
# NoPE: when False this layer skips RoPE (e.g. Llama4 global layers).
# Plain Python bool set at construction so torch.export folds it structurally.
self.use_rope: bool = kwargs.get("use_rope", True)
self._init_qk_norms(config, is_kv_shared_layer)

def _init_wo(self, config: ModelArgs) -> None:
Expand Down Expand Up @@ -966,6 +969,9 @@ def from_attention_mha(
scale_query_by=getattr(other, "scale_query_by", 1.0),
)

# Preserve NoPE: copy the source layer's use_rope so global (NoPE) layers
# continue to skip RoPE after conversion to StaticAttention.
kwargs.setdefault("use_rope", getattr(other, "use_rope", True))
instance = cls(
config=config,
layer_id=other.layer_id,
Expand Down Expand Up @@ -1123,6 +1129,8 @@ def _apply_rope(self, qs, ks, freqs_cos, freqs_sin):
freqs_cos (list): List of cosine frequencies.
freqs_sin (list): List of sine frequencies.
"""
if not self.use_rope:
return qs, ks
qs = [self.rope(q, freqs_cos, freqs_sin) for q in qs]
if ks is not None:
ks = [self.rope(k, freqs_cos, freqs_sin) for k in ks]
Expand Down Expand Up @@ -1339,7 +1347,8 @@ def _forward_mha(
if self.use_qk_norm and self.qk_norm_before_rope:
q = self.q_norm(q) * self.scale_query_by

q = self.rope(q, freqs_cos, freqs_sin)
if self.use_rope:
q = self.rope(q, freqs_cos, freqs_sin)

if self.use_qk_norm and not self.qk_norm_before_rope:
q = self.q_norm(q) * self.scale_query_by
Expand All @@ -1354,8 +1363,9 @@ def _forward_mha(
q = self.q_norm(q) * self.scale_query_by
k = self.k_norm(k)

q = self.rope(q, freqs_cos, freqs_sin)
k = self.rope(k, freqs_cos, freqs_sin)
if self.use_rope:
q = self.rope(q, freqs_cos, freqs_sin)
k = self.rope(k, freqs_cos, freqs_sin)

if self.use_qk_norm and not self.qk_norm_before_rope:
q = self.q_norm(q) * self.scale_query_by
Expand Down
Loading