attention-0.6b-adamw-wsd-step6500 / configuration_parallax.py
YifeiZuo's picture
Upload AdamW-WSD attention 0.6B step 6500
fe70ec2 verified
Raw History Blame Contribute Delete
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"]