Download modeling_parallax.py from YifeiZuo/attention-0.6b-adamw-wsd-step6500: direct link, hf CLI and curl.
- Browser
- Download file 16.1 kB
-
https://huggingface.co/YifeiZuo/attention-0.6b-adamw-wsd-step6500/resolve/main/modeling_parallax.py
- Command line
-
hf download hf://YifeiZuo/attention-0.6b-adamw-wsd-step6500/modeling_parallax.py
-
curl -L -o modeling_parallax.py https://huggingface.co/YifeiZuo/attention-0.6b-adamw-wsd-step6500/resolve/main/modeling_parallax.py
16.1 kB
| """Parallax HF modeling file, modular-style (inherits from Qwen3). | |
| This module is the second, separate HF interface for the Parallax family. | |
| It inherits almost everything from ``transformers.models.qwen3`` and only | |
| overrides the attention class to add the centered amortized local linear | |
| attention op (``attn_type="parallax"``). | |
| Design notes | |
| ------------ | |
| * All boilerplate — RMSNorm, MLP, rotary embeddings, decoder-layer scaffolding, | |
| the base ``PreTrainedModel`` class, ``Qwen3Model``, ``Qwen3ForCausalLM``, | |
| mask preparation, cache handling, ``GenerationMixin`` integration, etc. — | |
| is inherited from ``transformers.models.qwen3``. | |
| * The only new code is: | |
| - ``parallax_centered_attention_forward``: a pure-PyTorch attention kernel | |
| matching the math of ``torchtitan/models/flashlla/dev_ops/reference.py`` | |
| (``centered_dense_reference``), but batched over heads and taking an | |
| additive attention mask in the HF shape ``(B, 1, q_len, kv_len)``. | |
| - ``ParallaxAttention``: inherits from ``Qwen3Attention``, adds an | |
| ``r_proj`` / ``r_norm`` pair when ``config.attn_type == "parallax"``, | |
| and overrides ``forward`` to dispatch to the parallax kernel in that | |
| case. For ``attn_type == "sdpa"`` the forward simply delegates to | |
| ``Qwen3Attention.forward`` unchanged, so all modern HF attention | |
| backends (sdpa, flash_attention_2, flex_attention, eager) remain | |
| available for that path. | |
| - ``ParallaxDecoderLayer``, ``ParallaxModel``, ``ParallaxPreTrainedModel``, | |
| ``ParallaxForCausalLM``: thin subclasses that rewire parent classes to | |
| the ``ParallaxAttention`` / ``ParallaxConfig`` below. | |
| Why this file exists alongside ``torchtitan/hf_modeling/parallax/`` | |
| ------------------------------------------------------------------- | |
| The older ``torchtitan.hf_modeling.parallax`` package is fully self-contained | |
| (no imports from ``transformers.models.qwen3``). That makes it robust across | |
| transformers versions but also duplicates a lot of boilerplate that drifts | |
| from HF conventions. This ``parallax_qwen3`` package takes the opposite | |
| tradeoff: minimal code, surgical diffs from Qwen3, tracks HF updates for | |
| free, at the cost of depending on a stable ``transformers.models.qwen3`` API. | |
| Both produce numerically equivalent outputs on the same Parallax checkpoint | |
| weights; a parity test verifies this. | |
| """ | |
| from __future__ import annotations | |
| from typing import Optional | |
| import torch | |
| import torch.nn as nn | |
| from transformers.cache_utils import Cache | |
| from transformers.modeling_outputs import CausalLMOutputWithPast | |
| from transformers.models.qwen3.modeling_qwen3 import ( | |
| Qwen3Attention, | |
| Qwen3DecoderLayer, | |
| Qwen3ForCausalLM, | |
| Qwen3Model, | |
| Qwen3PreTrainedModel, | |
| Qwen3RMSNorm, | |
| apply_rotary_pos_emb, | |
| repeat_kv, | |
| rotate_half, | |
| ) | |
| from .configuration_parallax import ParallaxConfig | |
| # --------------------------------------------------------------------------- | |
| # Parallax centered attention kernel (pure PyTorch reference) | |
| # --------------------------------------------------------------------------- | |
| def parallax_centered_attention_forward( | |
| module: nn.Module, | |
| query: torch.Tensor, | |
| r: torch.Tensor, | |
| key: torch.Tensor, | |
| value: torch.Tensor, | |
| attention_mask: Optional[torch.Tensor], | |
| scaling: float, | |
| dropout: float = 0.0, | |
| **kwargs, | |
| ) -> tuple[torch.Tensor, None]: | |
| """Centered amortized local linear attention. | |
| This is a batched generalization of ``centered_dense_reference`` from | |
| ``torchtitan/models/flashlla/dev_ops/reference.py``: | |
| p = softmax((Q Kᵀ) * scale + mask) | |
| rk = R Kᵀ | |
| s = p V | |
| c = sum(p * rk, -1, keepdim=True) | |
| r_term = (p * rk) V | |
| out = s + c * s - r_term | |
| Shapes | |
| ------ | |
| query : (B, H, q_len, D) | |
| r : (B, H, q_len, D) — Parallax-specific extra projection | |
| key : (B, H_kv, kv_len, D) — GQA-compressed, expanded internally | |
| value : (B, H_kv, kv_len, D) | |
| attention_mask : (B, 1, q_len, kv_len) additive, 0 or -inf (HF convention) | |
| Returns | |
| ------- | |
| attn_output : (B, q_len, H, D) — transposed to match Qwen3 eager impl | |
| attn_weights : None — not exposed (centered attention has no | |
| meaningful "attention matrix" per head) | |
| """ | |
| # Expand GQA-compressed k/v to match q heads | |
| key = repeat_kv(key, module.num_key_value_groups) | |
| value = repeat_kv(value, module.num_key_value_groups) | |
| # Upcast to float32 for numerical stability (matches the training kernel, | |
| # which uses bfloat16 for math but accumulates in fp32 within the softmax). | |
| out_dtype = query.dtype | |
| qf = query.float() | |
| rf = r.float() | |
| kf = key.float() | |
| vf = value.float() | |
| qk = torch.einsum("bhqd,bhkd->bhqk", qf, kf) * scaling | |
| q_len = qk.shape[-2] | |
| kv_len = qk.shape[-1] | |
| if attention_mask is not None: | |
| # HF causal masks may cover a larger key span than our current | |
| # kv_len during decode; slice to match. | |
| qk = qk + attention_mask[..., :q_len, :kv_len].to(qk.dtype) | |
| elif q_len > 1: | |
| # HF's `create_causal_mask` returns None when the model's default | |
| # attention backend (e.g. sdpa) handles causality implicitly. Our | |
| # parallax kernel is hand-written and needs the explicit additive | |
| # mask, so build one here. Decode (q_len == 1) is naturally | |
| # unmasked: the single query can attend to every cached key. | |
| offset = kv_len - q_len | |
| row = torch.arange(q_len, device=qk.device)[:, None] + offset | |
| col = torch.arange(kv_len, device=qk.device)[None, :] | |
| causal = row >= col | |
| bias = torch.zeros(q_len, kv_len, dtype=qk.dtype, device=qk.device) | |
| bias = bias.masked_fill(~causal, float("-inf")) | |
| qk = qk + bias | |
| p = torch.softmax(qk, dim=-1) | |
| rk = torch.einsum("bhqd,bhkd->bhqk", rf, kf) | |
| s = torch.einsum("bhqk,bhkd->bhqd", p, vf) | |
| c = (p * rk).sum(dim=-1, keepdim=True) | |
| r_term = torch.einsum("bhqk,bhkd->bhqd", p * rk, vf) | |
| attn_output = (s + c * s - r_term).to(out_dtype) | |
| # Match Qwen3 eager_attention_forward: return transposed (B, q_len, H, D) | |
| attn_output = attn_output.transpose(1, 2).contiguous() | |
| return attn_output, None | |
| # --------------------------------------------------------------------------- | |
| # Attention / DecoderLayer / Model / ForCausalLM | |
| # --------------------------------------------------------------------------- | |
| class ParallaxAttention(Qwen3Attention): | |
| """Qwen3 attention plus an optional Parallax (centered amortized LLA) path.""" | |
| def __init__(self, config: ParallaxConfig, layer_idx: int) -> None: | |
| super().__init__(config, layer_idx) | |
| self.attn_type = config.attn_type | |
| self.amortized_rho_scale = config.amortized_rho_scale | |
| self.learn_amortized_rho_scale = config.learn_amortized_rho_scale | |
| self.rope_r = getattr(config, "rope_r", "correct") | |
| self._buggy_rope_cos: torch.Tensor | None = None | |
| self._buggy_rope_sin: torch.Tensor | None = None | |
| if self.attn_type == "parallax": | |
| # Extra projection + RMSNorm for the amortized local linear attn | |
| self.r_proj = nn.Linear( | |
| config.hidden_size, | |
| config.num_attention_heads * self.head_dim, | |
| bias=config.attention_bias, | |
| ) | |
| self.r_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps) | |
| # Learned per-head gate on r (sigmoid), mirroring torchtitan `wa`. | |
| if self.learn_amortized_rho_scale: | |
| self.a_proj = nn.Linear( | |
| config.hidden_size, | |
| config.num_attention_heads, | |
| bias=False, | |
| ) | |
| else: | |
| self.a_proj = None | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| position_embeddings: tuple[torch.Tensor, torch.Tensor], | |
| attention_mask: Optional[torch.Tensor], | |
| past_key_values: Optional[Cache] = None, | |
| **kwargs, | |
| ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: | |
| if self.attn_type != "parallax": | |
| # Plain Qwen3 attention — delegate to the parent, which supports | |
| # sdpa / flash_attention_2 / flex_attention / eager backends. | |
| return super().forward( | |
| hidden_states=hidden_states, | |
| position_embeddings=position_embeddings, | |
| attention_mask=attention_mask, | |
| past_key_values=past_key_values, | |
| **kwargs, | |
| ) | |
| # ---- Parallax path ---- | |
| input_shape = hidden_states.shape[:-1] | |
| hidden_shape = (*input_shape, -1, self.head_dim) | |
| query_states = self.q_norm( | |
| self.q_proj(hidden_states).view(hidden_shape) | |
| ).transpose(1, 2) | |
| key_states = self.k_norm( | |
| self.k_proj(hidden_states).view(hidden_shape) | |
| ).transpose(1, 2) | |
| value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2) | |
| # r: same num_heads as q (not num_key_value_heads), rope applied. | |
| # Note hidden_shape uses -1 for the heads dim, so r_proj (which produces | |
| # num_attention_heads*head_dim) will correctly view as (..., n_heads, D). | |
| r_states = self.r_norm( | |
| self.r_proj(hidden_states).view(hidden_shape) | |
| ).transpose(1, 2) | |
| cos, sin = position_embeddings | |
| query_states, key_states = apply_rotary_pos_emb( | |
| query_states, key_states, cos, sin | |
| ) | |
| # Apply rope to r according to rope_r mode. | |
| if self.rope_r == "correct": | |
| r_states, _ = apply_rotary_pos_emb(r_states, r_states, cos, sin) | |
| elif self.rope_r == "buggy": | |
| # Reproduce the old torchtitan bug: reshape_for_broadcast read | |
| # (B, H, S, D) as (bz, seqlen=H, _, head_dim=D), so it applied | |
| # positions 0..H-1 along the heads axis. Every token in the | |
| # sequence got the SAME rotation (the one for position h, where | |
| # h is the head index). This is a fixed per-head rotation that | |
| # does not vary with sequence position. | |
| # | |
| # We lazily compute cos/sin for positions 0..H-1 and cache them. | |
| # Shape: (1, H, 1, D) so head h gets rotation for position h, | |
| # broadcast across all sequence positions. | |
| if self._buggy_rope_cos is None or self._buggy_rope_cos.device != r_states.device: | |
| H = r_states.shape[1] | |
| # Reuse the same RoPE frequencies: extract from cos/sin for | |
| # positions 0..H-1. cos/sin are (B,1,T,D) during prefill | |
| # (T >= H always since T >= 1 and H = num_heads). | |
| # For robustness, recompute from config's rope_theta. | |
| head_dim = r_states.shape[-1] | |
| rope_theta = self.config.rope_parameters["rope_theta"] | |
| inv_freq = 1.0 / (rope_theta ** ( | |
| torch.arange(0, head_dim, 2, device=r_states.device, dtype=torch.float32) / head_dim | |
| )) | |
| pos = torch.arange(H, device=r_states.device, dtype=torch.float32) | |
| freqs = torch.outer(pos, inv_freq) # (H, D/2) | |
| emb = torch.cat([freqs, freqs], dim=-1) # (H, D) | |
| self._buggy_rope_cos = emb.cos().view(1, H, 1, head_dim).to(r_states.dtype) | |
| self._buggy_rope_sin = emb.sin().view(1, H, 1, head_dim).to(r_states.dtype) | |
| r_states = (r_states * self._buggy_rope_cos) + (rotate_half(r_states) * self._buggy_rope_sin) | |
| # else: rope_r == "none", skip RoPE on R entirely | |
| if self.a_proj is not None: | |
| # a: (B, n_heads, q_len), sigmoid per-head gate | |
| a = torch.sigmoid(self.a_proj(hidden_states)).transpose(1, 2).to(dtype=r_states.dtype) | |
| r_states = r_states * (self.amortized_rho_scale * a.unsqueeze(-1)) | |
| elif self.amortized_rho_scale != 1.0: | |
| r_states = r_states * self.amortized_rho_scale | |
| if past_key_values is not None: | |
| key_states, value_states = past_key_values.update( | |
| key_states, value_states, self.layer_idx | |
| ) | |
| attn_output, attn_weights = parallax_centered_attention_forward( | |
| self, | |
| query_states, | |
| r_states, | |
| key_states, | |
| value_states, | |
| attention_mask, | |
| dropout=0.0 if not self.training else self.attention_dropout, | |
| scaling=self.scaling, | |
| ) | |
| attn_output = attn_output.reshape(*input_shape, -1).contiguous() | |
| attn_output = self.o_proj(attn_output) | |
| return attn_output, attn_weights | |
| class ParallaxDecoderLayer(Qwen3DecoderLayer): | |
| def __init__(self, config: ParallaxConfig, layer_idx: int) -> None: | |
| # Call the grandparent (nn.Module) init to avoid Qwen3DecoderLayer | |
| # building a Qwen3Attention we'd immediately replace. | |
| nn.Module.__init__(self) | |
| self.hidden_size = config.hidden_size | |
| self.attention_type = config.layer_types[layer_idx] | |
| self.self_attn = ParallaxAttention(config=config, layer_idx=layer_idx) | |
| # Reuse the MLP + norms from Qwen3 via its own constructor logic. | |
| # Easiest is to rebuild them directly here to match Qwen3DecoderLayer. | |
| from transformers.models.qwen3.modeling_qwen3 import Qwen3MLP | |
| self.mlp = Qwen3MLP(config) | |
| self.input_layernorm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps) | |
| self.post_attention_layernorm = Qwen3RMSNorm( | |
| config.hidden_size, eps=config.rms_norm_eps | |
| ) | |
| class ParallaxPreTrainedModel(Qwen3PreTrainedModel): | |
| config: ParallaxConfig # type: ignore[assignment] | |
| config_class = ParallaxConfig | |
| base_model_prefix = "model" | |
| _no_split_modules = ["ParallaxDecoderLayer"] | |
| # Parallax path is hand-written and does not plug into ALL_ATTENTION_FUNCTIONS. | |
| # The sdpa path (delegated to super) supports the standard backends. | |
| _supports_flash_attn = True | |
| _supports_sdpa = True | |
| _supports_flex_attn = True | |
| class ParallaxModel(ParallaxPreTrainedModel, Qwen3Model): | |
| def __init__(self, config: ParallaxConfig) -> None: | |
| # Bypass Qwen3Model.__init__ (which builds Qwen3DecoderLayers) and | |
| # rebuild the module tree with ParallaxDecoderLayers. | |
| Qwen3PreTrainedModel.__init__(self, config) | |
| self.padding_idx = config.pad_token_id | |
| self.vocab_size = config.vocab_size | |
| self.embed_tokens = nn.Embedding( | |
| config.vocab_size, config.hidden_size, self.padding_idx | |
| ) | |
| self.layers = nn.ModuleList( | |
| [ | |
| ParallaxDecoderLayer(config, layer_idx) | |
| for layer_idx in range(config.num_hidden_layers) | |
| ] | |
| ) | |
| self.norm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps) | |
| # Qwen3Model constructs its rotary embedding from the config; reuse. | |
| from transformers.models.qwen3.modeling_qwen3 import Qwen3RotaryEmbedding | |
| self.rotary_emb = Qwen3RotaryEmbedding(config=config) | |
| self.gradient_checkpointing = False | |
| self.has_sliding_layers = "sliding_attention" in (config.layer_types or []) | |
| self.post_init() | |
| class ParallaxForCausalLM(ParallaxPreTrainedModel, Qwen3ForCausalLM): | |
| _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"} | |
| def __init__(self, config: ParallaxConfig) -> None: | |
| # Bypass Qwen3ForCausalLM.__init__ (which builds a Qwen3Model) and | |
| # rebuild with ParallaxModel. | |
| Qwen3PreTrainedModel.__init__(self, config) | |
| self.model = ParallaxModel(config) | |
| self.vocab_size = config.vocab_size | |
| self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) | |
| self.post_init() | |
| __all__ = [ | |
| "ParallaxAttention", | |
| "ParallaxDecoderLayer", | |
| "ParallaxForCausalLM", | |
| "ParallaxModel", | |
| "ParallaxPreTrainedModel", | |
| "parallax_centered_attention_forward", | |
| ] | |