DM-GDNMLA-1.7B-DM / configuration_bqalm.py
tturing's picture
Add Q35MLA_fp8base_md_ffn8_qkvo8_fp11_hd128_11785_fm128: GDN:MLA 1.7B, DM (match) arm, seed 11785 (standalone, trust_remote_code)
2ed6d1a verified
Raw History Blame Contribute Delete
11.9 kB
"""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,
)