File size: 3,822 Bytes
c6b1b88
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e14a9b6
 
 
 
c6b1b88
 
 
 
 
 
 
 
 
e14a9b6
c6b1b88
 
 
 
 
 
 
e14a9b6
c6b1b88
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# models/lm · LM 与 LMForCausalLM

## 类图

```
LM(nn.Module)                      # 纯主干
  ├─ embed_tokens
  ├─ layers: ModuleList[Block]
  ├─ norm: RMSNorm
  └─ freqs_cos / freqs_sin (buffer)

LMForCausalLM(PreTrainedModel, GenerationMixin)   # 可训练 / 可生成
  ├─ model: LM
  └─ lm_head: Linear(hidden, vocab)
```

## LM(主干)

```python
class LM(nn.Module):
    def __init__(self, config):
        self.embed_tokens = nn.Embedding(vocab, hidden)
        self.layers = nn.ModuleList([Block(l, config) for l in range(num_layers)])
        self.norm = RMSNorm(hidden)
        freqs_cos, freqs_sin = precompute_freqs_cis(...)
        self.register_buffer("freqs_cos", freqs_cos, persistent=False)
        self.register_buffer("freqs_sin", freqs_sin, persistent=False)

    def forward(self, input_ids, ...):
        h = self.dropout(self.embed_tokens(input_ids))
        if self.freqs_cos[0, 0] != 1.0:            # 未初始化时重算
            freqs_cos, freqs_sin = precompute_freqs_cis(...)
            self.freqs_cos, self.freqs_sin = freqs_cos.to(device=h.device, dtype=h.dtype), ...
        position_embeddings = (self.freqs_cos[start:], self.freqs_sin[start:])   # 按 past 长度切片
        presents = []
        for layer, past in zip(self.layers, past_key_values):
            h, present = layer(h, position_embeddings, past_key_value=past, ...)
            presents.append(present)
        h = self.norm(h)
        aux_loss = sum(l.mlp.aux_loss for l in self.layers if MoE)
        return h, presents, aux_loss
```

- `freqs_cos/sin``persistent=False` 注册,不进 state_dict(节省 3 MB 冗余);HF `from_pretrained` 创建模型时在 meta device,非持久 buffer 之后用 `torch.empty_like` 分配(未初始化内存)。forward 内 guard `if self.freqs_cos[0, 0] != 1.0` 捕获异常值并自动重算(cos(0)=1,未初始化的 0/NaN/随机值均会触发)。
- `aux_loss` 仅在 MoE 时非 0,供 `LMForCausalLM` 加到总损失。

## LMForCausalLM(带 head + 损失 + 生成)

```python
class LMForCausalLM(PreTrainedModel, GenerationMixin):
    config_class = LMConfig
    model_type = "omni"
    _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"}

    def __init__(self, config):
        self.model = LM(config)
        self.lm_head = nn.Linear(hidden, vocab, bias=False)
        if config.tie_word_embeddings:
            self.model.embed_tokens.weight = self.lm_head.weight   # 权重绑定

    def forward(self, input_ids, labels=None, logits_to_keep=0, ...):
        h, past, aux_loss = self.model(...)
        logits = self.lm_head(h[:, slice(-logits_to_keep, None)])
        loss = CE(logits[:, :-1], labels[:, 1:], ignore_index=-100) if labels else None
        return MoeCausalLMOutputWithPast(loss, aux_loss, logits, ...)

    @torch.inference_mode()
    def generate(self, input_ids, max_new_tokens=8192, temperature=0.85,
                 top_p=0.85, top_k=50, repetition_penalty=1.0, ...):
        # 自回归循环:每次只喂新 token,拼接 past_key_values
        # 支持 top_k / top_p 截断、repetition_penalty、eos 提前停止、streamer
```

## 要点(面试)

1. **权重绑定 (weight tying)**`lm_head``embed_tokens` 共享权重,省一半词表参数;`_tied_weights_keys``save_pretrained` 只存一份。
2. **`logits_to_keep`**:生成时只算最后若干 token 的 logits,省算力。
3. **损失平移**:用 `logits[:, :-1]``labels[:, 1:]` 对齐,标准next-token预测。
4. **`generate` 自回归**:增量解码 + KV-cache,每次只 forward 新 token;`repetition_penalty` 对历史 token 降权防重复。
5. 继承 `PreTrainedModel` → 可直接用 `from_pretrained` / `save_pretrained` / `generate`(HF 生态)。