Download tiny_gdn/config.py from kerzgrr/Haiku-base: direct link, hf CLI and curl.
- Browser
- Download file 8.14 kB
-
https://huggingface.co/kerzgrr/Haiku-base/resolve/main/tiny_gdn/config.py
- Command line
-
hf download hf://kerzgrr/Haiku-base/tiny_gdn/config.py
-
curl -L -o config.py https://huggingface.co/kerzgrr/Haiku-base/resolve/main/tiny_gdn/config.py
8.14 kB
| from __future__ import annotations | |
| import json | |
| from dataclasses import asdict, dataclass, fields | |
| from pathlib import Path | |
| from typing import Any | |
| class TinyGDNConfig: | |
| architecture: str = "TinyGDNForCausalLM" | |
| model_type: str = "tiny_gdn" | |
| vocab_size: int = 49_152 | |
| # Deep-thin sizing is deliberate: controlled sub-billion studies find | |
| # depth materially more valuable than width around the 125M-150M scale. | |
| hidden_size: int = 512 | |
| intermediate_size: int = 1_472 | |
| num_hidden_layers: int = 32 | |
| num_attention_heads: int = 4 | |
| num_key_value_heads: int = 1 | |
| attention_head_dim: int = 128 | |
| full_attention_interval: int = 4 | |
| attention_dropout: float = 0.0 | |
| partial_rotary_factor: float = 0.5 | |
| rope_theta: float = 1_000_000.0 | |
| linear_num_heads: int = 4 | |
| linear_num_value_heads: int = 4 | |
| linear_head_dim: int = 128 | |
| linear_expand_v: float = 1.0 | |
| linear_conv_kernel_dim: int = 4 | |
| allow_negative_eigenvalues: bool = False | |
| linear_layer_type: str = "gdn2" | |
| attention_layer_type: str = "full_attention" | |
| mlp_activation: str = "swiglu" | |
| use_block_attn_res: bool = False | |
| kda_safe_gate: bool = True | |
| kda_lower_bound: float = -5.0 | |
| mla_q_lora_rank: int | None = 512 | |
| mla_kv_lora_rank: int = 256 | |
| mla_qk_nope_head_dim: int = 128 | |
| mla_v_head_dim: int = 128 | |
| situ_gate_cap: float = 4.0 | |
| situ_up_cap: float = 25.0 | |
| max_position_embeddings: int = 32_768 | |
| training_sequence_length: int = 2_048 | |
| rms_norm_eps: float = 1e-6 | |
| initializer_range: float = 0.02 | |
| tie_word_embeddings: bool = True | |
| shared_layer_indices: tuple[int, ...] = () | |
| # MTP is an opt-in ablation at this scale; static MTP is not assumed to | |
| # improve a 150M model without a controlled pilot. | |
| mtp_num_heads: int = 0 | |
| mtp_adapter_rank: int = 128 | |
| mtp_loss_weight: float = 0.0 | |
| bos_token_id: int = 0 | |
| eos_token_id: int = 1 | |
| pad_token_id: int = 2 | |
| unk_token_id: int = 3 | |
| def __post_init__(self) -> None: | |
| if self.vocab_size <= 0 or self.vocab_size > 65_536: | |
| raise ValueError("vocab_size must fit the uint16 token dataset") | |
| if self.hidden_size != self.num_attention_heads * self.attention_head_dim: | |
| raise ValueError("hidden_size must equal num_attention_heads * attention_head_dim") | |
| if self.hidden_size != self.linear_num_heads * self.linear_head_dim: | |
| raise ValueError("hidden_size must equal linear_num_heads * linear_head_dim") | |
| if self.linear_num_value_heads < self.linear_num_heads: | |
| raise ValueError("linear_num_value_heads must be at least linear_num_heads") | |
| if self.linear_num_value_heads % self.linear_num_heads != 0: | |
| raise ValueError("linear_num_value_heads must be divisible by linear_num_heads") | |
| if self.num_attention_heads % self.num_key_value_heads != 0: | |
| raise ValueError("num_attention_heads must be divisible by num_key_value_heads") | |
| if self.num_hidden_layers % self.full_attention_interval != 0: | |
| raise ValueError("num_hidden_layers must be divisible by full_attention_interval") | |
| if not 0.0 < self.partial_rotary_factor <= 1.0: | |
| raise ValueError("partial_rotary_factor must be in (0, 1]") | |
| rotary_dim = int(self.attention_head_dim * self.partial_rotary_factor) | |
| if rotary_dim <= 0 or rotary_dim % 2: | |
| raise ValueError("The partial rotary dimension must be positive and even") | |
| if self.training_sequence_length > self.max_position_embeddings: | |
| raise ValueError("training_sequence_length exceeds max_position_embeddings") | |
| if len(set(self.shared_layer_indices)) != len(self.shared_layer_indices): | |
| raise ValueError("shared_layer_indices must be unique") | |
| if any( | |
| index < 0 or index >= self.num_hidden_layers | |
| for index in self.shared_layer_indices | |
| ): | |
| raise ValueError("shared_layer_indices contains an invalid layer") | |
| if self.mtp_num_heads < 0: | |
| raise ValueError("mtp_num_heads cannot be negative") | |
| if self.mtp_num_heads and self.mtp_adapter_rank <= 0: | |
| raise ValueError("mtp_adapter_rank must be positive when MTP is enabled") | |
| if not 0.0 <= self.mtp_loss_weight <= 1.0: | |
| raise ValueError("mtp_loss_weight must be between zero and one") | |
| for token_id in ( | |
| self.bos_token_id, | |
| self.eos_token_id, | |
| self.pad_token_id, | |
| self.unk_token_id, | |
| ): | |
| if not 0 <= token_id < self.vocab_size: | |
| raise ValueError(f"Special token ID {token_id} is outside the vocabulary") | |
| if self.linear_layer_type not in {"gdn2", "kda"}: | |
| raise ValueError("linear_layer_type must be 'gdn2' or 'kda'") | |
| if self.attention_layer_type not in {"full_attention", "gated_mla"}: | |
| raise ValueError("attention_layer_type must be 'full_attention' or 'gated_mla'") | |
| if self.mlp_activation not in {"swiglu", "situ_glu"}: | |
| raise ValueError("mlp_activation must be 'swiglu' or 'situ_glu'") | |
| if self.mla_kv_lora_rank <= 0: | |
| raise ValueError("mla_kv_lora_rank must be positive") | |
| if self.mla_q_lora_rank is not None and self.mla_q_lora_rank <= 0: | |
| raise ValueError("mla_q_lora_rank must be positive or None") | |
| if self.mla_qk_nope_head_dim <= 0 or self.mla_v_head_dim <= 0: | |
| raise ValueError("MLA head dimensions must be positive") | |
| if self.kda_lower_bound >= 0.0: | |
| raise ValueError("kda_lower_bound must be negative log-decay") | |
| if self.situ_gate_cap <= 0.0 or self.situ_up_cap <= 0.0: | |
| raise ValueError("SiTU caps must be positive") | |
| def layer_types(self) -> tuple[str, ...]: | |
| return tuple( | |
| self.attention_layer_type | |
| if (index + 1) % self.full_attention_interval == 0 | |
| else self.linear_layer_type | |
| for index in range(self.num_hidden_layers) | |
| ) | |
| def rotary_dim(self) -> int: | |
| return int(self.attention_head_dim * self.partial_rotary_factor) | |
| def effective_num_layers(self) -> int: | |
| return self.num_hidden_layers + len(self.shared_layer_indices) | |
| def to_dict(self) -> dict[str, Any]: | |
| payload = asdict(self) | |
| payload["layer_types"] = list(self.layer_types) | |
| return payload | |
| def save_json(self, path: Path) -> None: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| path.write_text( | |
| json.dumps(self.to_dict(), indent=2, sort_keys=True) + "\n", | |
| encoding="utf-8", | |
| ) | |
| def from_json(cls, path: Path) -> TinyGDNConfig: | |
| payload = json.loads(path.read_text(encoding="utf-8")) | |
| payload.pop("layer_types", None) | |
| if "shared_layer_indices" in payload: | |
| payload["shared_layer_indices"] = tuple(payload["shared_layer_indices"]) | |
| allowed = {item.name for item in fields(cls)} | |
| return cls(**{key: value for key, value in payload.items() if key in allowed}) | |
| def haiku_config(**overrides: Any) -> TinyGDNConfig: | |
| """650M Haiku: 3×KDA + 1×Gated-MLA, AttnRes, SiTU-GLU, MTP.""" | |
| payload = { | |
| "architecture": "HaikuForCausalLM", | |
| "model_type": "haiku", | |
| "vocab_size": 65_536, | |
| "hidden_size": 1024, | |
| "intermediate_size": 3840, | |
| "num_hidden_layers": 36, | |
| "num_attention_heads": 8, | |
| "num_key_value_heads": 8, | |
| "attention_head_dim": 128, | |
| "full_attention_interval": 4, | |
| "linear_num_heads": 8, | |
| "linear_num_value_heads": 8, | |
| "linear_head_dim": 128, | |
| "linear_layer_type": "kda", | |
| "attention_layer_type": "gated_mla", | |
| "mlp_activation": "situ_glu", | |
| "use_block_attn_res": True, | |
| "mtp_num_heads": 2, | |
| "mtp_adapter_rank": 128, | |
| "mtp_loss_weight": 0.2, | |
| "max_position_embeddings": 32_768, | |
| "training_sequence_length": 2048, | |
| } | |
| payload.update(overrides) | |
| return TinyGDNConfig(**payload) | |