MiniMind 学习笔记之 06 一张图、一百行:MiniMind 的骨架

上承〔05 · 从 768 维到下一个 token〕。前五篇把推理链路走完了:文字进来,过 8 层积木,出去变成下一个 token。这一篇做一个收拢:用一张架构图,把前五篇的所有零件各就各位;再用一百行伪代码,把 MiniMind 的实现意图压到最短。承上,也启下。

一、架构图:前五篇的零件各就各位

MiniMind 的整体数据流,从输入到输出,长这样:

ini 复制代码
输入文本
  │
  ▼
[BPE 分词器] ──→ token id 序列 [B, S]
  │
  ▼
[Embedding 层] (6400 × 768,与 lm_head 共享权重)
  │
  ▼
┌───────────────────────────────────────────────────────┐
│  MiniMindBlock × 8 层                                 │
│                                                       │
│  residual = hidden_states                            │
│  hidden_states = RMSNorm(hidden_states)              │
│  hidden_states = Attention(hidden_states)            │
│    ├─ QKV 投影 (GQA: 8 Q头 / 4 KV头)                 │
│    ├─ RoPE 旋转 Q, K                                 │
│    ├─ KV Cache 拼接                                  │
│    ├─ Flash Attention (is_causal=True)               │
│    └─ o_proj 输出                                    │
│  hidden_states = hidden_states + residual            │
│                                                       │
│  residual = hidden_states                            │
│  hidden_states = RMSNorm(hidden_states)              │
│  hidden_states = FeedForward(hidden_states)          │
│    ├─ gate_proj → SiLU                               │
│    ├─ up_proj                                        │
│    ├─ 逐元素相乘 (门控)                               │
│    └─ down_proj                                      │
│  hidden_states = hidden_states + residual            │
│                                                       │
└───────────────────────────────────────────────────────┘
  │
  ▼
[Final RMSNorm]
  │
  ▼
[lm_head] (768 → 6400,与 embed_tokens 共享权重)
  │
  ▼
[Softmax + 采样] ──→ 下一个 token

这张图里,前五篇的零件全部就位:

  • 02 篇:BPE 分词、Embedding 层。文字怎么变成向量。
  • 03 篇:Attention 里的 QKV 投影、因果掩码、多头。
  • 04 篇:RoPE、KV Cache、GQA、Flash Attention、FFN/SwiGLU、RMSNorm、残差连接。
  • 05 篇:Final RMSNorm、lm_head、Softmax、采样。

一张图,把五篇串完了。

二、一百行伪代码:骨架版的 MiniMind

真实代码 3500 行,里面大部分是工程细节:数据加载、训练循环、日志、检查点、分布式支持、模型变体。模型架构本身,一百行就够。

下面这份伪代码,去掉了所有工程噪音,数字全部硬编码。读起来清爽,记住几个关键数:768、96、6400、8、4、2432。

python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F

# 全局硬编码:768 维、96 头维、6400 词表、8 层、8 Q头、4 KV头、2432 FFN中间维
# RoPE 基频 1000000,最大位置 32768

class RMSNorm(nn.Module):
    def forward(self, x):
        # 除以均方根,不减均值
        return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + 1e-5)

def precompute_rope():
    # 48 个频率,从高到低;对应 96 维分 48 对
    freqs = 1.0 / (1000000.0 ** (torch.arange(0, 96, 2).float() / 96))
    t = torch.arange(32768)
    angles = torch.outer(t, freqs)                    # [32768, 48]
    return (torch.cat([torch.cos(angles)]*2, -1),     # [32768, 96]
            torch.cat([torch.sin(angles)]*2, -1))

def rope(x, cos, sin):
    # x: [B, H, S, 96],前后两半交换,前半取负
    x1, x2 = x[..., :48], x[..., 48:]
    return x * cos + torch.cat([-x2, x1], -1) * sin

