欢迎光临
我们一直在努力

nanochat 核心结构

一、项目定位

– **目标**:"$100 能买到的最好的 ChatGPT"
– **代码量**:约 6000+ 行(模块化设计)
– **状态**:活跃维护(2026)
– ** motto**:"The best ChatGPT that $100 can buy"

 二、文件结构

```
nanochat/
├── nanochat/               # 核心模块
│   ├── __init__.py
│   ├── gpt.py              # GPT Transformer 模型 (~555行)
│   ├── tokenizer.py        # BPE 分词器 (~279行)
│   ├── engine.py           # 推理引擎 + KV Cache (~352行)
│   ├── optim.py            # AdamW + Muon 优化器 (~459行)
│   ├── fp8.py              # FP8 精度支持 (~262行)
│   ├── flash_attention.py  # 自定义 Flash Attention
│   ├── dataloader.py       # 分布式数据加载
│   ├── dataset.py          # 数据集下载/读取工具
│   ├── loss_eval.py        # Bits per byte 评估
│   ├── core_eval.py        # DCLM CORE 评分
│   ├── checkpoint_manager.py  # 检查点管理
│   └── common.py           # 通用工具函数
├── scripts/                # 训练脚本
│   ├── base_train.py       # 预训练 (~604行)
│   ├── base_eval.py        # 基础模型评估
│   ├── chat_sft.py         # SFT 微调 (~499行)
│   ├── chat_rl.py          # RL 对齐 (~326行)
│   ├── chat_cli.py         # CLI 对话接口
│   ├── infer_bench.py      # 推理基准测试
│   ├── tok_train.py        # 分词器训练
│   └── tok_eval.py         # 分词器评估
├── runs/                   # 一键运行脚本
│   ├── speedrun.sh         # 快速复现 GPT-2
│   ├── scaling_laws.sh     # 扩展律实验
│   ├── miniseries.sh       # 小型训练
│   └── runcpu.sh           # CPU/MPS 运行示例
├── tasks/                  # 评估任务
│   ├── arc.py              # 多选科学问题
│   ├── gsm8k.py            # 小学数学
│   ├── mmlu.py             # 广泛主题多选
│   ├── humaneval.py        # Python 编程
│   └── smoltable.py        # 表格理解
├── tests/                  # 测试
│   ├── test_attentionFallback.py
│   ├── test_engine.py
│   ├── test_execution.py
│   ├── test_optim.py
│   ├── test_tasks.py
│   └── test_tokenizer.py
├── dev/                    # 开发文档
│   ├── LEADERBOARD.md
│   └── nanochat.png
├── pyproject.toml
└── uv.lock
```

三、核心模块详解

 3.1 GPT 模型 (gpt.py)

```python
class GPT(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.config = config
        
        # 嵌入层
        self.wte = nn.Embedding(config.vocab_size, config.n_embd)
        self.wpe = nn.Embedding(config.block_size, config.n_embd)
        
        # Transformer 层
        self.h = nn.ModuleList([Block(config) for _ in range(config.n_layer)])
        
        # 最终归一化 + 输出头
        self.ln_f = LayerNorm(config.n_embd)
        self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
        
        # 权重共享
        self.wte.weight = self.lm_head.weight

    def forward(self, idx, targets=None, return_loss=True):
        B, T = idx.size()
        
        # 位置嵌入
        pos = torch.arange(T, device=idx.device)
        pos_emb = self.wpe(pos)
        tok_emb = self.wte(idx)
        x = tok_emb + pos_emb
        
        # Transformer 层
        for block in self.h:
            x = block(x)
        x = self.ln_f(x)
        
        # 预测
        logits = self.lm_head(x)
        
        if targets is not None:
            loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1))
            return logits, loss
        else:
            return logits
```

 3.2 Block 模块

```python
class Block(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.ln_1 = LayerNorm(config.n_embd)
        self.attn = CausalSelfAttention(config)
        self.ln_2 = LayerNorm(config.n_embd)
        self.mlp = MLP(config)

    def forward(self, x):
        # 预归一化 + 残差连接
        x = x + self.attn(self.ln_1(x))
        x = x + self.mlp(self.ln_2(x))
        return x
```

 3.3 注意力机制 (CausalSelfAttention)

