File size: 4,980 Bytes
5e6d9f5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
"""模型配置。

一套 dataclass 描述整个模型,三档预设:nano / small / base。
命名尽量与 DeepSeek-V2/V3 技术报告保持一致,方便对照论文阅读。
"""

from dataclasses import dataclass, field, asdict
from typing import Literal, Optional


@dataclass
class ModelConfig:
    # ---------------- 基础规模 ----------------
    vocab_size: int = 8192
    dim: int = 256                      # 模型隐藏维度 d_model
    n_layers: int = 6
    max_seq_len: int = 512
    norm_eps: float = 1e-5
    dropout: float = 0.0
    tie_embeddings: bool = True         # 输出头与词嵌入共享权重

    # ---------------- MLA 多头潜在注意力 ----------------
    n_heads: int = 8
    q_lora_rank: int = 0                # 0 表示 Q 不做低秩分解(小模型没必要)
    kv_lora_rank: int = 64              # KV 压缩到的潜在维度 —— KV cache 只存这个
    qk_nope_head_dim: int = 32          # Q/K 中不加旋转位置编码的部分
    qk_rope_head_dim: int = 16          # Q/K 中加 RoPE 的部分(解耦 RoPE)
    v_head_dim: int = 32
    attn_impl: Literal["naive", "absorb"] = "naive"  # 训练用 naive(走 flash),解码可用 absorb

    # ---------------- RoPE ----------------
    rope_theta: float = 10000.0
    rope_scaling: float = 1.0           # >1 时做 NTK-aware 频率插值,用于外推长上下文

    # ---------------- DeepSeekMoE ----------------
    n_dense_layers: int = 1             # 前若干层用普通 FFN(V3 的做法,稳定训练早期)
    dense_inter_dim: int = 704          # 稠密层 FFN 中间维度(SwiGLU)
    moe_inter_dim: int = 256            # 单个细粒度专家的中间维度
    n_routed_experts: int = 8           # 路由专家数
    n_shared_experts: int = 1           # 共享专家数(每个 token 都过)
    n_activated_experts: int = 2        # 每 token 激活的路由专家数 top-k
    n_expert_groups: int = 1            # 专家分组数(V3 的 group-limited routing)
    n_limited_groups: int = 1           # 每 token 最多命中几个组
    score_func: Literal["softmax", "sigmoid"] = "sigmoid"   # V3 用 sigmoid
    route_scale: float = 1.0            # 路由权重缩放
    bias_update_speed: float = 1e-3     # 无辅助损失负载均衡:专家偏置更新步长 γ
    aux_loss_alpha: float = 1e-3        # 序列级辅助损失权重(V3 里权重很小,兜底用)

    # ---------------- MTP 多 token 预测 ----------------
    n_mtp: int = 1                      # MTP 深度,0 表示关闭
    mtp_loss_weight: float = 0.3        # MTP 损失权重 λ

    # -------------------------------------------------
    @property
    def qk_head_dim(self) -> int:
        return self.qk_nope_head_dim + self.qk_rope_head_dim

    def to_dict(self) -> dict:
        return asdict(self)

    @classmethod
    def from_dict(cls, d: dict) -> "ModelConfig":
        known = {f for f in cls.__dataclass_fields__}
        return cls(**{k: v for k, v in d.items() if k in known})


# ============================ 预设档位 ============================

def nano() -> ModelConfig:
    """~12M 总参 / ~5M 激活。笔记本 CPU 都能跑,用来验证架构是否正确。"""
    return ModelConfig(
        vocab_size=8192, dim=256, n_layers=6, n_heads=8, max_seq_len=512,
        q_lora_rank=0, kv_lora_rank=64,
        qk_nope_head_dim=32, qk_rope_head_dim=16, v_head_dim=32,
        n_dense_layers=1, dense_inter_dim=704,
        moe_inter_dim=256, n_routed_experts=8, n_shared_experts=1,
        n_activated_experts=2, n_expert_groups=1, n_limited_groups=1,
        n_mtp=1,
    )


def small() -> ModelConfig:
    """~110M 总参 / ~35M 激活。单卡或 M 系列 Mac 训几小时能出像样的效果。"""
    return ModelConfig(
        vocab_size=16384, dim=512, n_layers=12, n_heads=8, max_seq_len=1024,
        q_lora_rank=0, kv_lora_rank=128,
        qk_nope_head_dim=48, qk_rope_head_dim=16, v_head_dim=64,
        n_dense_layers=1, dense_inter_dim=1408,
        moe_inter_dim=640, n_routed_experts=8, n_shared_experts=1,
        n_activated_experts=2, n_expert_groups=1, n_limited_groups=1,
        n_mtp=1,
    )


def base() -> ModelConfig:
    """~370M 总参 / ~120M 激活。开始有 group-limited routing 和 Q 低秩分解。"""
    return ModelConfig(
        vocab_size=32000, dim=768, n_layers=16, n_heads=12, max_seq_len=2048,
        q_lora_rank=384, kv_lora_rank=192,
        qk_nope_head_dim=64, qk_rope_head_dim=32, v_head_dim=64,
        n_dense_layers=2, dense_inter_dim=2048,
        moe_inter_dim=1024, n_routed_experts=8, n_shared_experts=1,
        n_activated_experts=2, n_expert_groups=4, n_limited_groups=2,
        n_mtp=1,
    )


CONFIGS = {"nano": nano, "small": small, "base": base}


def get_config(name: str) -> ModelConfig:
    if name not in CONFIGS:
        raise KeyError(f"未知配置 {name!r},可选:{list(CONFIGS)}")
    return CONFIGS[name]()