Spaces:
Runtime error
Runtime error
| """ECHO-1-nano configuration. | |
| A from-scratch decoder-only transformer that collects the modern architecture | |
| tricks (RMSNorm, RoPE, GQA/MLA, SwiGLU, MoE, MTP, QK-norm, KV cache). Every trick | |
| is a toggle so each can be ablated and understood in isolation. This is a didactic | |
| instrument: tiny by design, the point is the architecture, not capability. | |
| """ | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| class NanoConfig: | |
| # core | |
| vocab_size: int = 256 | |
| d_model: int = 256 | |
| n_layers: int = 6 | |
| n_heads: int = 8 | |
| n_kv_heads: int = 2 # GQA: < n_heads shares KV; == n_heads is plain MHA | |
| d_ff: int = 0 # 0 -> auto SwiGLU width (~8/3 * d_model, /64 rounded) | |
| max_seq: int = 512 | |
| dropout: float = 0.0 | |
| tie_embeddings: bool = True | |
| # positional / stability | |
| rope_theta: float = 10000.0 | |
| qk_norm: bool = True # RMSNorm on per-head Q and K before RoPE | |
| # attention variant | |
| attn: str = "gqa" # "gqa" | "mla" | |
| # MLA (DeepSeek-V2) latent compression; used when attn == "mla" | |
| kv_lora_rank: int = 64 # compressed KV latent dim | |
| q_lora_rank: int = 0 # 0 -> no Q compression | |
| rope_head_dim: int = 0 # 0 -> head_dim // 2 decoupled RoPE dim | |
| # feed-forward variant | |
| ffn: str = "dense" # "dense" | "moe" | |
| # MoE (DeepSeek-V3) used when ffn == "moe" | |
| n_experts: int = 8 | |
| n_shared: int = 1 # always-on shared experts | |
| top_k: int = 2 # routed experts per token | |
| moe_aux_alpha: float = 0.0 # 0 -> aux-loss-free (bias-based) balancing | |
| # multi-token prediction (DeepSeek-V3) | |
| mtp_tokens: int = 0 # 0 -> off; predict this many future tokens | |
| mtp_lambda: float = 0.3 # weight of the averaged MTP loss | |
| def head_dim(self) -> int: | |
| assert self.d_model % self.n_heads == 0, "d_model must divide n_heads" | |
| return self.d_model // self.n_heads | |
| def ff_dim(self) -> int: | |
| if self.d_ff: | |
| return self.d_ff | |
| # SwiGLU: keep param count near a 4x dense FFN -> ~8/3 hidden, rounded. | |
| h = int(8 / 3 * self.d_model) | |
| return (h + 63) // 64 * 64 | |