```python
class CausalSelfAttention(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.n_head = config.n_head
        self.n_embd = config.n_embd
        self.head_dim = n_embd // n_head
        
        # QKV 投影
        self.c_attn = nn.Linear(n_embd, 3 * n_embd, bias=False)
        # 输出投影
        self.c_proj = nn.Linear(n_embd, n_embd, bias=False)
        
        # Flash Attention 支持
        self.use_flash = hasattr(F, 'scaled_dot_product_attention')

    def forward(self, x):
        B, T, C = x.size()
        
        # 1. 投影 QKV
        qkv = self.c_attn(x)
        q, k, v = qkv.split(self.n_embd, dim=2)
        
        # 2. 多头重塑
        k = k.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
        q = q.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
        v = v.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
        
        # 3. 高效注意力
        if self.use_flash:
            y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
        else:
            # 手动实现
            scores = q @ k.transpose(-2, -1) * (1.0 / math.sqrt(self.head_dim))
            scores = scores.masked_fill(self.causal_mask[:, :, :T, :T] == 0, float('-inf'))
            scores = F.softmax(scores, dim=-1)
            y = scores @ v
        
        # 4. 合并头 + 输出投影
        y = y.transpose(1, 2).contiguous().view(B, T, C)
        y = self.c_proj(y)
        return y
```

 3.4 MLP (前馈网络)

```python
class MLP(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd, bias=False)
        self.gelu = nn.GELU()
        self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd, bias=False)

    def forward(self, x):
        x = self.c_fc(x)
        x = self.gelu(x)
        x = self.c_proj(x)
        return x
```

 3.5 推理引擎 (engine.py) — 新增核心!

```python
class Engine:
    """高效推理引擎,带 KV Cache"""
    
    def __init__(self, model):
        self.model = model
        self.kv_cache = {}
        
    def forward_with_cache(self, idx, position_offset):
        """使用 KV Cache 加速自回归生成"""
        # 只计算当前 token 的 logits
        logits = self.model(idx[:, -1:])
        
        # 更新 KV Cache
        # … (详细实现见 engine.py)
        
        return logits
```

**KV Cache 原理**:
– 预计算并缓存之前的 Key/Value 矩阵
– 新 token 只需计算当前位置的 attention
– 避免重复计算,加速 N 倍

3.6 分词器 (tokenizer.py)

```python
class Tokenizer:
    """BPE 分词器,风格类似 GPT-4"""
    
    def __init__(self, vocab_path):
        self.encoder = load_bpe_vocab(vocab_path)
        self.decoder = {v: k for k, v in self.encoder.items()}
        
    def encode(self, text):
        """文本 → token IDs"""
        tokens = byte_pair_encode(text, self.encoder)
        return tokens
    
    def decode(self, tokens):
        """token IDs → 文本"""
        return ''.join([self.decoder[t] for t in tokens])
```

 3.7 优化器 (optim.py)

```python
class MuonAdamW(Optimizer):
    """AdamW + Muon 优化器,支持 1 GPU 和分布式"""
    
    def __init__(self, params, lr=6e-4, betas=(0.9, 0.95), weight_decay=1e-1):
        super().__init__(params, lr=lr)
        self.betas = betas
        self.weight_decay = weight_decay
        
    def step(self):
        for group in self.param_groups:
            for p in group['params']:
                state = self.state[p]
                
                # 一阶矩估计
                state['exp_avg'] = 0.9 * state.get('exp_avg', 0) + 0.1 * p.grad
                # 二阶矩估计
                state['exp_avg_sq'] = 0.95 * state.get('exp_avg_sq', 0) + 0.05 * p.grad ** 2
                
                # Muon 更新
                # … (复杂数学见 optim.py)
```

3.8 精度管理 (fp8.py)

```python
# 全局精度控制
COMPUTE_DTYPE = torch.bfloat16  # 默认 bf16

# 精度映射
dtype_map = {
    'float32': torch.float32,
    'bfloat16': torch.bfloat16,
    'float16': torch.float16,
    'fp8': torch.float8_e4m3fn
}

# FP8 支持
if hardware_supports_fp8:
    COMPUTE_DTYPE = torch.float8_e4m3fn
```

 四、训练脚本详解

4.1 base_train.py (预训练)

```python
# 核心参数
–depth 26          # Transformer 层数(单旋钮控制)
–run "speedrun"    # 运行名称
–model-tag "d26"   # 模型标签

# 自动计算其他超参数
# – 宽度 = f(depth)
# – 头数 = f(width)
# – 学习率 = f(depth)
# – 批次大小 = f(depth)
```

**训练循环**:
```python
for step in range(max_steps):
    # 数据加载
    batch = next(data_loader)
    
    # 前向传播
    logits, loss = model(batch)
    
    # 反向传播
    loss.backward()
    
    # 梯度累积
    if step % grad_accum_steps == 0:
        optimizer.step()
        optimizer.zero_grad()
    
    # 评估
    if step % eval_interval == 0:
        eval_metrics = evaluate(model)
    
    # 保存检查点
    if step % save_interval == 0:
        save_checkpoint(model, optimizer, step)
```

 4.2 chat_sft.py (SFT 微调)

```python
# SFT 训练脚本
python -m scripts.chat_sft \\
    –pretrained_path ./runs/speedrun/checkpoint.pt \\
    –dataset alpaca \\
    –lr 1e-5 \\
    –epochs 3
```

