File size: 11,894 Bytes
2ed6d1a | 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 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 | """HF `PretrainedConfig` for BqaLM, the hybrid decoder of the DeltaMatching study.
Self-contained copy of the bqa codebase's `src/pretrain/hf/configuration_bqalm.py`, with the backbone's
`ModelConfig` (`src/pretrain/modeling/config.py`) vendored below so the checkpoint loads through
`trust_remote_code` without the bqa source tree. Field names, defaults and `to_model_config` are unchanged, so a
`config.json` written by the bqa exporter round-trips exactly.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from transformers import PretrainedConfig
@dataclass
class ModelConfig:
"""Backbone architecture (the bqa `ModelConfig`). Only the fields the shipped mixers read matter here; the rest
are carried so the exporter's configs keep their exact meaning."""
# core dims
vocab_size: int = 32768
d_model: int = 1024
n_layers: int = 24
n_heads: int = 16
n_kv_heads: int = 4 # GQA: n_heads % n_kv_heads == 0; == n_heads is MHA
head_dim: int | None = None # default d_model // n_heads
# FFN (SwiGLU)
intermediate_size: int | None = None
ffn_mult: float = 8.0 / 3.0
ffn_multiple_of: int = 256
# positional / norm
rope_theta: float = 10000.0
rope_scaling: dict | None = None # {"type": "yarn", ...} for YaRN context extension, else None
partial_rotary_factor: float = 1.0 # rotate the leading head_dim * factor channels only
nope: bool = False # identity rotation (no positional encoding in attention)
norm_eps: float = 1e-5
max_seq_len: int = 2048
# mixers
mixer: str = "gqa"
attn_impl: str = "auto"
qk_norm: bool = False # per-head RMSNorm on q, k before RoPE
attn_output_gate: bool = False # q_proj emits 2 * q_dim; out * sigmoid(gate) before o_proj
layer_mixers: list[str] | None = None # per-layer pattern, repeated over n_layers
gdn_head_dim: int | None = None
gdn_num_heads: int | None = None
gdn_num_v_heads: int | None = None
kv_lora_rank: int | None = None
mamba2_headdim: int | None = None
mamba2_d_state: int = 128
mamba2_expand: int = 2
mamba2_ngroups: int = 1
mamba2_chunk_size: int = 256
kda_head_dim: int | None = None
kda_num_heads: int | None = None
kda_num_v_heads: int | None = None
kda_expand_v: float = 1.0
kda_conv_size: int = 4
kda_allow_neg_eigval: bool = False
kda_safe_gate: bool = False
kda_lower_bound: float | None = None
# numerics
tie_embeddings: bool = True
attn_dropout: float = 0.0
resid_dropout: float = 0.0
initializer_range: float = 0.02
z_loss_weight: float = 1e-4
rms_norm_in_fp32: bool = True
fused_rmsnorm: bool = False
fused_rope: bool = False
fused_swiglu: bool = False
# training-kernel and BQA-mixer knobs, carried for config fidelity only (not read by the shipped mixers)
sb_mode: str = "uniform"
sb_window: int = 512
sb_sink: bool = True
sb_impl: str = "triton"
sb_t0_dedup: bool = False
sb_fused_producers: bool = False
sb_fuse_tier2: str = "off"
hs_protect: str = ""
hs_impl: str = "compose"
bqa_window: int = 512
bqa_rotate: bool = True
bqa_nvfp4_block: int = 16
bqa_remote_topk: int = -1
bqa_local_precision: str = "fp8"
bqa_remote_k_precision: str = "fp8"
bqa_remote_v_precision: str = "nvfp4"
_resolved: bool = field(default=False, repr=False)
def __post_init__(self):
if self.head_dim is None:
assert self.d_model % self.n_heads == 0, "d_model must divide by n_heads when head_dim is None"
self.head_dim = self.d_model // self.n_heads
assert self.n_heads % self.n_kv_heads == 0, (
f"n_heads ({self.n_heads}) must be divisible by n_kv_heads ({self.n_kv_heads})")
if self.intermediate_size is None:
raw = self.ffn_mult * self.d_model
m = self.ffn_multiple_of
self.intermediate_size = int(((int(raw) + m - 1) // m) * m)
assert 0.0 < self.partial_rotary_factor <= 1.0, (
f"partial_rotary_factor must be in (0, 1], got {self.partial_rotary_factor}")
assert self.rotary_dim % 2 == 0 and self.rotary_dim > 0, (
f"rotary_dim = head_dim({self.head_dim}) * partial_rotary_factor"
f"({self.partial_rotary_factor}) = {self.rotary_dim}, which must be a positive even number")
if self.rotary_dim != self.head_dim and self.fused_rope:
self.fused_rope = False
self._resolved = True
@property
def rotary_dim(self) -> int:
return int(self.head_dim * self.partial_rotary_factor)
@property
def n_rep(self) -> int:
return self.n_heads // self.n_kv_heads
@property
def q_dim(self) -> int:
return self.n_heads * self.head_dim
@property
def kv_dim(self) -> int:
return self.n_kv_heads * self.head_dim
def mixer_for_layer(self, layer_idx: int) -> str:
if self.layer_mixers is not None:
return self.layer_mixers[layer_idx % len(self.layer_mixers)]
return self.mixer
class BqaLMConfig(PretrainedConfig):
model_type = "bqalm"
def __init__(
self,
vocab_size: int = 50257,
d_model: int = 1024,
n_layers: int = 24,
n_heads: int = 16,
n_kv_heads: int = 4,
head_dim: int | None = None,
intermediate_size: int | None = None,
ffn_mult: float = 8.0 / 3.0,
ffn_multiple_of: int = 256,
rope_theta: float = 10000.0,
rope_scaling: dict | None = None,
norm_eps: float = 1e-5,
rms_norm_in_fp32: bool = True,
qk_norm: bool = False,
attn_output_gate: bool = False,
partial_rotary_factor: float = 1.0,
nope: bool = False,
gdn_head_dim: int | None = None,
gdn_num_heads: int | None = None,
gdn_num_v_heads: int | None = None,
kv_lora_rank: int | None = None,
mamba2_headdim: int | None = None,
mamba2_d_state: int = 128,
mamba2_expand: int = 2,
mamba2_ngroups: int = 1,
mamba2_chunk_size: int = 256,
kda_head_dim: int | None = None,
kda_num_heads: int | None = None,
kda_num_v_heads: int | None = None,
kda_expand_v: float = 1.0,
kda_conv_size: int = 4,
kda_allow_neg_eigval: bool = False,
kda_safe_gate: bool = False,
kda_lower_bound: float | None = None,
max_seq_len: int = 2048,
mixer: str = "gqa",
attn_impl: str = "auto",
layer_mixers: list[str] | None = None,
tie_embeddings: bool = True,
z_loss_weight: float = 0.0,
bqa_window: int = 512,
bqa_remote_topk: int = 512,
bqa_local_precision: str = "fp8",
bqa_remote_k_precision: str = "fp8",
bqa_remote_v_precision: str = "nvfp4",
bqa_rotate: bool = True,
bqa_nvfp4_block: int = 16,
sb_impl: str = "triton",
**kwargs,
):
self.vocab_size = vocab_size
self.d_model = d_model
self.n_layers = n_layers
self.n_heads = n_heads
self.n_kv_heads = n_kv_heads
self.head_dim = head_dim
self.intermediate_size = intermediate_size
self.ffn_mult = ffn_mult
self.ffn_multiple_of = ffn_multiple_of
self.rope_theta = rope_theta
self.rope_scaling = rope_scaling
self.norm_eps = norm_eps
self.rms_norm_in_fp32 = rms_norm_in_fp32
self.qk_norm = qk_norm
self.attn_output_gate = attn_output_gate
self.partial_rotary_factor = partial_rotary_factor
self.nope = bool(nope)
self.gdn_head_dim = gdn_head_dim
self.gdn_num_heads = gdn_num_heads
self.gdn_num_v_heads = gdn_num_v_heads
self.kv_lora_rank = kv_lora_rank
self.mamba2_headdim = mamba2_headdim
self.mamba2_d_state = mamba2_d_state
self.mamba2_expand = mamba2_expand
self.mamba2_ngroups = mamba2_ngroups
self.mamba2_chunk_size = mamba2_chunk_size
self.kda_head_dim = kda_head_dim
self.kda_num_heads = kda_num_heads
self.kda_num_v_heads = kda_num_v_heads
self.kda_expand_v = kda_expand_v
self.kda_conv_size = kda_conv_size
self.kda_allow_neg_eigval = kda_allow_neg_eigval
self.kda_safe_gate = kda_safe_gate
self.kda_lower_bound = kda_lower_bound
self.max_seq_len = max_seq_len
# set before super().__init__: transformers' rope validation reads max_position_embeddings during init
self.max_position_embeddings = max_seq_len
self.mixer = mixer
self.attn_impl = attn_impl
self.layer_mixers = layer_mixers
self.tie_embeddings = tie_embeddings
self.z_loss_weight = z_loss_weight
self.bqa_window = bqa_window
self.bqa_remote_topk = bqa_remote_topk
self.bqa_local_precision = bqa_local_precision
self.bqa_remote_k_precision = bqa_remote_k_precision
self.bqa_remote_v_precision = bqa_remote_v_precision
self.bqa_rotate = bqa_rotate
self.bqa_nvfp4_block = bqa_nvfp4_block
self.sb_impl = sb_impl
kwargs.setdefault("max_position_embeddings", max_seq_len)
kwargs.setdefault("hidden_size", d_model)
kwargs.setdefault("num_hidden_layers", n_layers)
kwargs.setdefault("num_attention_heads", n_heads)
kwargs.setdefault("tie_word_embeddings", tie_embeddings)
super().__init__(**kwargs)
def to_model_config(self) -> ModelConfig:
return ModelConfig(
vocab_size=self.vocab_size, d_model=self.d_model, n_layers=self.n_layers,
n_heads=self.n_heads, n_kv_heads=self.n_kv_heads, head_dim=self.head_dim,
intermediate_size=self.intermediate_size, ffn_mult=self.ffn_mult,
ffn_multiple_of=self.ffn_multiple_of, rope_theta=self.rope_theta,
rope_scaling=self.rope_scaling,
norm_eps=self.norm_eps, rms_norm_in_fp32=self.rms_norm_in_fp32, qk_norm=self.qk_norm,
attn_output_gate=self.attn_output_gate,
partial_rotary_factor=self.partial_rotary_factor,
nope=bool(getattr(self, "nope", False)),
gdn_head_dim=self.gdn_head_dim, gdn_num_heads=self.gdn_num_heads,
gdn_num_v_heads=self.gdn_num_v_heads,
kv_lora_rank=getattr(self, "kv_lora_rank", None),
mamba2_headdim=getattr(self, "mamba2_headdim", None), mamba2_d_state=getattr(self, "mamba2_d_state", 128),
mamba2_expand=getattr(self, "mamba2_expand", 2), mamba2_ngroups=getattr(self, "mamba2_ngroups", 1),
mamba2_chunk_size=getattr(self, "mamba2_chunk_size", 256),
kda_head_dim=getattr(self, "kda_head_dim", None), kda_num_heads=getattr(self, "kda_num_heads", None),
kda_num_v_heads=getattr(self, "kda_num_v_heads", None), kda_expand_v=getattr(self, "kda_expand_v", 1.0),
kda_conv_size=getattr(self, "kda_conv_size", 4), kda_allow_neg_eigval=getattr(self, "kda_allow_neg_eigval", False),
kda_safe_gate=getattr(self, "kda_safe_gate", False), kda_lower_bound=getattr(self, "kda_lower_bound", None),
max_seq_len=self.max_seq_len, mixer=self.mixer,
attn_impl=self.attn_impl, layer_mixers=self.layer_mixers,
tie_embeddings=self.tie_embeddings, z_loss_weight=self.z_loss_weight,
bqa_window=self.bqa_window, bqa_remote_topk=self.bqa_remote_topk,
bqa_local_precision=self.bqa_local_precision,
bqa_remote_k_precision=self.bqa_remote_k_precision,
bqa_remote_v_precision=self.bqa_remote_v_precision,
bqa_rotate=self.bqa_rotate, bqa_nvfp4_block=self.bqa_nvfp4_block,
sb_impl=self.sb_impl,
)
|