"""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"]