**数据处理**:
```python
# Alpaca 格式
{
    "instruction": "What is the capital of France?",
    "input": "",
    "output": "Paris"
}

# 转换为对话格式
messages = [
    {"role": "user", "content": "What is the capital of France?"},
    {"role": "assistant", "content": "Paris"}
]
```

 4.3 chat_rl.py (RL 对齐)

```python
# RL 训练脚本
python -m scripts.chat_rl \\
    –policy_model ./runs/speedrun/checkpoint.pt \\
    –reference_model ./runs/speedrun/checkpoint.pt \\
    –reward_model anthropic_hhrl \\
    –algorithm ppo
```

**支持算法**:
– PPO (Proximal Policy Optimization)
– GRPO (Group Relative Policy Optimization)
– DPO (Direct Preference Optimization)

4.4 chat_cli.py (CLI 对话)

```python
# 启动对话界面
python -m scripts.chat_cli \\
    –model_path ./runs/speedrun/checkpoint.pt \\
    –temperature 0.8 \\
    –top_k 200
```

**对话示例**:
```
> You: Hello!
> Assistant: Hello! How can I help you today?
> You: Why is the sky blue?
> Assistant: The sky is blue due to Rayleigh scattering…
```

五、核心配置参数

5.1 模型架构参数

```python
# GPT-2 能力模型配置
depth = 26              # Transformer 层数
width = 3072            # 嵌入维度
n_head = 32             # 注意力头数
head_dim = 96           # 每头维度

# 计算验证
total_params ≈ depth × width² × 6
            ≈ 26 × 3072² × 6
            ≈ 1.47B 参数
```

5.2 训练超参数

| 参数 | 默认值 | 说明 |
|——|——–|——|
| `lr` | 6e-4 | 最大学习率 |
| `warmup_ratio` | 0.02 | Warmup 比例 |
| `weight_decay` | 1e-1 | 权重衰减 |
| `beta1` | 0.9 | Adam 一阶矩 |
| `beta2` | 0.95 | Adam 二阶矩 |
| `grad_clip` | 1.0 | 梯度裁剪 |
| `batch_size` | 32 | 每 GPU 批次 |
| `block_size` | 4096 | 序列长度 |

 六、速度排行榜

| # | 时间 | val_bpb | CORE | 描述 |
|—|——|———|——|——|
| 0 | 168h | – | 0.2565 | 原始 GPT-2 |
| 1 | 3.04h | 0.74833 | 0.2585 | d24 基线 |
| 2 | 2.91h | 0.74504 | 0.2578 | d26 + FP8 |
| 3 | 2.76h | 0.74645 | 0.2602 | 批次大小 1M |
| 4 | 2.02h | 0.71854 | 0.2571 | NVidia ClipMixture |
| 5 | 1.80h | 0.71808 | 0.2690 | 自主研究 Round 1 |
| 6 | 1.65h | 0.71800 | 0.2626 | 自主研究 Round 2 |

**目标**:GPT-2 CORE 分数 0.256525,当前最佳 0.2626

七、精度支持

| 硬件 | 默认 dtype | 说明 |
|——|———–|——|
| CUDA SM 80+ | bfloat16 | 原生 bf16 tensor 核心 |
| CUDA SM < 80 | float32 | 无 bf16 |
| CPU/MPS | float32 | 安全默认 |

**覆盖方式**:
```bash
NANochat_DTYPE=bfloat16 python -m scripts.base_train –depth=26
```

八、一键运行命令

 8.1 复现 GPT-2

```bash
bash runs/speedrun.sh
```

8.2 训练小模型 (CPU/MPS)

```bash
bash runs/runcpu.sh
```

8.3 扩展律实验

```bash
bash runs/scaling_laws.sh
```

 九、与 nanoGPT 对比总结

| 维度 | nanoGPT | nanochat |
|——|———|———-|
| **代码量** | ~750 行 | ~6000+ 行 |
| **架构** | 扁平 | 模块化 |
| **功能** | 预训练 + 微调 | 预训练 + SFT + RLHF + CLI |
| **多 GPU** | DDP | DDP + FSDP |
| **精度** | bf16/fp16 | bf16/fp16/fp32/FP8 |
| **推理** | 无优化 | KV Cache 引擎 |
| **效率** | 4 天 (8xA100) | 1.65 小时 (8xH100) |
| **成本** | ~$100+ | ~$48 |
| **状态** | 已弃用 | 活跃维护 |

 十、核心理念

1. **单旋钮设计**:`–depth` 控制所有超参数
2. **计算最优**:自动计算最优宽度、头数、学习率
3. **完整流水线**:从预训练到对话一键完成
4. **极致效率**:通过 Flash Attention、FP8、梯度累积优化
 

赞(0)
未经允许不得转载:171主机测评 » nanochat 核心结构
分享到: 更多 (0)

评论 抢沙发

  • 昵称 (必填)
  • 邮箱 (必填)
  • 网址