diff --git a/examples/models/llama/attention.py b/examples/models/llama/attention.py index d43533b5a70..336e7e6ef09 100644 --- a/examples/models/llama/attention.py +++ b/examples/models/llama/attention.py @@ -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( @@ -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: @@ -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) diff --git a/examples/models/llama/model_args.py b/examples/models/llama/model_args.py index a71b9857dbf..b14d68c3a99 100644 --- a/examples/models/llama/model_args.py +++ b/examples/models/llama/model_args.py @@ -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. diff --git a/examples/models/llama/static_attention.py b/examples/models/llama/static_attention.py index 8e985239651..ac966b3fc1a 100644 --- a/examples/models/llama/static_attention.py +++ b/examples/models/llama/static_attention.py @@ -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: @@ -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, @@ -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] @@ -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 @@ -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