File size: 3,724 Bytes
fe70ec2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
"""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"]