class Attention(nn.Module):
    def __init__(self):
        self.q_proj = nn.Linear(768, 768, bias=False)   # 8 头 × 96 维
        self.k_proj = nn.Linear(768, 384, bias=False)   # 4 头 × 96 维
        self.v_proj = nn.Linear(768, 384, bias=False)
        self.o_proj = nn.Linear(768, 768, bias=False)
        self.cos, self.sin = precompute_rope()

    def forward(self, x, kv_cache=None):
        B, S, _ = x.shape
        # 投影 + 拆头:[B, 8, S, 96] / [B, 4, S, 96]
        q = self.q_proj(x).view(B, S, 8, 96).transpose(1, 2)
        k = self.k_proj(x).view(B, S, 4, 96).transpose(1, 2)
        v = self.v_proj(x).view(B, S, 4, 96).transpose(1, 2)

        # RoPE 只旋转 q 和 k
        q = rope(q, self.cos[:S], self.sin[:S])
        k = rope(k, self.cos[:S], self.sin[:S])

        # KV Cache:拼上缓存
        if kv_cache is not None:
            k = torch.cat([kv_cache[0], k], dim=2)
            v = torch.cat([kv_cache[1], v], dim=2)

        # GQA:4 个 KV 头复制成 8 个,和 Q 头对齐
        k = k.repeat_interleave(2, dim=1)
        v = v.repeat_interleave(2, dim=1)

        # Flash Attention,内部处理因果掩码
        out = F.scaled_dot_product_attention(q, k, v, is_causal=True)

        # 合并头,输出投影
        out = out.transpose(1, 2).reshape(B, S, 768)
        return self.o_proj(out), (k, v)

class FFN(nn.Module):
    def __init__(self):
        self.gate = nn.Linear(768, 2432, bias=False)
        self.up   = nn.Linear(768, 2432, bias=False)
        self.down = nn.Linear(2432, 768, bias=False)

    def forward(self, x):
        # gate 过 SiLU 当门控,up 提供信息,逐元素相乘
        return self.down(F.silu(self.gate(x)) * self.up(x))

class Block(nn.Module):
    def __init__(self):
        self.norm1 = RMSNorm()
        self.attn  = Attention()
        self.norm2 = RMSNorm()
        self.ffn   = FFN()

    def forward(self, x, kv_cache=None):
        # 注意力 + 残差
        h, new_kv = self.attn(self.norm1(x), kv_cache)
        x = x + h
        # FFN + 残差
        x = x + self.ffn(self.norm2(x))
        return x, new_kv

class MiniMind(nn.Module):
    def __init__(self):
        self.embed = nn.Embedding(6400, 768)
        self.layers = nn.ModuleList([Block() for _ in range(8)])
        self.norm = RMSNorm()
        self.lm_head = nn.Linear(768, 6400, bias=False)
        self.lm_head.weight = self.embed.weight   # 权重绑定

    def forward(self, ids, kv_caches=None):
        x = self.embed(ids)                        # [B, S, 768]
        new_kvs = []
        for i, layer in enumerate(self.layers):
            cache = kv_caches[i] if kv_caches else None
            x, new_kv = layer(x, cache)
            new_kvs.append(new_kv)
        return self.lm_head(self.norm(x)), new_kvs  # [B, S, 6400]

    @torch.no_grad()
    def generate(self, ids, max_new=100, temperature=0.7, top_k=50, top_p=0.9):
        kv_caches = None
        for _ in range(max_new):
            inp = ids if kv_caches is None else ids[:, -1:]
            logits, kv_caches = self.forward(inp, kv_caches)
            logits = logits[:, -1, :] / temperature
            # Top-k + Top-p 过滤
            v, _ = torch.topk(logits, top_k)
            logits[logits < v[:, [-1]]] = -float('inf')
            sorted_logits, sorted_idx = torch.sort(logits, descending=True)
            cum = torch.cumsum(F.softmax(sorted_logits, -1), -1)
            sorted_logits[cum > top_p] = -float('inf')
            probs = F.softmax(sorted_logits, -1)
            # 采样
            next_id = sorted_idx.gather(-1, torch.multinomial(probs, 1))
            ids = torch.cat([ids, next_id], -1)
            if next_id.item() == 2:   # 假设 2 是结束符
                break
        return ids

