from __future__ import annotations from dataclasses import asdict, dataclass MODEL_NAME = "Sol Milkshake" ORGANIZATION = "Sol Labs" DEPLOYED_PARAMS = 2_990_000 VOCAB_SIZE = 2_048 D_MODEL = 192 MAX_CONTEXT = 2_048 N_PHYSICAL_BLOCKS = 5 N_Q_HEADS = 6 N_KV_HEADS = 2 HEAD_DIM = 32 FFN_HIDDEN = 512 TN_RANK = 26 TIED_EMBEDDING_HEAD = True XSA_PARAMETER_COUNT = 0 TOKENS_PER_PARAMETER = 544 GLOBAL_TOKENS = 1_626_560_000 GLOBAL_BATCH = 32_768 COMPLETE_UPDATES = 49_638 FINAL_PARTIAL_TOKENS = 22_016 assert DEPLOYED_PARAMS * TOKENS_PER_PARAMETER == GLOBAL_TOKENS assert COMPLETE_UPDATES * GLOBAL_BATCH + FINAL_PARTIAL_TOKENS == GLOBAL_TOKENS @dataclass(frozen=True) class SolConfig: model_name: str = MODEL_NAME vocab_size: int = VOCAB_SIZE d_model: int = D_MODEL max_context: int = MAX_CONTEXT n_blocks: int = N_PHYSICAL_BLOCKS n_q_heads: int = N_Q_HEADS n_kv_heads: int = N_KV_HEADS head_dim: int = HEAD_DIM ffn_hidden: int = FFN_HIDDEN tn_rank: int = TN_RANK rope_theta: float = 20_000.0 memory_width: int = 64 memory_slots: int = 32 chunk_size: int = 32 passes: int = 3 dropout: float = 0.0 def as_dict(self) -> dict: return asdict(self)