Download configuration_parallax.py from YifeiZuo/attention-0.6b-adamw-wsd-step1000: direct link, hf CLI and curl.
- Browser
- Download file 3.72 kB
-
https://huggingface.co/YifeiZuo/attention-0.6b-adamw-wsd-step1000/resolve/main/configuration_parallax.py
- Command line
-
hf download hf://YifeiZuo/attention-0.6b-adamw-wsd-step1000/configuration_parallax.py
-
curl -L -o configuration_parallax.py https://huggingface.co/YifeiZuo/attention-0.6b-adamw-wsd-step1000/resolve/main/configuration_parallax.py
3.72 kB
| """Parallax HF config, modular-style (inherits Qwen3Config). | |
| This is a second, separate HF interface for the Parallax model family. Unlike | |
| the self-contained ``torchtitan.hf_modeling.parallax`` package, this version | |
| inherits from ``transformers.models.qwen3.configuration_qwen3.Qwen3Config`` | |
| so it picks up all of Qwen3's modern conventions for free (strict dataclass, | |
| ``rope_parameters``, ``layer_types``, tp/pp plans, etc.). | |
| The only Parallax-specific additions are the four fields at the bottom that | |
| control the centered amortized local linear attention op. When | |
| ``attn_type == "sdpa"`` this config describes a plain Qwen3-equivalent model. | |
| """ | |
| from __future__ import annotations | |
| from transformers.models.qwen3.configuration_qwen3 import Qwen3Config | |
| class ParallaxConfig(Qwen3Config): | |
| r"""Config for the Parallax model (Qwen3-derived). | |
| Args: | |
| attn_type (`str`, *optional*, defaults to `"parallax"`): | |
| Attention implementation. Must be one of ``"sdpa"`` or | |
| ``"parallax"``. ``"sdpa"`` uses the plain Qwen3 causal attention | |
| path (softmax attention). ``"parallax"`` uses the centered | |
| amortized local linear attention op and requires the extra | |
| ``r_proj`` / ``r_norm`` per-layer weights. | |
| use_centering (`bool`, *optional*, defaults to `True`): | |
| Whether the amortized-LLA math uses the centered variant | |
| (``s + c*s - r_term``). Currently only the centered variant is | |
| implemented; kept for config-roundtrip parity with the training | |
| code, which also has a non-centered path. | |
| amortized_rho_scale (`float`, *optional*, defaults to `1.0`): | |
| Multiplicative scale applied to the ``r`` tensor after its | |
| rotary embedding, mirroring the training-time | |
| ``amortized_rho_scale`` hyperparameter. | |
| learn_amortized_rho_scale (`bool`, *optional*, defaults to `False`): | |
| Whether the model learns a per-head scalar ``a_proj`` modulating | |
| ``r``. When enabled, adds an ``a_proj: Linear(hidden, n_heads)`` | |
| whose sigmoid output gates ``r`` per head, mirroring the | |
| training-time ``wa`` module in torchtitan. | |
| rope_r (`str`, *optional*, defaults to `"correct"`): | |
| How to apply rotary positional embeddings to R. | |
| ``"correct"`` — standard RoPE on R (matches fixed torchtitan code). | |
| ``"none"`` — skip RoPE on R entirely. | |
| ``"buggy"`` — reproduce the old torchtitan bug where transpose | |
| was applied before RoPE, causing positions 0..H-1 to be applied | |
| along the heads axis instead of the sequence axis. Use this to | |
| faithfully evaluate checkpoints trained with the buggy code. | |
| """ | |
| model_type = "parallax" | |
| # ---- Parallax-specific extras ---- | |
| # "sdpa" : plain causal self-attention (identical to Qwen3). | |
| # "parallax": centered amortized local linear attention (adds r_proj + r_norm | |
| # and replaces softmax attention with the centered LLA op). | |
| attn_type: str = "parallax" | |
| use_centering: bool = True | |
| amortized_rho_scale: float = 1.0 | |
| learn_amortized_rho_scale: bool = False | |
| rope_r: str = "correct" | |
| def __post_init__(self, **kwargs): | |
| if self.attn_type not in ("sdpa", "parallax"): | |
| raise ValueError( | |
| f"attn_type must be 'sdpa' or 'parallax', got {self.attn_type!r}" | |
| ) | |
| if self.rope_r not in ("correct", "none", "buggy"): | |
| raise ValueError( | |
| f"rope_r must be 'correct', 'none', or 'buggy', got {self.rope_r!r}" | |
| ) | |
| super().__post_init__(**kwargs) | |
| __all__ = ["ParallaxConfig"] | |