约 100 行。每一块在前五篇都拆过:

代码块 前五篇对应
RMSNorm 04 篇 6.3--6.4 节
precompute_rope / rope 04 篇 1.5--1.6 节
Attention 里的 q/k/v 投影 03 篇 5.1 节
repeat_interleave (GQA) 04 篇 3.3--3.5 节
scaled_dot_product_attention 04 篇 4.2 节
FFN 里的 gate/up/down 04 篇 5.5--5.7 节
Block 里的残差连接 04 篇 6.2 节
embed 和 lm_head 权重绑定 05 篇 1.3 节
generate 里的采样 05 篇三、四节

三、从 100 行到 3500 行:多出来的 3400 行是什么

一百行是骨架,跑不起来------它没有训练、没有数据、没有保存、没有日志。真实代码的 3500 行里,多出来的部分大概是这几类:

类别 占比 内容
训练循环 约 15% 前向、损失、反向、优化器、学习率调度、梯度累积、日志、检查点
数据工程 约 25% 数据加载、清洗、分词、格式化、批处理、padding、mask
模型变体 约 20% MoE 版本、VLM 版本、Omni 版本、不同尺寸的配置
工具代码 约 30% 模型转换、推理部署、Web 演示、评估脚本、tokenizer 训练
模型架构 约 10% 这一百行,加上各种辅助函数和边界处理

模型架构只占一成。 剩下九成,是把模型变成一个能跑、能训、能用的项目所需的工程。

这不是 MiniMind 独有的比例。几乎所有大模型项目都是这样:模型本身很简单,围绕它的工程很复杂。这也是为什么读完前五篇,再看真实代码,会发现"核心部分其实都见过了"。

另外,现代 LLM 基本不用 Dropout。MiniMind 的配置里也预留了这个参数,但默认值是 0。所以架构图和伪代码里都没画它。

四、承上,启下

这一篇起的作用是两个。

承上:把前五篇的所有零件,收拢到一张图和一百行代码里。如果前五篇是拆零件,这一篇是把零件装回去,让你看到整机长什么样。

启下 :这一百行里,全是 forward,没有 backward。模型怎么从随机初始化,被拧到能猜词,这一百行里没有答案。

下一个问题自然就出来了:那些旋钮,是怎么被拧动的?

这就要讲训练了。下一篇,我们从这一百行的 forward 出发,补上它的另一半:backward。

相关推荐
用户7341681035481 小时前
只盯被动响应做状态监测?论“控制信号”在变工况故障诊断中的重要作用
算法
Ivanqhz1 小时前
MQA、GQA、稀疏/滑动窗口注意力及混合范式
人工智能·算法·机器学习
罗湖老棍子1 小时前
荒岛野人(信息学奥赛一本通- P1637)(洛谷-P2421)
算法·数论·枚举·裴蜀定理·扩展欧几里得·扩展中国剩余定理
朝朝辞暮i1 小时前
C++ 第 40 章:ROS2 Publisher + Timer + Subscriber + Callback + Executor 完整闭环
开发语言·c++·算法·ros2
400分1 小时前
具身智能系列之从零到「机械臂会放试管」:π0.5 VLA 真机部署完全指南【第一节】
算法
小宋10211 小时前
一套服务托管多个 LoRA:适配器加载、租户隔离与热切换实战
大数据·人工智能·算法
Rosanci2 小时前
Codex 下载与本地部署实战:从安装到运行全流程指南
开发语言·前端·算法·chatgpt·codex
我爱工作&工作love我2 小时前
P14358 [CSP-J 2025] 座位
算法
旖旎夜光2 小时前
LeetCode 1576: 替换所有的问号(模拟) —— 题解
c++·学习·算法·leetcode·力控