Download configuration_bqalm.py from tturing/DM-GDNMLA-1.7B-MP: direct link, hf CLI and curl.
- Browser
- Download file 11.9 kB
-
https://huggingface.co/tturing/DM-GDNMLA-1.7B-MP/resolve/main/configuration_bqalm.py
- Command line
-
hf download hf://tturing/DM-GDNMLA-1.7B-MP/configuration_bqalm.py
-
curl -L -o configuration_bqalm.py https://huggingface.co/tturing/DM-GDNMLA-1.7B-MP/resolve/main/configuration_bqalm.py
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 | |
| 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 | |
| def rotary_dim(self) -> int: | |
| return int(self.head_dim * self.partial_rotary_factor) | |
| def n_rep(self) -> int: | |
| return self.n_heads // self.n_kv_heads | |
| def q_dim(self) -> int: | |
| return self.n_heads * self.head_dim | |
| 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, | |
| ) | |