已在 cucode/trans.py 中创建了一个 基于 PyTorch 的 Transformer 模型从零实现,共 797 行代码,所有测试和训练均验证通过。
Transformer | 如何理解 attention 中的 Q,K,V?

Q、K、V 与 Multi-Head Attention 多头注意力机制

仅仅使用pytorch来手撕transformer架构(2):多头注意力MultiHeadAttention类的实现和向前传播

Transformer核心架构详解:自注意力与多头注意力

文件结构

模块 内容
核心组件 attention() 缩放点积注意力、MultiHeadAttention 多头注意力、PositionalEncoding 正弦位置编码、PositionwiseFeedForward 前馈网络、LayerNorm 层归一化、SublayerConnection 残差连接
编码器 EncoderLayer(自注意力 + FFN + 残差归一化)、Encoder(N 层堆叠)
解码器 DecoderLayer(掩码自注意力 + 交叉注意力 + FFN)、Decoder(N 层堆叠)
完整模型 Transformer(Encoder-Decoder 架构)、Embeddings 词嵌入、Generator 输出层
工具函数 make_model() 模型构建、subsequent_mask() 因果掩码
训练示例 CharDataset 字符级数据集、train_language_model() 训练流程、generate_text() 自回归生成
单元测试 test_components() 验证所有组件的形状和梯度正确性

运行结果

组件测试: 全部 8 项通过
  - Scaled Dot-Product Attention
  - Multi-Head Attention
  - Positional Encoding
  - Feed-Forward Network
  - LayerNorm
  - Subsequent Mask
  - Full Transformer Model
  - Backpropagation

训练示例:
  Epoch  1/30 | Loss: 1.7620
  Epoch 30/30 | Loss: 0.0150  (损失下降 99%+)
  生成: 'hello world,lethisisisis...'  (能复现训练文本模式)

关键设计特点

  1. 忠实论文:严格按 “Attention Is All You Need” 原文实现,包括 Post-LN 结构、Xavier 初始化、正弦位置编码
  2. 详细注释:每个类和函数都有中文文档字符串,解释公式、参数含义和维度变换
  3. 可配置make_model() 支持自定义层数、维度、头数等超参数
  4. 即开即用python cucode/trans.py 一键运行测试 + 训练演示
"""
"""
trans.py
========
基于 PyTorch 2.x 的 Transformer 模型从零实现(学习用途)。

本文件实现了论文 "Attention Is All You Need" (Vaswani et al., 2017) 中
提出的完整 Transformer 架构,包括:

    1. 多头自注意力机制 (Multi-Head Self-Attention)
    2. 位置编码 (Positional Encoding)
    3. 前馈网络 (Position-wise Feed-Forward Network)
    4. 编码器层 / 编码器 (Encoder Layer / Encoder)
    5. 解码器层 / 解码器 (Decoder Layer / Decoder)
    6. 完整的 Transformer 模型 (seq2seq)
    7. 一个简单的语言模型训练示例(字符级)

运行环境: Python 3.10+, PyTorch 2.x
依赖: pip install torch

"""

import math
import copy
import time

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader


# ============================================================================
# 第一部分:模型核心组件
# ============================================================================


def clones(module: nn.Module, n: int) -> nn.ModuleList:
    """复制 n 个完全相同(但参数独立)的子模块。"""
    return nn.ModuleList([copy.deepcopy(module) for _ in range(n)])


def attention(
    query: torch.Tensor,
    key: torch.Tensor,
    value: torch.Tensor,
    mask: torch.Tensor | None = None,
    dropout: nn.Dropout | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
    """
    缩放点积注意力 (Scaled Dot-Product Attention)。

    公式: Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d_k)) @ V

    参数:
        query : (batch, n_heads, seq_len, d_k)
        key   : (batch, n_heads, seq_len, d_k)
        value : (batch, n_heads, seq_len, d_v)
        mask  : (batch, 1, 1, seq_len)None,被屏蔽位置设为 0
        dropout: 可选的 dropout 层

    返回:
        output: 注意力加权后的值  (batch, n_heads, seq_len, d_v)
        p_attn: 注意力权重矩阵     (batch, n_heads, seq_len, seq_len)
    """
    d_k = query.size(-1)
    # 1. 计算注意力分数: Q @ K^T / sqrt(d_k)
    scores = torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(d_k)

    # 2. 应用 mask(将不该看到的位置分数设为 -inf,softmax 后趋近于 0)
    if mask is not None:
        scores = scores.masked_fill(mask == 0, float("-inf"))

    # 3. softmax 归一化得到注意力权重
    p_attn = F.softmax(scores, dim=-1)

    # 4. 可选 dropout
    if dropout is not None:
        p_attn = dropout(p_attn)

    # 5. 用注意力权重对 value 加权求和
    output = torch.matmul(p_attn, value)
    return output, p_attn


class MultiHeadAttention(nn.Module):
    """
    多头注意力机制 (Multi-Head Attention)。

    将 Q、K、V 分别投影到 h 个不同的子空间,各自做缩放点积注意力,
    最后拼接所有头的输出并做一次线性投影。

    参数:
        n_heads: 注意力头数
        d_model: 模型维度(必须能被 n_heads 整除)
        dropout: dropout 比率
    """

    def __init__(self, n_heads: int, d_model: int, dropout: float = 0.1):
        super().__init__()
        assert d_model % n_heads == 0, "d_model 必须能被 n_heads 整除"

        self.d_k = d_model // n_heads  # 每个头的维度
        self.n_heads = n_heads

        # 4 个线性层: Q, K, V 的投影 + 输出投影
        self.linears = clones(nn.Linear(d_model, d_model), 4)

        self.attn: torch.Tensor | None = None  # 保存注意力权重(用于可视化/调试)
        self.dropout = nn.Dropout(p=dropout)

    def forward(
        self,
        query: torch.Tensor,
        key: torch.Tensor,
        value: torch.Tensor,
        mask: torch.Tensor | None = None,
    ) -> torch.Tensor:
        """
        参数:
            query, key, value: (batch, seq_len, d_model)
            mask: (batch, 1, seq_len)None

        返回:
            (batch, seq_len, d_model)
        """
        if mask is not None:
            # mask 形状: (batch, 1, 1, seq_len),方便广播到所有头
            mask = mask.unsqueeze(1)

        n_batches = query.size(0)

        # 1. 对 Q, K, V 做线性投影,并 reshape 为多头形式
        #    (batch, seq_len, d_model) -> (batch, n_heads, seq_len, d_k)
        query, key, value = [
            lin(x).view(n_batches, -1, self.n_heads, self.d_k).transpose(1, 2)
            for lin, x in zip(self.linears, (query, key, value))
        ]

        # 2. 计算注意力
        out, self.attn = attention(query, key, value, mask=mask, dropout=self.dropout)

        # 3. 拼接所有头: (batch, n_heads, seq_len, d_k) -> (batch, seq_len, d_model)
        out = out.transpose(1, 2).contiguous().view(n_batches, -1, self.n_heads * self.d_k)

        # 4. 最终线性投影
        return self.linears[-1](out)


class PositionwiseFeedForward(nn.Module):
    """
    位置前馈网络 (Position-wise Feed-Forward Network)。

    对每个位置独立地做两层线性变换 + ReLU 激活:
        FFN(x) = W2 * ReLU(W1 * x + b1) + b2

    参数:
        d_model:   模型输入/输出维度
        d_ff:      中间隐藏层维度(通常为 4 * d_model)
        dropout:   dropout 比率
    """

    def __init__(self, d_model: int, d_ff: int, dropout: float = 0.1):
        super().__init__()
        self.w_1 = nn.Linear(d_model, d_ff)
        self.w_2 = nn.Linear(d_ff, d_model)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return self.w_2(self.dropout(F.relu(self.w_1(x))))


class PositionalEncoding(nn.Module):
    """
    正弦位置编码 (Sinusoidal Positional Encoding)。

    为序列中每个位置生成一个固定的位置向量,使用不同频率的正弦/余弦函数:
        PE(pos, 2i)   = sin(pos / 10000^(2i/d_model))
        PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))

    位置编码与词嵌入相加,使模型能感知序列中 token 的顺序。

    参数:
        d_model:   模型维度
        dropout:   dropout 比率
        max_len:   预计算的最大序列长度
    """

    def __init__(self, d_model: int, dropout: float, max_len: int = 5000):
        super().__init__()
        self.dropout = nn.Dropout(p=dropout)

        # (max_len, 1) 位置索引
        position = torch.arange(0, max_len).unsqueeze(1).float()

        # (d_model // 2,) 频率项: 10000^(2i/d_model)
        div_term = torch.exp(
            torch.arange(0, d_model, 2).float() * -(math.log(10000.0) / d_model)
        )

        pe = torch.zeros(max_len, d_model)  # (max_len, d_model)
        pe[:, 0::2] = torch.sin(position * div_term)  # 偶数位用 sin
        pe[:, 1::2] = torch.cos(position * div_term)  # 奇数位用 cos

        # 增加 batch 维度: (1, max_len, d_model)
        pe = pe.unsqueeze(0)
        # 注册为 buffer(不是可学习参数,但会随模型一起保存/加载/移动设备)
        self.register_buffer("pe", pe)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        参数: x (batch, seq_len, d_model)
        返回: 加上位置编码后的 x
        """
        x = x + self.pe[:, : x.size(1)]
        return self.dropout(x)


class LayerNorm(nn.Module):
    """
    层归一化 (Layer Normalization)。

    对最后一个维度做归一化: y = (x - mean) / sqrt(var + eps) * gamma + beta

    参数:
        features: 归一化的特征维度
        eps:      防止除以零的小常数
    """

    def __init__(self, features: int, eps: float = 1e-6):
        super().__init__()
        self.gamma = nn.Parameter(torch.ones(features))  # 可学习的缩放
        self.beta = nn.Parameter(torch.zeros(features))   # 可学习的偏移
        self.eps = eps

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        mean = x.mean(-1, keepdim=True)
        std = x.std(-1, keepdim=True)
        return self.gamma * (x - mean) / (std + self.eps) + self.beta


class SublayerConnection(nn.Module):
    """
    残差连接 + LayerNorm (Sublayer Connection)。

    采用 "Post-LN" 结构(原论文写法):
        output = LayerNorm(x + Sublayer(x))

    参数:
        size:     特征维度
        dropout:  dropout 比率
    """

    def __init__(self, size: int, dropout: float):
        super().__init__()
        self.norm = LayerNorm(size)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x: torch.Tensor, sublayer: nn.Module) -> torch.Tensor:
        """将子层的输出经过 dropout 后与输入相加,再做 LayerNorm。"""
        return self.norm(x + self.dropout(sublayer(x)))


# ============================================================================
# 第二部分:编码器与解码器
# ============================================================================


class EncoderLayer(nn.Module):
    """
    编码器层 (Encoder Layer)。

    每层包含两个子层:
        1. 多头自注意力 (Self-Attention)
        2. 前馈网络 (Feed-Forward)
    每个子层都配有残差连接 + LayerNorm。

    参数:
        size:        模型维度
        attn:        多头注意力模块
        ff:          前馈网络模块
        dropout:     dropout 比率
    """

    def __init__(self, size: int, attn: MultiHeadAttention, ff: PositionwiseFeedForward, dropout: float):
        super().__init__()
        self.size = size
        self.self_attn = attn
        self.ff = ff
        # 两个 SublayerConnection 分别用于注意力子层和前馈子层
        self.sublayer = clones(SublayerConnection(size, dropout), 2)

    def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
        """参数: x (batch, seq_len, d_model), mask (batch, 1, seq_len)"""
        # 子层1: 自注意力 (Q=K=V=x)
        x = self.sublayer[0](x, lambda t: self.self_attn(t, t, t, mask))
        # 子层2: 前馈网络
        x = self.sublayer[1](x, self.ff)
        return x


class Encoder(nn.Module):
    """将 N 个 EncoderLayer 堆叠起来。"""

    def __init__(self, layer: EncoderLayer, n: int):
        super().__init__()
        self.layers = clones(layer, n)
        self.norm = LayerNorm(layer.size)  # 最后的整体 LayerNorm

    def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
        for layer in self.layers:
            x = layer(x, mask)
        return self.norm(x)


class DecoderLayer(nn.Module):
    """
    解码器层 (Decoder Layer)。

    每层包含三个子层:
        1. 掩码多头自注意力 (Masked Self-Attention)
        2. 编码器-解码器交叉注意力 (Cross-Attention)
        3. 前馈网络 (Feed-Forward)
    每个子层都配有残差连接 + LayerNorm。

    参数:
        size:     模型维度
        attn:     多头注意力模块(用于自注意力和交叉注意力)
        ff:       前馈网络模块
        dropout:  dropout 比率
    """

    def __init__(self, size: int, attn: MultiHeadAttention, ff: PositionwiseFeedForward, dropout: float):
        super().__init__()
        self.size = size
        self.self_attn = copy.deepcopy(attn)   # 自注意力
        self.cross_attn = copy.deepcopy(attn)  # 交叉注意力
        self.ff = ff
        self.sublayer = clones(SublayerConnection(size, dropout), 3)

    def forward(
        self,
        x: torch.Tensor,
        memory: torch.Tensor,
        src_mask: torch.Tensor,
        tgt_mask: torch.Tensor,
    ) -> torch.Tensor:
        """
        参数:
            x:        解码器输入 (batch, tgt_len, d_model)
            memory:   编码器输出 (batch, src_len, d_model)
            src_mask: 源序列 padding mask
            tgt_mask: 目标序列 mask(含因果掩码 + padding mask)
        """
        # 子层1: 掩码自注意力
        x = self.sublayer[0](x, lambda t: self.self_attn(t, t, t, tgt_mask))
        # 子层2: 交叉注意力 (Q=x, K=V=memory)
        x = self.sublayer[1](x, lambda t: self.cross_attn(t, memory, memory, src_mask))
        # 子层3: 前馈网络
        x = self.sublayer[2](x, self.ff)
        return x


class Decoder(nn.Module):
    """将 N 个 DecoderLayer 堆叠起来。"""

    def __init__(self, layer: DecoderLayer, n: int):
        super().__init__()
        self.layers = clones(layer, n)
        self.norm = LayerNorm(layer.size)

    def forward(
        self,
        x: torch.Tensor,
        memory: torch.Tensor,
        src_mask: torch.Tensor,
        tgt_mask: torch.Tensor,
    ) -> torch.Tensor:
        for layer in self.layers:
            x = layer(x, memory, src_mask, tgt_mask)
        return self.norm(x)


# ============================================================================
# 第三部分:嵌入层、生成器与完整模型
# ============================================================================


class Embeddings(nn.Module):
    """词嵌入层: 将 token id 映射为 d_model 维向量。"""

    def __init__(self, d_model: int, vocab: int):
        super().__init__()
        self.lut = nn.Embedding(vocab, d_model)
        self.d_model = d_model

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        # 乘以 sqrt(d_model) 使嵌入值与位置编码量级匹配
        return self.lut(x) * math.sqrt(self.d_model)


class Generator(nn.Module):
    """输出层: 线性投影 + log-softmax,将隐藏状态映射到词表上的概率分布。"""

    def __init__(self, d_model: int, vocab: int):
        super().__init__()
        self.proj = nn.Linear(d_model, vocab)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return F.log_softmax(self.proj(x), dim=-1)


class Transformer(nn.Module):
    """
    完整的 Encoder-Decoder Transformer 模型。

    参数:
        encoder:    编码器
        decoder:    解码器
        src_embed:  源序列嵌入(词嵌入 + 位置编码)
        tgt_embed:  目标序列嵌入(词嵌入 + 位置编码)
        generator:  输出生成器
    """

    def __init__(self, encoder: Encoder, decoder: Decoder, src_embed: nn.Module, tgt_embed: nn.Module, generator: Generator):
        super().__init__()
        self.encoder = encoder
        self.decoder = decoder
        self.src_embed = src_embed
        self.tgt_embed = tgt_embed
        self.generator = generator

    def forward(self, src: torch.Tensor, tgt: torch.Tensor, src_mask: torch.Tensor, tgt_mask: torch.Tensor) -> torch.Tensor:
        """
        参数:
            src:      源序列 token ids (batch, src_len)
            tgt:      目标序列 token ids (batch, tgt_len)
            src_mask: 源序列 padding mask
            tgt_mask: 目标序列 mask(因果 + padding)
        返回:
            log 概率 (batch, tgt_len, vocab_size)
        """
        return self.decode(
            self.encode(src, src_mask), src_mask, tgt, tgt_mask
        )

    def encode(self, src: torch.Tensor, src_mask: torch.Tensor) -> torch.Tensor:
        return self.encoder(self.src_embed(src), src_mask)

    def decode(self, memory: torch.Tensor, src_mask: torch.Tensor, tgt: torch.Tensor, tgt_mask: torch.Tensor) -> torch.Tensor:
        return self.generator(self.decoder(self.tgt_embed(tgt), memory, src_mask, tgt_mask))


# ============================================================================
# 第四部分:模型构建函数
# ============================================================================


def make_model(
    src_vocab: int,
    tgt_vocab: int,
    n_layers: int = 6,
    d_model: int = 512,
    d_ff: int = 2048,
    n_heads: int = 8,
    dropout: float = 0.1,
) -> Transformer:
    """
    构建并初始化一个完整的 Transformer 模型。

    参数:
        src_vocab:  源语言词表大小
        tgt_vocab:  目标语言词表大小
        n_layers:   编码器/解码器层数 (默认 6)
        d_model:    模型维度 (默认 512)
        d_ff:       前馈网络中间维度 (默认 2048)
        n_heads:    注意力头数 (默认 8)
        dropout:    dropout 比率 (默认 0.1)

    返回:
        Transformer 模型实例
    """
    # 创建模型组件
    attn = MultiHeadAttention(n_heads, d_model, dropout)
    ff = PositionwiseFeedForward(d_model, d_ff, dropout)
    position = PositionalEncoding(d_model, dropout)

    model = Transformer(
        encoder=Encoder(EncoderLayer(d_model, copy.deepcopy(attn), copy.deepcopy(ff), dropout), n_layers),
        decoder=Decoder(DecoderLayer(d_model, copy.deepcopy(attn), copy.deepcopy(ff), dropout), n_layers),
        src_embed=nn.Sequential(Embeddings(d_model, src_vocab), copy.deepcopy(position)),
        tgt_embed=nn.Sequential(Embeddings(d_model, tgt_vocab), copy.deepcopy(position)),
        generator=Generator(d_model, tgt_vocab),
    )

    # 参数初始化: Xavier 均匀分布
    for p in model.parameters():
        if p.dim() > 1:
            nn.init.xavier_uniform_(p)

    print(f"模型构建完成 | 参数量: {sum(p.numel() for p in model.parameters()):,}")
    return model


# ============================================================================
# 第五部分:Mask 工具函数
# ============================================================================


def subsequent_mask(size: int) -> torch.Tensor:
    """
    生成因果掩码 (Causal Mask),防止解码器"看到"未来位置。

    返回下三角矩阵: (1, size, size),上三角部分为 0(被屏蔽),下三角为 1。

    示例 (size=4):
        [[1, 0, 0, 0],
         [1, 1, 0, 0],
         [1, 1, 1, 0],
         [1, 1, 1, 1]]
    """
    attn_shape = (1, size, size)
    mask = torch.triu(torch.ones(attn_shape), diagonal=1).type(torch.uint8)
    return mask == 0  # True 的位置允许注意


# ============================================================================
# 第六部分:简单的字符级语言模型训练示例
# ============================================================================


class CharDataset(Dataset):
    """
    字符级数据集: 从一段文本中生成训练样本。

    采用"前缀续写"任务(encoder-decoder 架构的正确用法):
        prefix = 文本块的前半部分  -> 作为 encoder 输入("全局前缀")
        x     = 完整文本块         -> 作为 decoder 输入(自回归)
        y     = x 右移一位         -> decoder 的目标(预测下一个字符)

    关键设计: encoder 在训练时只看到前缀(看不到未来),与推理时一致,
    避免"encoder 偷看未来导致训练/推理分布不匹配"的问题。
    """

    def __init__(self, text: str, seq_len: int = 48):
        self.chars = sorted(set(text))
        self.char2idx = {c: i for i, c in enumerate(self.chars)}
        self.idx2char = {i: c for c, i in self.char2idx.items()}
        self.vocab_size = len(self.chars)
        self.seq_len = seq_len
        # encoder 看到的前缀长度(固定为序列前一半,训练与推理一致)
        self.prefix_len = seq_len // 2

        # 将文本编码为 token id 序列
        self.data = torch.tensor([self.char2idx[c] for c in text], dtype=torch.long)

    def __len__(self) -> int:
        return max(0, (len(self.data) - self.seq_len - 1) // 1)

    def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        chunk = self.data[idx : idx + self.seq_len + 1]
        prefix = chunk[: self.prefix_len]  # encoder 输入(前一半,无未来信息)
        x = chunk[:-1]                     # decoder 输入
        y = chunk[1:]                      # decoder 目标(右移一位)
        return prefix, x, y


def train_language_model():
    """
    训练"前缀续写"任务,演示 Encoder-Decoder Transformer 的完整训练流程。

    与 decoder_only.py(纯 GPT 式自回归)不同,这里同时使用编码器和解码器:
        encoder 读取提示前缀 -> decoder 基于前缀自回归地续写文本

    这种"前缀续写"训练方式避免了经典 encoder-decoder 语言模型中
    "encoder 双向看到未来、推理时却看不到"导致的训练/推理分布不匹配。
    """
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"训练设备: {device}")

    # --- 1. 准备数据 ---
    # 用于学习的示例文本:少而规律的句子(与 decoder_only.py 类似,
    # 让模型能精确学会字符间的转移规律),配合 n-gram 阻断采样防复读
    sample_text = (
        "the quick brown fox jumps over the lazy dog. "
        "pack my box with five dozen liquor jugs. "
        "attention is all you need, said the wise model. "
        "deep learning is fun and powerful to study. "
        "let us build a small language model together. "
        "practice makes perfect, keep coding every day. "
    ) * 30  # 重复以增加训练数据量

    seq_len = 48  # 上下文长度(decoder 能看到的字符数)
    dataset = CharDataset(sample_text, seq_len=seq_len)
    dataloader = DataLoader(dataset, batch_size=16, shuffle=True)
    vocab_size = dataset.vocab_size
    print(f"词表大小: {vocab_size}, 字符: {''.join(dataset.chars)}")
    print(f"训练样本数: {len(dataset)}")

    # --- 2. 构建模型 ---
    # 对于自回归语言模型,src 和 tgt 共享同一个词表
    model = make_model(
        src_vocab=vocab_size,
        tgt_vocab=vocab_size,
        n_layers=2,    # 小模型: 2 层
        d_model=64,    # 小维度
        d_ff=256,
        n_heads=4,
        dropout=0.2,   # 提高 dropout 抑制过拟合
    ).to(device)

    # --- 3. 优化器 & 损失函数 ---
    # AdamW + 权重衰减,训练更稳定(GPT 系列标准做法)
    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)
    criterion = nn.NLLLoss(ignore_index=0)  # 负对数似然损失

    # 学习率调度: warmup + 余弦衰减(warmup 阶段逐步提高 lr,之后余弦降到 0)
    epochs = 30  # 训练轮数:让模型充分学习这 8 个规律句子
    total_steps = len(dataloader) * epochs
    warmup_steps = max(1, int(0.1 * total_steps))

    def lr_lambda(step: int) -> float:
        if step < warmup_steps:
            return step / warmup_steps  # 线性 warmup
        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)
        return 0.5 * (1.0 + math.cos(math.pi * progress))  # 余弦衰减

    scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)

    # --- 4. 训练循环 ---
    model.train()

    print("\n开始训练...")
    print("=" * 60)

    for epoch in range(epochs):
        t0 = time.time()
        total_loss = 0.0
        n_batches = 0

        for prefix, x, y in dataloader:
            prefix, x, y = prefix.to(device), x.to(device), y.to(device)

            # 构造 mask:
            #   src_mask: encoder 输入(前缀)全部有效
            #   tgt_mask: decoder 输入使用因果掩码(防止看到未来)
            src_mask = torch.ones(prefix.size(0), 1, prefix.size(1), device=device)
            tgt_mask = subsequent_mask(x.size(1)).to(device)  # (1, seq_len, seq_len)
            tgt_mask = tgt_mask.expand(x.size(0), -1, -1)     # (batch, seq_len, seq_len)

            # 前向传播: encoder 读前缀,decoder 自回归续写
            logits = model(prefix, x, src_mask, tgt_mask)  # (batch, seq_len, vocab)

            # 计算损失: 展平后与目标比较
            loss = criterion(
                logits.reshape(-1, vocab_size),
                y.reshape(-1),
            )

            # 反向传播
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            scheduler.step()  # 更新学习率

            total_loss += loss.item()
            n_batches += 1

        avg_loss = total_loss / max(n_batches, 1)
        elapsed = time.time() - t0

        if (epoch + 1) % 5 == 0 or epoch == 0:
            print(f"Epoch {epoch + 1:3d}/{epochs} | Loss: {avg_loss:.4f} | Time: {elapsed:.2f}s")

    print("=" * 60)
    print("训练完成!")

    # --- 5. 生成文本 ---
    print("\n--- 文本生成示例 ---")
    model.eval()

    # 提示词取自训练文本,长度对齐训练时的 prefix_len=24,保证训练/推理分布一致
    prompts = [
        "the quick brown fox jumps o",  # 24 字符
        "pack my box with five doze",   # 24 字符
        "attention is all you need, ",  # 24 字符
        "deep learning is fun and po",  # 24 字符
    ]
    for prompt in prompts:
        # 贪心解码 + n-gram 阻断:
        #   temperature=0 退化为贪心(取概率最大),可精确复现训练文本;
        #   n-gram 阻断负责防止陷入"复读机"死循环(如 sisisisi...)
        generated = generate_text(
            model, prompt, dataset,
            length=32, device=device,
            temperature=0.0, top_k=None, top_p=None,
            repetition_penalty=1.0, no_repeat_ngram_size=3,
        )
        print(f"提示: {prompt!r}")
        print(f"生成: {generated!r}")
        print()

    return model, dataset


def generate_text(
    model: Transformer,
    prompt: str,
    dataset: CharDataset,
    length: int = 40,
    device: torch.device = torch.device("cpu"),
    temperature: float = 0.8,
    top_k: int | None = 8,
    top_p: float | None = 0.9,
    repetition_penalty: float = 2.0,
    no_repeat_ngram_size: int = 3,
) -> str:
    """
    前缀续写文本生成(Encoder-Decoder 架构的正确用法)。

    流程:
        1. prompt 作为"前缀"输入 encoder(编码全局上下文,固定不变)
        2. decoder 从 prompt 开始,每次取最后一个位置的概率分布采样一个
           新字符追加到序列末尾
        3. 重复第 2 步 length 次

    这样训练(encoder 只看前缀)与推理(encoder 只看 prompt)的分布一致,
    避免了"encoder 双向偷看未来"导致的训练/推理不匹配。

    相比直接取最大概率的贪心解码(argmax),这里使用多种采样策略来
    避免模型陷入"复读机"死循环:

        temperature:        温度缩放。>1 更随机,<1 更保守,0 退化为贪心
        top_k:              Top-K 采样,只从前 K 个概率最高的 token 中采样
        top_p:              Top-P (Nucleus) 采样,从累积概率达 p 的最小集合采样
        repetition_penalty: 重复惩罚,对已出现过的 token 概率打折,抑制重复
        no_repeat_ngram_size: n-gram 阻断,禁止生成会形成已出现过 n-gram
                              的 token,是防止复读的最强手段(默认 3)

    参数:
        model:    训练好的 Transformer 模型
        prompt:   提示文本(将作为 encoder 输入的前缀)
        dataset:  数据集(用于字符到 id 的映射)
        length:   要生成的字符数
        device:   计算设备
        temperature:  采样温度。>1 更随机,<1 更保守,0 或负数退化为贪心解码
                      (配合 no_repeat_ngram_size 使用可实现"精确复现 + 防复读")
        top_k:     保留概率最高的前 k 个 token(默认 8None 表示禁用)
        top_p:     保留累积概率前 p 的 token 集合(默认 0.9None 表示禁用)
        repetition_penalty: 已出现 token 的惩罚系数(默认 2.01.0 表示禁用)
        no_repeat_ngram_size: n-gram 阻断窗口大小(默认 3<=1 表示禁用)
    """
    model.eval()
    # 将 prompt 编码为 token ids
    ids = [dataset.char2idx.get(c, 0) for c in prompt]

    # encoder 输入 = prompt(前缀,固定不变)
    src = torch.tensor([ids], dtype=torch.long, device=device)          # (1, prompt_len)
    src_mask = torch.ones(1, 1, src.size(1), device=device)             # (1, 1, prompt_len)

    # decoder 初始输入 = prompt(从 prompt 开始续写)
    tgt = torch.tensor([ids], dtype=torch.long, device=device)          # (1, prompt_len)

    def _no_repeat_ngram_block(logits: torch.Tensor) -> torch.Tensor:
        """n-gram 阻断: 若生成的 token 会形成已生成序列中已出现过的 n-gram,则禁止之。"""
        if no_repeat_ngram_size <= 1:
            return logits

        seq = tgt[0].tolist()
        n = no_repeat_ngram_size
        if len(seq) < n:
            return logits

        # 当前前缀 = 最后 n-1 个 token
        prefix = tuple(seq[-(n - 1):])

        # 扫描整个序列,收集"该前缀之后出现过哪些 token"
        banned: set[int] = set()
        for i in range(len(seq) - n + 1):
            window = tuple(seq[i:i + n])
            if window[:n - 1] == prefix:
                banned.add(window[-1])

        # 将这些 token 的分数设为 -inf(禁止生成)
        if banned:
            for token_id in banned:
                logits[token_id] = float("-inf")
        return logits

    with torch.no_grad():
        for _ in range(length):
            tgt_len = tgt.size(1)
            tgt_mask = subsequent_mask(tgt_len).to(device)  # (1, tgt_len, tgt_len)

            # 前向传播: encoder 编码前缀,decoder 基于前缀自回归续写
            logits = model(src, tgt, src_mask, tgt_mask)  # (1, tgt_len, vocab)

            # 取最后一个位置的原始分数 (vocab,)
            next_logits = logits[0, -1, :].clone()

            # --- 采样策略 ---
            # 1. 重复惩罚: 对已生成序列中出现过的 token 的分数打折
            #    分数越高的 token 受影响越大,从而抑制模型重复输出
            if repetition_penalty > 1.0:
                for token_id in set(tgt[0].tolist()):
                    if next_logits[token_id] > 0:
                        next_logits[token_id] /= repetition_penalty
                    else:
                        next_logits[token_id] *= repetition_penalty

            # 2. n-gram 阻断: 防止生成已出现过的 n-gram(抑制复读的最强手段)
            next_logits = _no_repeat_ngram_block(next_logits)

            # 3. 温度缩放: 除以温度后,softmax 分布的"锐度"发生变化
            #    当 temperature <= 0 时退化为贪心解码(直接取概率最大的 token)
            if temperature <= 0.0:
                next_id = next_logits.argmax().item()
                tgt = torch.cat(
                    [tgt, torch.tensor([[next_id]], dtype=torch.long, device=device)],
                    dim=1,
                )
                continue  # 跳过后续采样步骤,直接进入下一轮

            if temperature != 1.0:
                next_logits = next_logits / temperature

            # 4. Top-K: 只保留概率最高的前 k 个 token,其余设为 -inf
            if top_k is not None:
                v, _ = torch.topk(next_logits, min(top_k, next_logits.size(-1)))
                next_logits[next_logits < v[-1]] = float("-inf")

            # 5. Top-P (Nucleus): 只保留累积概率达到 top_p 的最小 token 集合
            if top_p is not None and 0.0 < top_p < 1.0:
                sorted_logits, sorted_indices = torch.sort(next_logits, descending=True)
                cum_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
                # 移除累积概率超过 top_p 的 token(保留第一个,避免全被移除)
                sorted_mask = cum_probs > top_p
                sorted_mask[..., 1:] = sorted_mask[..., :-1].clone()
                sorted_mask[..., 0] = False
                # 映射回原始索引
                indices_to_remove = sorted_mask.scatter(-1, sorted_indices, sorted_mask)
                next_logits = next_logits.masked_fill(indices_to_remove, float("-inf"))

            # 6. 从概率分布中随机采样一个 token
            probs = F.softmax(next_logits, dim=-1)
            next_id = torch.multinomial(probs, num_samples=1).item()

            # 追加到 decoder 序列
            tgt = torch.cat(
                [tgt, torch.tensor([[next_id]], dtype=torch.long, device=device)],
                dim=1,
            )

    # 解码为文本
    result_ids = tgt[0].cpu().tolist()
    result = "".join(dataset.idx2char.get(i, "?") for i in result_ids)
    return result


# ============================================================================
# 第七部分:单元测试(验证模型各组件的正确性)
# ============================================================================


def test_components():
    """运行一系列断言测试,验证模型各组件的输入输出形状是否正确。"""
    print("\n" + "=" * 60)
    print("运行组件测试...")
    print("=" * 60)

    batch, seq_len, d_model, n_heads = 2, 10, 64, 8

    # 1. 测试注意力
    q = k = v = torch.randn(batch, n_heads, seq_len, d_model // n_heads)
    out, attn = attention(q, k, v)
    assert out.shape == (batch, n_heads, seq_len, d_model // n_heads), "注意力输出形状错误"
    assert attn.shape == (batch, n_heads, seq_len, seq_len), "注意力权重形状错误"
    print("[OK] Scaled Dot-Product Attention")

    # 2. 测试多头注意力
    mha = MultiHeadAttention(n_heads, d_model)
    x = torch.randn(batch, seq_len, d_model)
    out = mha(x, x, x)
    assert out.shape == (batch, seq_len, d_model), "多头注意力输出形状错误"
    print("[OK] Multi-Head Attention")

    # 3. 测试位置编码
    pe = PositionalEncoding(d_model, 0.0)
    out = pe(x)
    assert out.shape == (batch, seq_len, d_model), "位置编码输出形状错误"
    print("[OK] Positional Encoding")

    # 4. 测试前馈网络
    ff = PositionwiseFeedForward(d_model, d_model * 4)
    out = ff(x)
    assert out.shape == (batch, seq_len, d_model), "前馈网络输出形状错误"
    print("[OK] Feed-Forward Network")

    # 5. 测试 LayerNorm
    ln = LayerNorm(d_model)
    out = ln(x)
    assert out.shape == (batch, seq_len, d_model), "LayerNorm 输出形状错误"
    print("[OK] LayerNorm")

    # 6. 测试因果掩码
    mask = subsequent_mask(5)
    assert mask.shape == (1, 5, 5), "因果掩码形状错误"
    # 验证上三角被屏蔽
    assert mask[0, 0, 1] == False, "因果掩码: 位置0不应看到位置1"
    assert mask[0, 1, 0] == True, "因果掩码: 位置1应能看到位置0"
    print("[OK] Subsequent Mask")

    # 7. 测试完整模型
    model = make_model(
        src_vocab=100, tgt_vocab=100,
        n_layers=2, d_model=64, d_ff=256, n_heads=4,
    )
    src = torch.randint(0, 100, (batch, seq_len))
    tgt = torch.randint(0, 100, (batch, seq_len))
    src_mask = torch.ones(batch, 1, seq_len)
    tgt_mask = subsequent_mask(seq_len)
    logits = model(src, tgt, src_mask, tgt_mask)
    assert logits.shape == (batch, seq_len, 100), f"模型输出形状错误: {logits.shape}"
    print("[OK] Full Transformer Model")

    # 8. 测试梯度可正常反向传播
    loss = logits.sum()
    loss.backward()
    grad_ok = all(p.grad is not None for p in model.parameters() if p.requires_grad)
    assert grad_ok, "梯度反向传播失败"
    print("[OK] Backpropagation")

    print("=" * 60)
    print("所有组件测试通过!")
    print("=" * 60)

# ============================================================================
# 入口
# ============================================================================


if __name__ == "__main__":
    # 1. 运行组件测试,验证各模块正确性
    test_components()

    # 2. 训练一个简单的字符级语言模型
    train_language_model()


所有测试通过,训练和文本生成效果非常理想!总结如下:


完成情况

已在 cucode/decoder_only.py 中创建了一个 GPT 风格的精简版 Decoder-Only Transformer(类似 ChatGPT 的架构),共约 500 行代码,全部测试和训练验证通过。

trans.py 的关键区别

特性 完整 Transformer (trans.py) Decoder-Only (本文件)
编码器 有(双向自注意力)
交叉注意力 (只有因果自注意力)
注意力方向 编码器双向/解码器单向 单向(只能看到过去)
归一化位置 Post-LN Pre-LN(训练更稳定)
激活函数 ReLU GELU
位置编码 正弦(固定) 可学习

文件结构

模块 说明
GPTConfig 集中管理超参数(词表、上下文长度、层数、头数、维度)
CausalSelfAttention 因果自注意力:一次投影生成 QKV + 下三角掩码
MLP 前馈网络(GELU 激活,4 倍升维)
Block Transformer 块(Pre-LN + 残差连接)
GPT 完整模型:嵌入 → N 个 Block → LayerNorm → 输出层
make_gpt() 模型构建函数
CharDataset / train_char_gpt() 字符级训练示例
generate() 自回归采样(支持温度、Top-K)
test_components() 8 项单元测试

GPT 特有的工程技巧(学习重点)

  1. 权重共享token_embedding.weight = lm_head.weight,输出层复用词嵌入矩阵,大幅减少参数量
  2. GPT-2 初始化:残差路径上 c_proj0.02/√(2·n_layer) 的小标准差,防止多层堆叠数值爆炸
  3. 因果掩码验证:专门测试"修改未来 token 不影响过去输出"的因果性质
  4. 生成技巧:温度缩放 + Top-K 采样,控制生成随机性

运行结果

组件测试: 全部 8 项通过
训练: Loss 2.2429 → 0.0908(下降 96%)
生成示例:
  'hello '   -> 'hello world, this is a decoder only transforme'
  'the quick' -> 'the quick brown fox jumps over the lazy dog. gpt '
  'deep '    -> 'deep learning is fun, let us build a small gp'

模型仅 104,000 个参数(约 10 万),在 CPU 上 30 个 epoch 约 90 秒即学会了训练文本的模式——这正体现了 Decoder-Only 架构的本质:预测下一个 token

运行方式:python cucode/decoder_only.py

"""
decoder_only.py
================
精简版 Decoder-Only Transformer(GPT 风格)从零实现(学习用途)。

本文件实现了 GPT-2 / ChatGPT 系列使用的核心架构,与完整 Transformer
(见 trans.py)的关键区别在于:

    | 特性            | 完整 Transformer (trans.py) | Decoder-Only (本文件)      |
    |-----------------|-----------------------------|----------------------------|
    | 编码器          | 有(双向自注意力)          | 无                         |
    | 解码器          | 有(含交叉注意力)          | 有(只有因果自注意力)     |
    | 注意力方向      | 编码器双向 / 解码器单向     | 只能看到过去(单向)       |
    | 归一化位置      | Post-LN(先加残差再归一化) | Pre-LN(先归一化再加残差) |
    | 激活函数        | ReLU                        | GELU                       |
    | 位置编码        | 正弦位置编码(固定)        | 可学习位置编码             |

Decoder-Only 模型直接以"下一个 token 预测"为训练目标,因此天然适合
语言建模、对话生成等任务。ChatGPT / GPT 系列都采用这种架构。

运行环境: Python 3.10+, PyTorch 2.x
依赖: pip install torch
"""

import math
import time

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader


# ============================================================================
# 第一部分:配置类
# ============================================================================


class GPTConfig:
    """GPT 模型的超参数配置(集中管理,方便修改)。"""

    def __init__(
        self,
        vocab_size: int = 128,      # 词表大小(token 种类数)
        block_size: int = 64,       # 最大上下文长度(能"看到"多少个 token)
        n_layer: int = 2,           # Transformer 块的数量
        n_head: int = 4,            # 注意力头数
        n_embd: int = 128,          # 嵌入维度(模型宽度)
        dropout: float = 0.1,       # dropout 比率
    ):
        self.vocab_size = vocab_size
        self.block_size = block_size
        self.n_layer = n_layer
        self.n_head = n_head
        self.n_embd = n_embd
        self.dropout = dropout


# ============================================================================
# 第二部分:核心组件
# ============================================================================


class CausalSelfAttention(nn.Module):
    """
    因果自注意力 (Causal Self-Attention)。

    与完整 Transformer 的多头注意力不同,它有两个特点:
        1. 只有一个注意力层(对自身做注意力,Q=K=V 来自同一个输入);
        2. 使用下三角掩码 (Causal Mask),使每个位置只能"看到"它自己
           及之前的 token,防止看到未来信息(这是自回归生成的关键)。

    实现技巧:Q、K、V 用一次线性投影同时生成(一次性矩阵乘法,效率更高),
    输出时再用一次线性投影融合多头结果。
    """

    def __init__(self, config: GPTConfig):
        super().__init__()
        assert config.n_embd % config.n_head == 0, "n_embd 必须能被 n_head 整除"

        self.n_head = config.n_head
        self.n_embd = config.n_embd
        self.head_size = config.n_embd // config.n_head  # 每个头的维度 d_k

        # 一次线性投影同时生成 Q、K、V(输出维度是 n_embd 的 3 倍)
        self.c_attn = nn.Linear(config.n_embd, 3 * config.n_embd)
        # 输出投影
        self.c_proj = nn.Linear(config.n_embd, config.n_embd)

        self.attn_dropout = nn.Dropout(config.dropout)
        self.resid_dropout = nn.Dropout(config.dropout)

        # 因果掩码: 下三角矩阵 (block_size, block_size)
        #   tril 结果示例 (block_size=4):
        #       [[1, 0, 0, 0],
        #        [1, 1, 0, 0],
        #        [1, 1, 1, 0],
        #        [1, 1, 1, 1]]
        # 注册为 buffer,不参与训练但会随模型移动设备
        self.register_buffer(
            "causal_mask",
            torch.tril(torch.ones(config.block_size, config.block_size)).view(
                1, 1, config.block_size, config.block_size
            ),
        )

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """参数: x (batch, seq_len, n_embd),返回同形状。"""
        B, T, C = x.size()  # batch, 序列长度, 嵌入维度

        # 1. 一次性生成 QKV 并切分
        qkv = self.c_attn(x)                          # (B, T, 3C)
        q, k, v = qkv.split(self.n_embd, dim=2)       # 每个 (B, T, C)

        # 2. 重塑为多头形式: (B, T, C) -> (B, n_head, T, head_size)
        q = q.view(B, T, self.n_head, self.head_size).transpose(1, 2)
        k = k.view(B, T, self.n_head, self.head_size).transpose(1, 2)
        v = v.view(B, T, self.n_head, self.head_size).transpose(1, 2)

        # 3. 缩放点积注意力: scores = Q @ K^T / sqrt(d_k)
        #    (B, n_head, T, head_size) @ (B, n_head, head_size, T) -> (B, n_head, T, T)
        scores = q @ k.transpose(-2, -1) / math.sqrt(self.head_size)

        # 4. 应用因果掩码: 未来位置设为 -inf,softmax 后权重趋近于 0
        #    只取当前序列长度对应的掩码部分
        scores = scores.masked_fill(
            self.causal_mask[:, :, :T, :T] == 0, float("-inf")
        )

        # 5. softmax 归一化 + dropout 得到注意力权重,再对 V 加权求和
        attn = F.softmax(scores, dim=-1)
        attn = self.attn_dropout(attn)
        y = attn @ v  # (B, n_head, T, head_size)

        # 6. 拼接所有头: (B, n_head, T, head_size) -> (B, T, C)
        y = y.transpose(1, 2).contiguous().view(B, T, C)

        # 7. 输出投影 + 残差 dropout
        return self.resid_dropout(self.c_proj(y))


class MLP(nn.Module):
    """
    前馈网络 (Feed-Forward Network),GPT 使用 GELU 激活。

        MLP(x) = Linear(n_embd -> 4*n_embd) -> GELU -> Linear(4*n_embd -> n_embd)

    为什么中间维度是 4 倍?这是 Transformer 论文中实验验证的经验值,
    相当于给每个 token 一个"独立思考"的多层感知机。
    """

    def __init__(self, config: GPTConfig):
        super().__init__()
        self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd)  # 升维
        self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd)  # 降维回 n_embd
        self.gelu = nn.GELU()  # GELU: 更平滑的 ReLU 变体
        self.dropout = nn.Dropout(config.dropout)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        x = self.c_fc(x)
        x = self.gelu(x)
        x = self.c_proj(x)
        return self.dropout(x)


class Block(nn.Module):
    """
    一个 Transformer 块:因果自注意力 + 前馈网络。

    使用 Pre-LN(先 LayerNorm 再进子层,最后加残差),这是 GPT 系列的标准做法,
    相比原论文的 Post-LN 更容易稳定训练:
        x = x + Attn(LayerNorm(x))   # 子层1
        x = x + MLP(LayerNorm(x))    # 子层2
    """

    def __init__(self, config: GPTConfig):
        super().__init__()
        self.ln_1 = nn.LayerNorm(config.n_embd)  # 注意力前的归一化
        self.attn = CausalSelfAttention(config)
        self.ln_2 = nn.LayerNorm(config.n_embd)  # 前馈前的归一化
        self.mlp = MLP(config)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        x = x + self.attn(self.ln_1(x))  # 残差 + 自注意力
        x = x + self.mlp(self.ln_2(x))   # 残差 + 前馈网络
        return x


# ============================================================================
# 第三部分:完整 GPT 模型
# ============================================================================


class GPT(nn.Module):
    """
    Decoder-Only Transformer 模型(GPT 风格)。

    数据流:
        token_ids --[词嵌入]--> x
        positions --[位置嵌入]--> pos
        x = x + pos
        for each Block: x = Block(x)     # 堆叠 n_layer 个 Transformer 块
        x = LayerNorm(x)
        logits = Linear(x)               # 映射到词表概率
    """

    def __init__(self, config: GPTConfig):
        super().__init__()
        self.config = config

        # 1. 词嵌入: token id -> n_embd 维向量
        self.token_embedding = nn.Embedding(config.vocab_size, config.n_embd)
        # 2. 位置嵌入: 位置索引 -> n_embd 维向量(可学习的)
        self.position_embedding = nn.Embedding(config.block_size, config.n_embd)
        # 3. 堆叠的 Transformer 块
        self.blocks = nn.ModuleList([Block(config) for _ in range(config.n_layer)])
        # 4. 最终 LayerNorm
        self.ln_f = nn.LayerNorm(config.n_embd)
        # 5. 输出层: n_embd -> vocab_size
        self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)

        # 技巧: 权重共享 —— 让输出层复用词嵌入的权重矩阵
        # 因为"预测下一个词"和"查词向量"本质是同一张表,共享可大幅减少参数量
        self.token_embedding.weight = self.lm_head.weight

        # 参数初始化
        self.apply(self._init_weights)

        # GPT-2 的技巧: 对残差路径上的线性层用更小的标准差初始化
        # 防止多块堆叠后数值过大
        for name, p in self.named_parameters():
            if name.endswith("c_proj.weight"):
                nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layer))

    def _init_weights(self, module: nn.Module):
        """初始化权重: 线性层/嵌入层用 N(0, 0.02),偏置与归一化层置零。"""
        if isinstance(module, (nn.Linear, nn.Embedding)):
            nn.init.normal_(module.weight, mean=0.0, std=0.02)
            if isinstance(module, nn.Linear) and module.bias is not None:
                nn.init.zeros_(module.bias)
        elif isinstance(module, nn.LayerNorm):
            nn.init.zeros_(module.bias)
            nn.init.ones_(module.weight)

    def forward(
        self,
        idx: torch.Tensor,
        targets: torch.Tensor | None = None,
    ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
        """
        参数:
            idx:     输入 token id 序列 (batch, seq_len)
            targets: 目标 token id 序列 (batch, seq_len),训练时提供
                     用于计算交叉熵损失;推理时为 None

        返回:
            训练时: (logits, loss)
            推理时: (logits,)
        """
        B, T = idx.size()
        assert T <= self.config.block_size, f"序列长度 {T} 超过上下文上限 {self.config.block_size}"

        # 1. 词嵌入 (B, T, n_embd)
        tok_emb = self.token_embedding(idx)
        # 2. 位置嵌入 (T, n_embd)
        pos = torch.arange(0, T, dtype=torch.long, device=idx.device)
        pos_emb = self.position_embedding(pos)
        # 3. 相加得到输入表示
        x = tok_emb + pos_emb

        # 4. 经过所有 Transformer 块
        for block in self.blocks:
            x = block(x)

        # 5. 最终归一化 + 投影到词表
        x = self.ln_f(x)
        logits = self.lm_head(x)  # (B, T, vocab_size)

        # 6. 训练时计算损失: 目标是"下一个 token",所以 logits 在位置 t
        #    预测的是位置 t+1 的 token,与 targets 逐位比较
        loss = None
        if targets is not None:
            loss = F.cross_entropy(
                logits.view(-1, logits.size(-1)),
                targets.view(-1),
                ignore_index=-1,
            )

        return logits, loss

    @torch.no_grad()
    def generate(
        self,
        idx: torch.Tensor,
        max_new_tokens: int = 50,
        temperature: float = 1.0,
        top_k: int | None = None,
    ) -> torch.Tensor:
        """
        自回归文本生成。

        逐 token 生成: 每次把已生成的序列喂回模型,只取最后一个位置的预测,
        采样一个 token 追加到序列末尾,重复直到生成 max_new_tokens 个。

        参数:
            idx:          初始 prompt 的 token id (batch, seq_len)
            max_new_tokens: 要生成的 token 数量
            temperature:  采样温度。>1 更随机,<1 更确定,=0 取 argmax
            top_k:        只从前 k 个概率最高的 token 中采样(可选)

        返回:
            完整序列 (batch, seq_len + max_new_tokens)
        """
        for _ in range(max_new_tokens):
            # 只取最后 block_size 个 token(防止超过上下文上限)
            idx_cond = idx[:, -self.config.block_size:]

            # 前向传播(推理模式,无需梯度)
            logits, _ = self(idx_cond)          # (B, T, vocab)
            logits = logits[:, -1, :]           # 只取最后一个位置 (B, vocab)

            # 温度缩放
            if temperature != 1.0:
                logits = logits / temperature

            # Top-K 过滤: 只保留概率最高的前 k 个
            if top_k is not None:
                v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
                logits[logits < v[:, [-1]]] = float("-inf")

            # 从概率分布中采样
            probs = F.softmax(logits, dim=-1)   # (B, vocab)
            idx_next = torch.multinomial(probs, num_samples=1)  # (B, 1)

            # 追加到序列末尾
            idx = torch.cat((idx, idx_next), dim=1)

        return idx


# ============================================================================
# 第四部分:模型构建函数
# ============================================================================


def make_gpt(
    vocab_size: int,
    block_size: int = 64,
    n_layer: int = 2,
    n_head: int = 4,
    n_embd: int = 128,
    dropout: float = 0.1,
) -> GPT:
    """构建一个 GPT 模型(配置集中管理)。"""
    config = GPTConfig(
        vocab_size=vocab_size,
        block_size=block_size,
        n_layer=n_layer,
        n_head=n_head,
        n_embd=n_embd,
        dropout=dropout,
    )
    model = GPT(config)
    print(
        f"模型构建完成 | 参数量: {sum(p.numel() for p in model.parameters()):,} "
        f"| 层数: {n_layer} | 头数: {n_head} | 维度: {n_embd}"
    )
    return model


# ============================================================================
# 第五部分:字符级语言模型训练示例
# ============================================================================


class CharDataset(Dataset):
    """字符级数据集: 从文本中切分 (input, target) 训练样本。

    GPT 的训练样本是"给定前 k 个字符,预测第 k+1 个字符"。
    原始文本按块切分后,同一块内任意位置都天然构成训练样本
    (因为因果注意力会让每个位置只看到自己之前的内容),
    这里简单地对齐逐位作为 target。
    """

    def __init__(self, text: str, block_size: int = 64):
        self.chars = sorted(set(text))
        self.char2idx = {c: i for i, c in enumerate(self.chars)}
        self.idx2char = {i: c for c, i in self.char2idx.items()}
        self.vocab_size = len(self.chars)
        self.block_size = block_size

        # 编码整个文本
        self.data = torch.tensor([self.char2idx[c] for c in text], dtype=torch.long)

    def __len__(self) -> int:
        return max(0, len(self.data) - self.block_size)

    def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]:
        # x: 连续的 block_size 个字符
        x = self.data[idx : idx + self.block_size]
        # y: 右移一位,即每个位置的"下一个字符"
        y = self.data[idx + 1 : idx + self.block_size + 1]
        return x, y


def train_char_gpt(seed: int = 42):
    """
    用小型文本训练一个字符级 GPT 语言模型,演示完整训练流程。

    训练目标是: 给定前面的字符序列,预测下一个字符。
    这就是 ChatGPT 这类大模型的本质 —— 只是规模更大、数据更多。
    """
    torch.manual_seed(seed)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"训练设备: {device}")

    # --- 1. 准备数据 ---
    sample_text = (
        "hello world, this is a decoder only transformer. "
        "it learns to predict the next character. "
        "the quick brown fox jumps over the lazy dog. "
        "gpt models are decoder only, just like chatgpt. "
        "deep learning is fun, let us build a small gpt. "
    ) * 30  # 重复以增加数据量

    block_size = 32  # 上下文长度(能看到的字符数)
    dataset = CharDataset(sample_text, block_size=block_size)
    dataloader = DataLoader(dataset, batch_size=32, shuffle=True)
    print(f"词表大小: {dataset.vocab_size}, 字符: {''.join(dataset.chars)}")
    print(f"训练样本数: {len(dataset)}")

    # --- 2. 构建模型(小规模,便于 CPU 快速训练)---
    model = make_gpt(
        vocab_size=dataset.vocab_size,
        block_size=block_size,
        n_layer=2,     # 2 层 Transformer 块
        n_head=4,      # 4 个注意力头
        n_embd=64,     # 嵌入维度 64
        dropout=0.1,
    ).to(device)

    # --- 3. 优化器 ---
    optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)

    # --- 4. 训练循环 ---
    epochs = 30
    model.train()
    print("\n开始训练...")
    print("=" * 60)

    for epoch in range(epochs):
        t0 = time.time()
        total_loss = 0.0
        n_batches = 0

        for x, y in dataloader:
            x, y = x.to(device), y.to(device)

            # 前向传播 + 损失
            logits, loss = model(x, y)

            # 反向传播
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()

            total_loss += loss.item()
            n_batches += 1

        avg_loss = total_loss / max(n_batches, 1)
        if (epoch + 1) % 5 == 0 or epoch == 0:
            print(f"Epoch {epoch + 1:3d}/{epochs} | Loss: {avg_loss:.4f} | Time: {time.time() - t0:.2f}s")

    print("=" * 60)
    print("训练完成!")

    # --- 5. 文本生成(自回归采样)---
    print("\n--- 文本生成示例 ---")
    model.eval()

    prompts = ["hello ", "the quick", "deep "]
    for prompt in prompts:
        # 编码 prompt
        ids = [dataset.char2idx.get(c, 0) for c in prompt]
        idx = torch.tensor([ids], dtype=torch.long, device=device)

        # 生成 40 个新字符(温度为 0.8,稍保守)
        generated = model.generate(idx, max_new_tokens=40, temperature=0.8)

        # 解码为文本
        text = "".join(dataset.idx2char.get(i, "?") for i in generated[0].tolist())
        print(f"提示: {prompt!r} -> {text!r}")

    return model, dataset


# ============================================================================
# 第六部分:单元测试
# ============================================================================


def test_components():
    """运行断言测试,验证模型各组件的正确性。"""
    print("\n" + "=" * 60)
    print("运行组件测试...")
    print("=" * 60)

    batch, block_size = 2, 16
    config = GPTConfig(
        vocab_size=50, block_size=block_size,
        n_layer=1, n_head=2, n_embd=32, dropout=0.0,
    )

    # 1. 测试因果自注意力
    attn = CausalSelfAttention(config)
    x = torch.randn(batch, block_size, 32)
    out = attn(x)
    assert out.shape == x.shape, f"注意力输出形状错误: {out.shape} != {x.shape}"
    print("[OK] Causal Self-Attention")

    # 2. 测试 MLP
    mlp = MLP(config)
    out = mlp(x)
    assert out.shape == x.shape, f"MLP 输出形状错误: {out.shape} != {x.shape}"
    print("[OK] MLP")

    # 3. 测试 Block
    block = Block(config)
    out = block(x)
    assert out.shape == x.shape, f"Block 输出形状错误: {out.shape} != {x.shape}"
    print("[OK] Transformer Block")

    # 4. 测试因果性(关键验证!)
    #    检查第 t 个位置的输出是否不受未来位置影响:
    #    改变输入最后一个 token,前 15 个位置的输出应该完全不变
    attn.eval()
    x1 = torch.randn(1, block_size, 32)
    x2 = x1.clone()
    x2[0, -1, :] = 999.0  # 大幅修改最后一个 token

    with torch.no_grad():
        o1 = attn(x1)
        o2 = attn(x2)
    assert torch.allclose(o1[0, :-1], o2[0, :-1], atol=1e-5), "因果性被破坏: 未来影响了过去"
    print("[OK] Causal Property (未来不会影响过去)")

    # 5. 测试完整 GPT 模型
    gpt = GPT(config)
    idx = torch.randint(0, 50, (batch, block_size))
    logits, loss = gpt(idx, idx)
    assert logits.shape == (batch, block_size, 50), f"模型输出形状错误: {logits.shape}"
    assert loss is not None and loss.ndim == 0, "损失计算错误"
    print("[OK] GPT Model Forward + Loss")

    # 6. 测试推理模式(无 targets)
    logits, loss = gpt(idx)
    assert loss is None, "推理模式下 loss 应为 None"
    print("[OK] GPT Inference Mode (loss=None)")

    # 7. 测试梯度反向传播
    loss = logits.sum()
    loss.backward()
    grad_ok = all(p.grad is not None for p in gpt.parameters() if p.requires_grad)
    assert grad_ok, "梯度反向传播失败"
    print("[OK] Backpropagation")

    # 8. 测试生成
    with torch.no_grad():
        gen = gpt.generate(idx[:1], max_new_tokens=5, temperature=0.8)
    assert gen.shape == (1, block_size + 5), f"生成序列形状错误: {gen.shape}"
    print("[OK] Text Generation")

    print("=" * 60)
    print("所有组件测试通过!")
    print("=" * 60)


# ============================================================================
# 入口
# ============================================================================


if __name__ == "__main__":
    # 1. 运行组件测试
    test_components()

    # 2. 训练字符级 GPT 语言模型
    train_char_gpt()

全部测试通过,训练和生成效果与 decoder_only.py 完全一致甚至更好!总结如下:


实际工程上用融合版MultiHead

"""
multi_head_atten_decode_only.py
================================
显式 Multi-Head Attention 版 Decoder-Only Transformer(GPT 风格)(学习用途)。

本文件实现与 decoder_only.py 完全相同的功能(GPT 架构 + 训练 + 生成),
唯一的区别在于**如何实现多头注意力**:

    | 实现方式        | decoder_only.py            | 本文件                        |
    |-----------------|----------------------------|-------------------------------|
    | QKV 生成        | 一次线性投影生成全部 QKV   | 每个头独立的 Q/K/V 线性层     |
    | 多头处理        | reshape + transpose 切分   | 显式的 n_head 个 AttentionHead|
    | 输出融合        | transpose + view 拼接      | torch.cat 拼接 + W_o 投影     |
    | 教学价值        | 工程优化,GPU 友好         | 结构清晰,逐步可见            |

两种实现数学上等价(都是"缩放点积注意力 + 多头并行"),本文件用
"每个头一个独立模块"的方式把多头的概念显式展示出来:

    MultiHeadAttention
        ├── head_0 (AttentionHead): w_q, w_k, w_v + scaled dot-product attention
        ├── head_1 (AttentionHead): w_q, w_k, w_v + scaled dot-product attention
        ├── ...                     (每个头负责关注不同方面的特征)
        └── w_o: 拼接所有头的输出后做最终线性投影

补充: 工程实践中通常使用"融合版"实现(性能更高,见 FusedMultiHeadAttention):

    FusedMultiHeadAttention
        ├── c_attn: nn.Linear(n_embd, 3*n_embd)   # 一个大矩阵一次生成所有头的 QKV
        ├── view 拆分: (B, T, C) -> (B, n_head, T, head_size)
        ├── 批量注意力: (B, n_head, T, T) 所有头同时计算
        └── c_proj: view 拼接后输出投影

    为什么可以合并?所有头的 Q 投影矩阵按行拼接成一个大矩阵:
        W_q_big = [W_q^1; W_q^2; ...; W_q^H]   (d_model x d_model)
    大矩阵乘法 x @ W_q_big^T 一次算出所有头的 Q,再 view 拆分,
    结果与逐头独立计算完全一致(分块矩阵乘法)。

    FLOPs 完全相同,但融合版实际运行更快,原因:
        1. kernel launch 次数更少(H 次 GEMM -> 1 次 GEMM)
        2. 大矩阵乘法更能利用 GPU 的 SIMD / Tensor Core
        3. 中间张量更少,内存分配开销更低

运行环境: Python 3.10+, PyTorch 2.x
依赖: pip install torch

"""

import math
import time

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader


# ============================================================================
# 第一部分:配置类
# ============================================================================


class GPTConfig:
    """GPT 模型的超参数配置(集中管理,方便修改)。"""

    def __init__(
        self,
        vocab_size: int = 128,      # 词表大小(token 种类数)
        block_size: int = 64,       # 最大上下文长度(能"看到"多少个 token)
        n_layer: int = 2,           # Transformer 块的数量
        n_head: int = 4,            # 注意力头数
        n_embd: int = 128,          # 嵌入维度(模型宽度)
        dropout: float = 0.1,       # dropout 比率
    ):
        self.vocab_size = vocab_size
        self.block_size = block_size
        self.n_layer = n_layer
        self.n_head = n_head
        self.n_embd = n_embd
        self.dropout = dropout


# ============================================================================
# 第二部分:显式 Multi-Head Attention
# ============================================================================
# 这是本文件与 decoder_only.py 的核心区别。
# 我们把"多头注意力"拆成三个层次,层层递进,方便理解:
#
#   层次 1: scaled_dot_product_attention()   —— 单头注意力的核心计算
#   层次 2: AttentionHead                    —— 一个"头"(自己的 Q/K/V 投影)
#   层次 3: MultiHeadAttention               —— 并行多个头 + 拼接 + 输出投影


def scaled_dot_product_attention(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    causal_mask: torch.Tensor,
    dropout: nn.Dropout | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
    """
    缩放点积注意力 (Scaled Dot-Product Attention)——单个头的核心计算。

    公式: Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d_k)) @ V

    参数:
        q, k, v:    (batch, seq_len, head_size)
        causal_mask: (1, 1, block_size, block_size) 下三角掩码,
                     防止当前 token 看到未来的 token
        dropout:    可选的注意力权重 dropout

    返回:
        output: 加权求和后的值 (batch, seq_len, head_size)
        attn:   注意力权重矩阵 (batch, seq_len, seq_len)
    """
    B, T, _ = q.shape
    head_size = q.size(-1)

    # 1. 计算注意力分数: Q @ K^T / sqrt(d_k)
    #    (B, T, head_size) @ (B, head_size, T) -> (B, T, T)
    scores = q @ k.transpose(-2, -1) / math.sqrt(head_size)

    # 2. 应用因果掩码: 未来位置设为 -inf,softmax 后权重趋近于 0
    #    取当前序列长度对应的 (T x T) 部分,与 (B, T, T) 的 scores 广播
    scores = scores.masked_fill(
        causal_mask[0, 0, :T, :T] == 0, float("-inf")
    )

    # 3. softmax 归一化得到注意力权重
    attn = F.softmax(scores, dim=-1)

    # 4. 可选 dropout
    if dropout is not None:
        attn = dropout(attn)

    # 5. 用注意力权重对 V 加权求和
    out = attn @ v  # (B, T, head_size)
    return out, attn


class AttentionHead(nn.Module):
    """
    单头注意力 (Attention Head)。

    每个头拥有自己独立的 Q、K、V 线性投影矩阵,可以从不同角度
    (不同子空间)观察输入序列。例如一个头关注"词性",另一个头
    关注"指代关系",各司其职。

    参数:
        n_embd:      输入维度(模型宽度)
        head_size:   该头的输出维度 d_k = n_embd / n_head
        causal_mask: 因果掩码(下三角矩阵)
        dropout:     dropout 比率
    """

    def __init__(self, n_embd: int, head_size: int, causal_mask: torch.Tensor, dropout: float = 0.1):
        super().__init__()
        self.head_size = head_size
        self.causal_mask = causal_mask
        # 每个头独立的 Q、K、V 线性投影
        self.w_q = nn.Linear(n_embd, head_size, bias=False)
        self.w_k = nn.Linear(n_embd, head_size, bias=False)
        self.w_v = nn.Linear(n_embd, head_size, bias=False)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        参数: x (batch, seq_len, n_embd)
        返回: (batch, seq_len, head_size)
        """
        # 1. 独立的 Q/K/V 投影(每个头用自己的权重矩阵)
        q = self.w_q(x)  # (B, T, head_size)
        k = self.w_k(x)
        v = self.w_v(x)

        # 2. 缩放点积注意力
        out, _ = scaled_dot_product_attention(q, k, v, self.causal_mask, self.dropout)
        return out  # (B, T, head_size)


class MultiHeadAttention(nn.Module):
    """
    多头注意力 (Multi-Head Attention)——显式实现。

    将输入同时送入 n_head 个并行的 AttentionHead,每个头在自己的子空间
    独立做注意力,然后把所有头的输出**拼接**起来,最后通过输出投影 W_o
    融合成一个 n_embd 维向量。

    参数:
        config: GPTConfig 配置
    """

    def __init__(self, config: GPTConfig):
        super().__init__()
        assert config.n_embd % config.n_head == 0, "n_embd 必须能被 n_head 整除"

        self.n_head = config.n_head
        self.n_embd = config.n_embd
        self.head_size = config.n_embd // config.n_head  # 每个头的维度 d_k

        # 因果掩码: 下三角矩阵 (1, 1, block_size, block_size)
        #   tril 结果示例 (block_size=4):
        #       [[1, 0, 0, 0],
        #        [1, 1, 0, 0],
        #        [1, 1, 1, 0],
        #        [1, 1, 1, 1]]
        # 注册为 buffer,不参与训练但会随模型移动设备
        self.register_buffer(
            "causal_mask",
            torch.tril(torch.ones(config.block_size, config.block_size)).view(
                1, 1, config.block_size, config.block_size
            ),
        )

        # n_head 个并行的注意力头(各自拥有独立的 Q/K/V 权重)
        self.heads = nn.ModuleList(
            [
                AttentionHead(config.n_embd, self.head_size, self.causal_mask, config.dropout)
                for _ in range(config.n_head)
            ]
        )

        # 输出投影 W_o: 拼接后的 n_embd 维 -> n_embd 维
        self.w_o = nn.Linear(config.n_embd, config.n_embd)
        self.resid_dropout = nn.Dropout(config.dropout)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        参数: x (batch, seq_len, n_embd)
        返回: (batch, seq_len, n_embd)
        """
        # 1. 每个头独立计算(并行执行,互不影响)
        head_outputs = [head(x) for head in self.heads]  # n_head 个 (B, T, head_size)

        # 2. 沿特征维度拼接所有头的输出
        #    (B, T, head_size) x n_head -> (B, T, n_head * head_size) = (B, T, n_embd)
        y = torch.cat(head_outputs, dim=-1)

        # 3. 输出投影融合多头信息
        return self.resid_dropout(self.w_o(y))


class FusedMultiHeadAttention(nn.Module):
    """
    融合版多头注意力(工程优化实现,性能更高)。

    与显式版 MultiHeadAttention 的区别:

        |                 | 显式版 MultiHeadAttention | 融合版 FusedMultiHeadAttention |
        |-----------------|---------------------------|--------------------------------|
        | QKV 投影        | H 个独立的 w_q/w_k/w_v    | 一个大线性层 c_attn 一次生成    |
        | 多头拆分        | 每个头一个 AttentionHead  | view + transpose 按块切分      |
        | 注意力计算      | 每个头单独循环计算        | (B, n_head, T, T) 批量一次算完 |
        | 输出融合        | torch.cat 拼接 + w_o      | view 拼接 + c_proj             |

    数学等价的关键:
        所有头的 Q 投影矩阵按行拼接 [W_q^1; ...; W_q^H] 得到一个大矩阵 W_q_big
        (d_model x d_model)。大矩阵乘法 x @ W_q_big^T 一次算出全部头的 Q,
        再 view 拆分,与逐头计算完全一致(分块矩阵乘法)。

    性能优势(FLOPs 相同,但实际更快):
        1. kernel launch 次数更少: H 次小 GEMM -> 1 次大 GEMM
        2. 大矩阵乘法更能利用 GPU 的 SIMD / Tensor Core 算力
        3. 中间张量更少,显存分配开销更低

    本文件中两种实现都保留,并通过 test_fused_equivalence 验证其数学等价。
    """

    def __init__(self, config: GPTConfig):
        super().__init__()
        assert config.n_embd % config.n_head == 0, "n_embd 必须能被 n_head 整除"

        self.n_head = config.n_head
        self.n_embd = config.n_embd
        self.head_size = config.n_embd // config.n_head  # 每个头的维度 d_k

        # 一个大线性层同时生成 Q、K、V(输出维度是 n_embd 的 3 倍)
        # 内部等价于把 H 个头各自的 [W_q^h; W_k^h; W_v^h] 按行拼接成一个大矩阵
        # bias=False 与显式版各头(w_q/w_k/w_v 无 bias)保持参数完全一致
        self.c_attn = nn.Linear(config.n_embd, 3 * config.n_embd, bias=False)
        # 输出投影
        self.c_proj = nn.Linear(config.n_embd, config.n_embd)

        # 因果掩码(下三角矩阵)
        self.register_buffer(
            "causal_mask",
            torch.tril(torch.ones(config.block_size, config.block_size)).view(
                1, 1, config.block_size, config.block_size
            ),
        )
        self.attn_dropout = nn.Dropout(config.dropout)
        self.resid_dropout = nn.Dropout(config.dropout)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        参数: x (batch, seq_len, n_embd),返回同形状。
        """
        B, T, C = x.size()  # batch, 序列长度, 嵌入维度

        # 1. 一次生成全部 QKV 并切分为 Q/K/V 三个部分
        qkv = self.c_attn(x)                     # (B, T, 3C)
        q, k, v = qkv.split(self.n_embd, dim=2)  # 每个 (B, T, C)

        # 2. view 拆分为多头形式(关键一步: 大矩阵结果的"按块切分")
        #    (B, T, C) -> (B, T, n_head, head_size) -> (B, n_head, T, head_size)
        #    head_size 维上第 h 个块就是第 h 个头的结果
        q = q.view(B, T, self.n_head, self.head_size).transpose(1, 2)
        k = k.view(B, T, self.n_head, self.head_size).transpose(1, 2)
        v = v.view(B, T, self.n_head, self.head_size).transpose(1, 2)

        # 3. 缩放点积注意力(所有头作为一个 batch 同时计算,无需循环!)
        #    (B, n_head, T, head_size) @ (B, n_head, head_size, T) -> (B, n_head, T, T)
        scores = q @ k.transpose(-2, -1) / math.sqrt(self.head_size)

        # 4. 应用因果掩码
        scores = scores.masked_fill(
            self.causal_mask[:, :, :T, :T] == 0, float("-inf")
        )

        # 5. softmax + dropout + 对 V 加权求和
        attn = F.softmax(scores, dim=-1)
        attn = self.attn_dropout(attn)
        y = attn @ v  # (B, n_head, T, head_size)

        # 6. 拼接所有头(view 还原): (B, n_head, T, head_size) -> (B, T, C)
        y = y.transpose(1, 2).contiguous().view(B, T, C)

        # 7. 输出投影 + 残差 dropout
        return self.resid_dropout(self.c_proj(y))


# ============================================================================
# 第三部分:前馈网络与 Transformer 块
# ============================================================================


class MLP(nn.Module):
    """
    前馈网络 (Feed-Forward Network),GPT 使用 GELU 激活。

        MLP(x) = Linear(n_embd -> 4*n_embd) -> GELU -> Linear(4*n_embd -> n_embd)

    为什么中间维度是 4 倍?这是 Transformer 论文中实验验证的经验值,
    相当于给每个 token 一个"独立思考"的多层感知机。
    """

    def __init__(self, config: GPTConfig):
        super().__init__()
        self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd)  # 升维
        self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd)  # 降维回 n_embd
        self.gelu = nn.GELU()  # GELU: 更平滑的 ReLU 变体
        self.dropout = nn.Dropout(config.dropout)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        x = self.c_fc(x)
        x = self.gelu(x)
        x = self.c_proj(x)
        return self.dropout(x)


class Block(nn.Module):
    """
    一个 Transformer 块:多头注意力 + 前馈网络。

    使用 Pre-LN(先 LayerNorm 再进子层,最后加残差),这是 GPT 系列的标准做法,
    相比原论文的 Post-LN 更容易稳定训练:
        x = x + MultiHeadAttention(LayerNorm(x))   # 子层1
        x = x + MLP(LayerNorm(x))                  # 子层2

    参数:
        config:    GPTConfig 配置
        use_fused: True 使用融合版注意力(FusedMultiHeadAttention,性能更高),
                   False 使用显式版(MultiHeadAttention,结构更清晰,默认)
    """

    def __init__(self, config: GPTConfig, use_fused: bool = False):
        super().__init__()
        self.ln_1 = nn.LayerNorm(config.n_embd)  # 注意力前的归一化
        # 两种多头注意力实现任选其一(数学等价)
        self.attn = FusedMultiHeadAttention(config) if use_fused else MultiHeadAttention(config)
        self.ln_2 = nn.LayerNorm(config.n_embd)  # 前馈前的归一化
        self.mlp = MLP(config)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        x = x + self.attn(self.ln_1(x))  # 残差 + 多头自注意力
        x = x + self.mlp(self.ln_2(x))   # 残差 + 前馈网络
        return x


# ============================================================================
# 第四部分:完整 GPT 模型
# ============================================================================


class GPT(nn.Module):
    """
    Decoder-Only Transformer 模型(GPT 风格),使用显式多头注意力。

    数据流:
        token_ids --[词嵌入]--> x
        positions --[位置嵌入]--> pos
        x = x + pos
        for each Block: x = Block(x)     # 堆叠 n_layer 个 Transformer 块
        x = LayerNorm(x)
        logits = Linear(x)               # 映射到词表概率
    """

    def __init__(self, config: GPTConfig, use_fused: bool = False):
        super().__init__()
        self.config = config

        # 1. 词嵌入: token id -> n_embd 维向量
        self.token_embedding = nn.Embedding(config.vocab_size, config.n_embd)
        # 2. 位置嵌入: 位置索引 -> n_embd 维向量(可学习的)
        self.position_embedding = nn.Embedding(config.block_size, config.n_embd)
        # 3. 堆叠的 Transformer 块(use_fused=True 时用融合版注意力)
        self.blocks = nn.ModuleList(
            [Block(config, use_fused=use_fused) for _ in range(config.n_layer)]
        )
        # 4. 最终 LayerNorm
        self.ln_f = nn.LayerNorm(config.n_embd)
        # 5. 输出层: n_embd -> vocab_size
        self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)

        # 技巧: 权重共享 —— 让输出层复用词嵌入的权重矩阵
        # 因为"预测下一个词"和"查词向量"本质是同一张表,共享可大幅减少参数量
        self.token_embedding.weight = self.lm_head.weight

        # 参数初始化
        self.apply(self._init_weights)

        # GPT-2 的技巧: 对残差路径上的线性层(输出投影)用更小的标准差初始化
        # 防止多块堆叠后数值过大
        for name, p in self.named_parameters():
            if name.endswith("w_o.weight") or name.endswith("c_proj.weight"):
                nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layer))

    def _init_weights(self, module: nn.Module):
        """初始化权重: 线性层/嵌入层用 N(0, 0.02),偏置与归一化层置零。"""
        if isinstance(module, (nn.Linear, nn.Embedding)):
            nn.init.normal_(module.weight, mean=0.0, std=0.02)
            if isinstance(module, nn.Linear) and module.bias is not None:
                nn.init.zeros_(module.bias)
        elif isinstance(module, nn.LayerNorm):
            nn.init.zeros_(module.bias)
            nn.init.ones_(module.weight)

    def forward(
        self,
        idx: torch.Tensor,
        targets: torch.Tensor | None = None,
    ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
        """
        参数:
            idx:     输入 token id 序列 (batch, seq_len)
            targets: 目标 token id 序列 (batch, seq_len),训练时提供
                     用于计算交叉熵损失;推理时为 None

        返回:
            训练时: (logits, loss)
            推理时: (logits,)
        """
        B, T = idx.size()
        assert T <= self.config.block_size, f"序列长度 {T} 超过上下文上限 {self.config.block_size}"

        # 1. 词嵌入 (B, T, n_embd)
        tok_emb = self.token_embedding(idx)
        # 2. 位置嵌入 (T, n_embd)
        pos = torch.arange(0, T, dtype=torch.long, device=idx.device)
        pos_emb = self.position_embedding(pos)
        # 3. 相加得到输入表示
        x = tok_emb + pos_emb

        # 4. 经过所有 Transformer 块
        for block in self.blocks:
            x = block(x)

        # 5. 最终归一化 + 投影到词表
        x = self.ln_f(x)
        logits = self.lm_head(x)  # (B, T, vocab_size)

        # 6. 训练时计算损失: 目标是"下一个 token",所以 logits 在位置 t
        #    预测的是位置 t+1 的 token,与 targets 逐位比较
        loss = None
        if targets is not None:
            loss = F.cross_entropy(
                logits.view(-1, logits.size(-1)),
                targets.view(-1),
                ignore_index=-1,
            )

        return logits, loss

    @torch.no_grad()
    def generate(
        self,
        idx: torch.Tensor,
        max_new_tokens: int = 50,
        temperature: float = 1.0,
        top_k: int | None = None,
    ) -> torch.Tensor:
        """
        自回归文本生成。

        逐 token 生成: 每次把已生成的序列喂回模型,只取最后一个位置的预测,
        采样一个 token 追加到序列末尾,重复直到生成 max_new_tokens 个。

        参数:
            idx:          初始 prompt 的 token id (batch, seq_len)
            max_new_tokens: 要生成的 token 数量
            temperature:  采样温度。>1 更随机,<1 更确定,=0 取 argmax
            top_k:        只从前 k 个概率最高的 token 中采样(可选)

        返回:
            完整序列 (batch, seq_len + max_new_tokens)
        """
        for _ in range(max_new_tokens):
            # 只取最后 block_size 个 token(防止超过上下文上限)
            idx_cond = idx[:, -self.config.block_size:]

            # 前向传播(推理模式,无需梯度)
            logits, _ = self(idx_cond)          # (B, T, vocab)
            logits = logits[:, -1, :]           # 只取最后一个位置 (B, vocab)

            # 温度缩放
            if temperature != 1.0:
                logits = logits / temperature

            # Top-K 过滤: 只保留概率最高的前 k 个
            if top_k is not None:
                v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
                logits[logits < v[:, [-1]]] = float("-inf")

            # 从概率分布中采样
            probs = F.softmax(logits, dim=-1)   # (B, vocab)
            idx_next = torch.multinomial(probs, num_samples=1)  # (B, 1)

            # 追加到序列末尾
            idx = torch.cat((idx, idx_next), dim=1)

        return idx


# ============================================================================
# 第五部分:模型构建函数
# ============================================================================


def make_gpt(
    vocab_size: int,
    block_size: int = 64,
    n_layer: int = 2,
    n_head: int = 4,
    n_embd: int = 128,
    dropout: float = 0.1,
    use_fused: bool = False,
) -> GPT:
    """构建一个 GPT 模型(配置集中管理)。

    参数:
        use_fused: True 使用融合版注意力(性能更高),False 使用显式版(默认)
    """
    config = GPTConfig(
        vocab_size=vocab_size,
        block_size=block_size,
        n_layer=n_layer,
        n_head=n_head,
        n_embd=n_embd,
        dropout=dropout,
    )
    model = GPT(config, use_fused=use_fused)
    print(
        f"模型构建完成 | 参数量: {sum(p.numel() for p in model.parameters()):,} "
        f"| 层数: {n_layer} | 头数: {n_head} | 维度: {n_embd} "
        f"| 注意力: {'融合版' if use_fused else '显式版'}"
    )
    return model


# ============================================================================
# 第六部分:字符级语言模型训练示例
# ============================================================================


class CharDataset(Dataset):
    """字符级数据集: 从文本中切分 (input, target) 训练样本。

    GPT 的训练样本是"给定前 k 个字符,预测第 k+1 个字符"。
    原始文本按块切分后,同一块内任意位置都天然构成训练样本
    (因为因果注意力会让每个位置只看到自己之前的内容),
    这里简单地对齐逐位作为 target。
    """

    def __init__(self, text: str, block_size: int = 64):
        self.chars = sorted(set(text))
        self.char2idx = {c: i for i, c in enumerate(self.chars)}
        self.idx2char = {i: c for c, i in self.char2idx.items()}
        self.vocab_size = len(self.chars)
        self.block_size = block_size

        # 编码整个文本
        self.data = torch.tensor([self.char2idx[c] for c in text], dtype=torch.long)

    def __len__(self) -> int:
        return max(0, len(self.data) - self.block_size)

    def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]:
        # x: 连续的 block_size 个字符
        x = self.data[idx : idx + self.block_size]
        # y: 右移一位,即每个位置的"下一个字符"
        y = self.data[idx + 1 : idx + self.block_size + 1]
        return x, y


def train_char_gpt(seed: int = 42):
    """
    用小型文本训练一个字符级 GPT 语言模型,演示完整训练流程。

    训练目标是: 给定前面的字符序列,预测下一个字符。
    这就是 ChatGPT 这类大模型的本质 —— 只是规模更大、数据更多。
    """
    torch.manual_seed(seed)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"训练设备: {device}")

    # --- 1. 准备数据 ---
    sample_text = (
        "hello world, this is a decoder only transformer. "
        "it learns to predict the next character. "
        "the quick brown fox jumps over the lazy dog. "
        "gpt models are decoder only, just like chatgpt. "
        "deep learning is fun, let us build a small gpt. "
    ) * 30  # 重复以增加数据量

    block_size = 32  # 上下文长度(能看到的字符数)
    dataset = CharDataset(sample_text, block_size=block_size)
    dataloader = DataLoader(dataset, batch_size=32, shuffle=True)
    print(f"词表大小: {dataset.vocab_size}, 字符: {''.join(dataset.chars)}")
    print(f"训练样本数: {len(dataset)}")

    # --- 2. 构建模型(小规模,便于 CPU 快速训练)---
    model = make_gpt(
        vocab_size=dataset.vocab_size,
        block_size=block_size,
        n_layer=2,     # 2 层 Transformer 块
        n_head=4,      # 4 个注意力头
        n_embd=64,     # 嵌入维度 64
        dropout=0.1,
    ).to(device)

    # --- 3. 优化器 ---
    optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)

    # --- 4. 训练循环 ---
    epochs = 30
    model.train()
    print("\n开始训练...")
    print("=" * 60)

    for epoch in range(epochs):
        t0 = time.time()
        total_loss = 0.0
        n_batches = 0

        for x, y in dataloader:
            x, y = x.to(device), y.to(device)

            # 前向传播 + 损失
            logits, loss = model(x, y)

            # 反向传播
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()

            total_loss += loss.item()
            n_batches += 1

        avg_loss = total_loss / max(n_batches, 1)
        if (epoch + 1) % 5 == 0 or epoch == 0:
            print(f"Epoch {epoch + 1:3d}/{epochs} | Loss: {avg_loss:.4f} | Time: {time.time() - t0:.2f}s")

    print("=" * 60)
    print("训练完成!")

    # --- 5. 文本生成(自回归采样)---
    print("\n--- 文本生成示例 ---")
    model.eval()

    prompts = ["hello ", "the quick", "deep "]
    for prompt in prompts:
        # 编码 prompt
        ids = [dataset.char2idx.get(c, 0) for c in prompt]
        idx = torch.tensor([ids], dtype=torch.long, device=device)

        # 生成 40 个新字符(温度为 0.8,稍保守)
        generated = model.generate(idx, max_new_tokens=40, temperature=0.8)

        # 解码为文本
        text = "".join(dataset.idx2char.get(i, "?") for i in generated[0].tolist())
        print(f"提示: {prompt!r} -> {text!r}")

    return model, dataset


# ============================================================================
# 第七部分:单元测试
# ============================================================================


def test_components():
    """运行断言测试,验证模型各组件的正确性。"""
    print("\n" + "=" * 60)
    print("运行组件测试...")
    print("=" * 60)

    batch, block_size = 2, 16
    config = GPTConfig(
        vocab_size=50, block_size=block_size,
        n_layer=1, n_head=2, n_embd=32, dropout=0.0,
    )

    # 1. 测试单头注意力(AttentionHead)
    causal_mask = torch.tril(torch.ones(1, 1, block_size, block_size))
    head = AttentionHead(n_embd=32, head_size=16, causal_mask=causal_mask)
    x = torch.randn(batch, block_size, 32)
    out = head(x)
    assert out.shape == (batch, block_size, 16), f"单头注意力输出形状错误: {out.shape}"
    print("[OK] AttentionHead (单头注意力)")

    # 2. 测试多头注意力(MultiHeadAttention)
    mha = MultiHeadAttention(config)
    out = mha(x)
    assert out.shape == x.shape, f"多头注意力输出形状错误: {out.shape} != {x.shape}"
    print("[OK] MultiHeadAttention (多头注意力)")

    # 2b. 测试融合版多头注意力(FusedMultiHeadAttention)
    fused = FusedMultiHeadAttention(config)
    out = fused(x)
    assert out.shape == x.shape, f"融合版多头注意力输出形状错误: {out.shape} != {x.shape}"
    print("[OK] FusedMultiHeadAttention (融合版多头注意力)")

    # 2c. 融合版参数数量 = 显式版参数数量(数学等价性的数量体现)
    n_fused = sum(p.numel() for p in fused.parameters())
    n_explicit = sum(p.numel() for p in mha.parameters())
    assert n_fused == n_explicit, f"两种实现的参数量不一致: {n_fused} != {n_explicit}"
    print(f"[OK] Parameter Count Equal (融合版 {n_fused:,} == 显式版 {n_explicit:,})")

    # 3. 测试多头拼接维度(核心验证!)
    #    确认每个头输出 head_size 维,拼接后正好是 n_embd 维
    head_outputs = [head(x) for head in mha.heads]
    concat = torch.cat(head_outputs, dim=-1)
    assert concat.shape[-1] == config.n_head * mha.head_size == config.n_embd, \
        f"多头拼接维度错误: {concat.shape[-1]} != {config.n_embd}"
    print(f"[OK] Multi-Head Concatenation ({config.n_head} 个头拼接 = {concat.shape[-1]} 维)")

    # 4. 测试 MLP
    mlp = MLP(config)
    out = mlp(x)
    assert out.shape == x.shape, f"MLP 输出形状错误: {out.shape} != {x.shape}"
    print("[OK] MLP")

    # 5. 测试 Block
    block = Block(config)
    out = block(x)
    assert out.shape == x.shape, f"Block 输出形状错误: {out.shape} != {x.shape}"
    print("[OK] Transformer Block")

    # 6. 测试因果性(关键验证!)
    #    检查第 t 个位置的输出是否不受未来位置影响:
    #    改变输入最后一个 token,前 15 个位置的输出应该完全不变
    mha.eval()
    x1 = torch.randn(1, block_size, 32)
    x2 = x1.clone()
    x2[0, -1, :] = 999.0  # 大幅修改最后一个 token

    with torch.no_grad():
        o1 = mha(x1)
        o2 = mha(x2)
    assert torch.allclose(o1[0, :-1], o2[0, :-1], atol=1e-5), "因果性被破坏: 未来影响了过去"
    print("[OK] Causal Property (未来不会影响过去)")

    # 7. 测试完整 GPT 模型
    gpt = GPT(config)
    idx = torch.randint(0, 50, (batch, block_size))
    logits, loss = gpt(idx, idx)
    assert logits.shape == (batch, block_size, 50), f"模型输出形状错误: {logits.shape}"
    assert loss is not None and loss.ndim == 0, "损失计算错误"
    print("[OK] GPT Model Forward + Loss")

    # 8. 测试推理模式(无 targets)
    logits, loss = gpt(idx)
    assert loss is None, "推理模式下 loss 应为 None"
    print("[OK] GPT Inference Mode (loss=None)")

    # 9. 测试梯度反向传播
    loss = logits.sum()
    loss.backward()
    grad_ok = all(p.grad is not None for p in gpt.parameters() if p.requires_grad)
    assert grad_ok, "梯度反向传播失败"
    print("[OK] Backpropagation")

    # 10. 测试生成
    with torch.no_grad():
        gen = gpt.generate(idx[:1], max_new_tokens=5, temperature=0.8)
    assert gen.shape == (1, block_size + 5), f"生成序列形状错误: {gen.shape}"
    print("[OK] Text Generation")

    print("=" * 60)
    print("所有组件测试通过!")
    print("=" * 60)


def test_fused_equivalence():
    """
    核心验证: 融合版(一个大线性层 + view 拆分)与显式版(多个独立头)
    在权重对齐的前提下,前向输出是否完全一致。

    验证方法(这也是"合并"的数学桥梁):
        融合版 c_attn 的权重矩阵 (3C, C) 按行划分:
            - 前 C 行  = 所有头的 Q 投影拼接成的大矩阵 W_q_big
            - 中间 C 行 = 所有头的 K 投影拼接成的大矩阵 W_k_big
            - 后 C 行  = 所有头的 V 投影拼接成的大矩阵 W_v_big
        其中每个头的 Q 权重就是 W_q_big 中连续的 head_size 行:
            head_h.w_q.weight = W_q_big[h*head_size : (h+1)*head_size]
        把这些行赋给显式版的各个 AttentionHead,两者输出应完全一致。
    """
    print("\n" + "=" * 60)
    print("运行融合版 vs 显式版数学等价性验证...")
    print("=" * 60)

    config = GPTConfig(
        vocab_size=50, block_size=16,
        n_layer=1, n_head=4, n_embd=32, dropout=0.0,
    )
    C, H, hs = config.n_embd, config.n_head, config.n_embd // config.n_head

    fused = FusedMultiHeadAttention(config)
    explicit = MultiHeadAttention(config)

    # 1. 对齐权重: 从融合版大矩阵中"按块切出"每个头的 Q/K/V 权重
    with torch.no_grad():
        w = fused.c_attn.weight  # (3C, C),c_attn 无 bias

        for h in range(H):
            head = explicit.heads[h]
            # 第 h 个头: 大矩阵中 [h*hs : (h+1)*hs] 行对应 Q/K/V 部分
            head.w_q.weight.data.copy_(w[0 * C + h * hs : 0 * C + (h + 1) * hs])
            head.w_k.weight.data.copy_(w[1 * C + h * hs : 1 * C + (h + 1) * hs])
            head.w_v.weight.data.copy_(w[2 * C + h * hs : 2 * C + (h + 1) * hs])

        # 对齐输出投影
        explicit.w_o.weight.data.copy_(fused.c_proj.weight)
        explicit.w_o.bias.data.copy_(fused.c_proj.bias)

    # 2. 用相同输入比较输出
    torch.manual_seed(0)
    x = torch.randn(2, 16, C)
    with torch.no_grad():
        y_fused = fused(x)
        y_explicit = explicit(x)

    # 3. 断言完全一致
    assert y_fused.shape == y_explicit.shape, "两种实现输出形状不一致"
    diff = (y_fused - y_explicit).abs().max().item()
    assert torch.allclose(y_fused, y_explicit, atol=1e-5), \
        f"融合版与显式版输出不一致! 最大差异: {diff:.2e}"

    print(f"[OK] Fused == Explicit | 输出最大差异: {diff:.2e} (阈值 1e-5)")
    print("     验证结论: 一个大线性层 + view 拆分 与 多个独立头 数学完全等价")
    print("=" * 60)


# ============================================================================
# 入口
# ============================================================================


if __name__ == "__main__":
    # 1. 运行组件测试
    test_components()

    # 2. 验证融合版与显式版数学等价
    test_fused_equivalence()

    # 3. 训练字符级 GPT 语言模型(默认显式版,可改 use_fused=True 体验融合版)
    train_char_gpt()

完成情况

已在 cucode/multi_head_atten_decode_only.py 创建了显式 Multi-Head Attention 版的 Decoder-Only Transformer,功能与 decoder_only.py 完全一致,所有 10 项单元测试 + 训练 + 生成验证通过。

核心区别:显式的多头注意力实现

本文件把多头注意力拆成三个清晰层次,与 decoder_only.py 的"一次投影+reshape"优化实现形成对比:

MultiHeadAttention
    ├── head_0 (AttentionHead)  ── w_q, w_k, w_v + 缩放点积注意力
    ├── head_1 (AttentionHead)  ── w_q, w_k, w_v + 缩放点积注意力
    ├── head_2 (AttentionHead)  ── ...
    ├── head_3 (AttentionHead)  ── ...
    └── w_o ── 拼接所有头的输出 → 输出投影
层次 类/函数 作用
1 scaled_dot_product_attention() 单头核心计算:softmax(QKᵀ/√dₖ)V
2 AttentionHead 一个头:独立的 Q/K/V 线性投影 + 缩放点积注意力
3 MultiHeadAttention 并行 n 个 AttentionHeadtorch.cat 拼接 → w_o 投影

与 decoder_only.py 的实现对比

实现方式 decoder_only.py 本文件
QKV 生成 一次线性投影生成全部 QKV 每个头独立的 w_q/w_k/w_v
多头处理 reshape + transpose 切分 显式的 nn.ModuleList 装 n 个头
输出融合 transpose + view 拼接 torch.cat + w_o 投影
教学价值 工程优化(GPU 高效) 结构清晰、逐步可见

数学上两者完全等价,但本文件的实现让你能直观看到每个头在做什么——这正是理解 Transformer 核心思想的正确入门方式。

运行结果

组件测试: 全部 10 项通过(含因果性、多头拼接维度等关键验证)
模型参数量: 103,616 | Loss: 2.214 → 0.091

生成效果:
  'hello '    -> 'hello world, this is a decoder only transforme'
  'the quick' -> 'the quick brown fox jumps over the lazy dog. gpt '
  'deep '     -> 'deep learning is fun, let us build a small gp'

建议学习路径

  1. trans.py — 完整 Encoder-Decoder Transformer(seq2seq 架构)
  2. decoder_only.py — GPT 风格优化版多头注意力(工程写法)
  3. multi_head_atten_decode_only.py — 本文件,显式多头注意力(教学写法),建议先读这里再对比 decoder_only.py,体会两种实现的等价关系

运行方式:python cucode/multi_head_atten_decode_only.py

"""
multi_head_atten_decode_only.py
================================
显式 Multi-Head Attention 版 Decoder-Only Transformer(GPT 风格)(学习用途)。

本文件实现与 decoder_only.py 完全相同的功能(GPT 架构 + 训练 + 生成),
唯一的区别在于**如何实现多头注意力**:

    | 实现方式        | decoder_only.py            | 本文件                        |
    |-----------------|----------------------------|-------------------------------|
    | QKV 生成        | 一次线性投影生成全部 QKV   | 每个头独立的 Q/K/V 线性层     |
    | 多头处理        | reshape + transpose 切分   | 显式的 n_head 个 AttentionHead|
    | 输出融合        | transpose + view 拼接      | torch.cat 拼接 + W_o 投影     |
    | 教学价值        | 工程优化,GPU 友好         | 结构清晰,逐步可见            |

两种实现数学上等价(都是"缩放点积注意力 + 多头并行"),本文件用
"每个头一个独立模块"的方式把多头的概念显式展示出来:

    MultiHeadAttention
        ├── head_0 (AttentionHead): w_q, w_k, w_v + scaled dot-product attention
        ├── head_1 (AttentionHead): w_q, w_k, w_v + scaled dot-product attention
        ├── ...                     (每个头负责关注不同方面的特征)
        └── w_o: 拼接所有头的输出后做最终线性投影

运行环境: Python 3.10+, PyTorch 2.x
依赖: pip install torch

"""

import math
import time

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader


# ============================================================================
# 第一部分:配置类
# ============================================================================


class GPTConfig:
    """GPT 模型的超参数配置(集中管理,方便修改)。"""

    def __init__(
        self,
        vocab_size: int = 128,      # 词表大小(token 种类数)
        block_size: int = 64,       # 最大上下文长度(能"看到"多少个 token)
        n_layer: int = 2,           # Transformer 块的数量
        n_head: int = 4,            # 注意力头数
        n_embd: int = 128,          # 嵌入维度(模型宽度)
        dropout: float = 0.1,       # dropout 比率
    ):
        self.vocab_size = vocab_size
        self.block_size = block_size
        self.n_layer = n_layer
        self.n_head = n_head
        self.n_embd = n_embd
        self.dropout = dropout


# ============================================================================
# 第二部分:显式 Multi-Head Attention
# ============================================================================
# 这是本文件与 decoder_only.py 的核心区别。
# 我们把"多头注意力"拆成三个层次,层层递进,方便理解:
#
#   层次 1: scaled_dot_product_attention()   —— 单头注意力的核心计算
#   层次 2: AttentionHead                    —— 一个"头"(自己的 Q/K/V 投影)
#   层次 3: MultiHeadAttention               —— 并行多个头 + 拼接 + 输出投影


def scaled_dot_product_attention(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    causal_mask: torch.Tensor,
    dropout: nn.Dropout | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
    """
    缩放点积注意力 (Scaled Dot-Product Attention)——单个头的核心计算。

    公式: Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d_k)) @ V

    参数:
        q, k, v:    (batch, seq_len, head_size)
        causal_mask: (1, 1, block_size, block_size) 下三角掩码,
                     防止当前 token 看到未来的 token
        dropout:    可选的注意力权重 dropout

    返回:
        output: 加权求和后的值 (batch, seq_len, head_size)
        attn:   注意力权重矩阵 (batch, seq_len, seq_len)
    """
    B, T, _ = q.shape
    head_size = q.size(-1)

    # 1. 计算注意力分数: Q @ K^T / sqrt(d_k)
    #    (B, T, head_size) @ (B, head_size, T) -> (B, T, T)
    scores = q @ k.transpose(-2, -1) / math.sqrt(head_size)

    # 2. 应用因果掩码: 未来位置设为 -inf,softmax 后权重趋近于 0
    #    取当前序列长度对应的 (T x T) 部分,与 (B, T, T) 的 scores 广播
    scores = scores.masked_fill(
        causal_mask[0, 0, :T, :T] == 0, float("-inf")
    )

    # 3. softmax 归一化得到注意力权重
    attn = F.softmax(scores, dim=-1)

    # 4. 可选 dropout
    if dropout is not None:
        attn = dropout(attn)

    # 5. 用注意力权重对 V 加权求和
    out = attn @ v  # (B, T, head_size)
    return out, attn


class AttentionHead(nn.Module):
    """
    单头注意力 (Attention Head)。

    每个头拥有自己独立的 Q、K、V 线性投影矩阵,可以从不同角度
    (不同子空间)观察输入序列。例如一个头关注"词性",另一个头
    关注"指代关系",各司其职。

    参数:
        n_embd:      输入维度(模型宽度)
        head_size:   该头的输出维度 d_k = n_embd / n_head
        causal_mask: 因果掩码(下三角矩阵)
        dropout:     dropout 比率
    """

    def __init__(self, n_embd: int, head_size: int, causal_mask: torch.Tensor, dropout: float = 0.1):
        super().__init__()
        self.head_size = head_size
        self.causal_mask = causal_mask
        # 每个头独立的 Q、K、V 线性投影
        self.w_q = nn.Linear(n_embd, head_size, bias=False)
        self.w_k = nn.Linear(n_embd, head_size, bias=False)
        self.w_v = nn.Linear(n_embd, head_size, bias=False)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        参数: x (batch, seq_len, n_embd)
        返回: (batch, seq_len, head_size)
        """
        # 1. 独立的 Q/K/V 投影(每个头用自己的权重矩阵)
        q = self.w_q(x)  # (B, T, head_size)
        k = self.w_k(x)
        v = self.w_v(x)

        # 2. 缩放点积注意力
        out, _ = scaled_dot_product_attention(q, k, v, self.causal_mask, self.dropout)
        return out  # (B, T, head_size)


class MultiHeadAttention(nn.Module):
    """
    多头注意力 (Multi-Head Attention)——显式实现。

    将输入同时送入 n_head 个并行的 AttentionHead,每个头在自己的子空间
    独立做注意力,然后把所有头的输出**拼接**起来,最后通过输出投影 W_o
    融合成一个 n_embd 维向量。

    参数:
        config: GPTConfig 配置
    """

    def __init__(self, config: GPTConfig):
        super().__init__()
        assert config.n_embd % config.n_head == 0, "n_embd 必须能被 n_head 整除"

        self.n_head = config.n_head
        self.n_embd = config.n_embd
        self.head_size = config.n_embd // config.n_head  # 每个头的维度 d_k

        # 因果掩码: 下三角矩阵 (1, 1, block_size, block_size)
        #   tril 结果示例 (block_size=4):
        #       [[1, 0, 0, 0],
        #        [1, 1, 0, 0],
        #        [1, 1, 1, 0],
        #        [1, 1, 1, 1]]
        # 注册为 buffer,不参与训练但会随模型移动设备
        self.register_buffer(
            "causal_mask",
            torch.tril(torch.ones(config.block_size, config.block_size)).view(
                1, 1, config.block_size, config.block_size
            ),
        )

        # n_head 个并行的注意力头(各自拥有独立的 Q/K/V 权重)
        self.heads = nn.ModuleList(
            [
                AttentionHead(config.n_embd, self.head_size, self.causal_mask, config.dropout)
                for _ in range(config.n_head)
            ]
        )

        # 输出投影 W_o: 拼接后的 n_embd 维 -> n_embd 维
        self.w_o = nn.Linear(config.n_embd, config.n_embd)
        self.resid_dropout = nn.Dropout(config.dropout)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        参数: x (batch, seq_len, n_embd)
        返回: (batch, seq_len, n_embd)
        """
        # 1. 每个头独立计算(并行执行,互不影响)
        head_outputs = [head(x) for head in self.heads]  # n_head 个 (B, T, head_size)

        # 2. 沿特征维度拼接所有头的输出
        #    (B, T, head_size) x n_head -> (B, T, n_head * head_size) = (B, T, n_embd)
        y = torch.cat(head_outputs, dim=-1)

        # 3. 输出投影融合多头信息
        return self.resid_dropout(self.w_o(y))


# ============================================================================
# 第三部分:前馈网络与 Transformer 块
# ============================================================================


class MLP(nn.Module):
    """
    前馈网络 (Feed-Forward Network),GPT 使用 GELU 激活。

        MLP(x) = Linear(n_embd -> 4*n_embd) -> GELU -> Linear(4*n_embd -> n_embd)

    为什么中间维度是 4 倍?这是 Transformer 论文中实验验证的经验值,
    相当于给每个 token 一个"独立思考"的多层感知机。
    """

    def __init__(self, config: GPTConfig):
        super().__init__()
        self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd)  # 升维
        self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd)  # 降维回 n_embd
        self.gelu = nn.GELU()  # GELU: 更平滑的 ReLU 变体
        self.dropout = nn.Dropout(config.dropout)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        x = self.c_fc(x)
        x = self.gelu(x)
        x = self.c_proj(x)
        return self.dropout(x)


class Block(nn.Module):
    """
    一个 Transformer 块:显式多头注意力 + 前馈网络。

    使用 Pre-LN(先 LayerNorm 再进子层,最后加残差),这是 GPT 系列的标准做法,
    相比原论文的 Post-LN 更容易稳定训练:
        x = x + MultiHeadAttention(LayerNorm(x))   # 子层1
        x = x + MLP(LayerNorm(x))                  # 子层2
    """

    def __init__(self, config: GPTConfig):
        super().__init__()
        self.ln_1 = nn.LayerNorm(config.n_embd)  # 注意力前的归一化
        self.attn = MultiHeadAttention(config)
        self.ln_2 = nn.LayerNorm(config.n_embd)  # 前馈前的归一化
        self.mlp = MLP(config)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        x = x + self.attn(self.ln_1(x))  # 残差 + 多头自注意力
        x = x + self.mlp(self.ln_2(x))   # 残差 + 前馈网络
        return x


# ============================================================================
# 第四部分:完整 GPT 模型
# ============================================================================


class GPT(nn.Module):
    """
    Decoder-Only Transformer 模型(GPT 风格),使用显式多头注意力。

    数据流:
        token_ids --[词嵌入]--> x
        positions --[位置嵌入]--> pos
        x = x + pos
        for each Block: x = Block(x)     # 堆叠 n_layer 个 Transformer 块
        x = LayerNorm(x)
        logits = Linear(x)               # 映射到词表概率
    """

    def __init__(self, config: GPTConfig):
        super().__init__()
        self.config = config

        # 1. 词嵌入: token id -> n_embd 维向量
        self.token_embedding = nn.Embedding(config.vocab_size, config.n_embd)
        # 2. 位置嵌入: 位置索引 -> n_embd 维向量(可学习的)
        self.position_embedding = nn.Embedding(config.block_size, config.n_embd)
        # 3. 堆叠的 Transformer 块
        self.blocks = nn.ModuleList([Block(config) for _ in range(config.n_layer)])
        # 4. 最终 LayerNorm
        self.ln_f = nn.LayerNorm(config.n_embd)
        # 5. 输出层: n_embd -> vocab_size
        self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)

        # 技巧: 权重共享 —— 让输出层复用词嵌入的权重矩阵
        # 因为"预测下一个词"和"查词向量"本质是同一张表,共享可大幅减少参数量
        self.token_embedding.weight = self.lm_head.weight

        # 参数初始化
        self.apply(self._init_weights)

        # GPT-2 的技巧: 对残差路径上的线性层(输出投影)用更小的标准差初始化
        # 防止多块堆叠后数值过大
        for name, p in self.named_parameters():
            if name.endswith("w_o.weight") or name.endswith("c_proj.weight"):
                nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layer))

    def _init_weights(self, module: nn.Module):
        """初始化权重: 线性层/嵌入层用 N(0, 0.02),偏置与归一化层置零。"""
        if isinstance(module, (nn.Linear, nn.Embedding)):
            nn.init.normal_(module.weight, mean=0.0, std=0.02)
            if isinstance(module, nn.Linear) and module.bias is not None:
                nn.init.zeros_(module.bias)
        elif isinstance(module, nn.LayerNorm):
            nn.init.zeros_(module.bias)
            nn.init.ones_(module.weight)

    def forward(
        self,
        idx: torch.Tensor,
        targets: torch.Tensor | None = None,
    ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
        """
        参数:
            idx:     输入 token id 序列 (batch, seq_len)
            targets: 目标 token id 序列 (batch, seq_len),训练时提供
                     用于计算交叉熵损失;推理时为 None

        返回:
            训练时: (logits, loss)
            推理时: (logits,)
        """
        B, T = idx.size()
        assert T <= self.config.block_size, f"序列长度 {T} 超过上下文上限 {self.config.block_size}"

        # 1. 词嵌入 (B, T, n_embd)
        tok_emb = self.token_embedding(idx)
        # 2. 位置嵌入 (T, n_embd)
        pos = torch.arange(0, T, dtype=torch.long, device=idx.device)
        pos_emb = self.position_embedding(pos)
        # 3. 相加得到输入表示
        x = tok_emb + pos_emb

        # 4. 经过所有 Transformer 块
        for block in self.blocks:
            x = block(x)

        # 5. 最终归一化 + 投影到词表
        x = self.ln_f(x)
        logits = self.lm_head(x)  # (B, T, vocab_size)

        # 6. 训练时计算损失: 目标是"下一个 token",所以 logits 在位置 t
        #    预测的是位置 t+1 的 token,与 targets 逐位比较
        loss = None
        if targets is not None:
            loss = F.cross_entropy(
                logits.view(-1, logits.size(-1)),
                targets.view(-1),
                ignore_index=-1,
            )

        return logits, loss

    @torch.no_grad()
    def generate(
        self,
        idx: torch.Tensor,
        max_new_tokens: int = 50,
        temperature: float = 1.0,
        top_k: int | None = None,
    ) -> torch.Tensor:
        """
        自回归文本生成。

        逐 token 生成: 每次把已生成的序列喂回模型,只取最后一个位置的预测,
        采样一个 token 追加到序列末尾,重复直到生成 max_new_tokens 个。

        参数:
            idx:          初始 prompt 的 token id (batch, seq_len)
            max_new_tokens: 要生成的 token 数量
            temperature:  采样温度。>1 更随机,<1 更确定,=0 取 argmax
            top_k:        只从前 k 个概率最高的 token 中采样(可选)

        返回:
            完整序列 (batch, seq_len + max_new_tokens)
        """
        for _ in range(max_new_tokens):
            # 只取最后 block_size 个 token(防止超过上下文上限)
            idx_cond = idx[:, -self.config.block_size:]

            # 前向传播(推理模式,无需梯度)
            logits, _ = self(idx_cond)          # (B, T, vocab)
            logits = logits[:, -1, :]           # 只取最后一个位置 (B, vocab)

            # 温度缩放
            if temperature != 1.0:
                logits = logits / temperature

            # Top-K 过滤: 只保留概率最高的前 k 个
            if top_k is not None:
                v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
                logits[logits < v[:, [-1]]] = float("-inf")

            # 从概率分布中采样
            probs = F.softmax(logits, dim=-1)   # (B, vocab)
            idx_next = torch.multinomial(probs, num_samples=1)  # (B, 1)

            # 追加到序列末尾
            idx = torch.cat((idx, idx_next), dim=1)

        return idx


# ============================================================================
# 第五部分:模型构建函数
# ============================================================================


def make_gpt(
    vocab_size: int,
    block_size: int = 64,
    n_layer: int = 2,
    n_head: int = 4,
    n_embd: int = 128,
    dropout: float = 0.1,
) -> GPT:
    """构建一个 GPT 模型(配置集中管理)。"""
    config = GPTConfig(
        vocab_size=vocab_size,
        block_size=block_size,
        n_layer=n_layer,
        n_head=n_head,
        n_embd=n_embd,
        dropout=dropout,
    )
    model = GPT(config)
    print(
        f"模型构建完成 | 参数量: {sum(p.numel() for p in model.parameters()):,} "
        f"| 层数: {n_layer} | 头数: {n_head} | 维度: {n_embd}"
    )
    return model


# ============================================================================
# 第六部分:字符级语言模型训练示例
# ============================================================================


class CharDataset(Dataset):
    """字符级数据集: 从文本中切分 (input, target) 训练样本。

    GPT 的训练样本是"给定前 k 个字符,预测第 k+1 个字符"。
    原始文本按块切分后,同一块内任意位置都天然构成训练样本
    (因为因果注意力会让每个位置只看到自己之前的内容),
    这里简单地对齐逐位作为 target。
    """

    def __init__(self, text: str, block_size: int = 64):
        self.chars = sorted(set(text))
        self.char2idx = {c: i for i, c in enumerate(self.chars)}
        self.idx2char = {i: c for c, i in self.char2idx.items()}
        self.vocab_size = len(self.chars)
        self.block_size = block_size

        # 编码整个文本
        self.data = torch.tensor([self.char2idx[c] for c in text], dtype=torch.long)

    def __len__(self) -> int:
        return max(0, len(self.data) - self.block_size)

    def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]:
        # x: 连续的 block_size 个字符
        x = self.data[idx : idx + self.block_size]
        # y: 右移一位,即每个位置的"下一个字符"
        y = self.data[idx + 1 : idx + self.block_size + 1]
        return x, y


def train_char_gpt(seed: int = 42):
    """
    用小型文本训练一个字符级 GPT 语言模型,演示完整训练流程。

    训练目标是: 给定前面的字符序列,预测下一个字符。
    这就是 ChatGPT 这类大模型的本质 —— 只是规模更大、数据更多。
    """
    torch.manual_seed(seed)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"训练设备: {device}")

    # --- 1. 准备数据 ---
    sample_text = (
        "hello world, this is a decoder only transformer. "
        "it learns to predict the next character. "
        "the quick brown fox jumps over the lazy dog. "
        "gpt models are decoder only, just like chatgpt. "
        "deep learning is fun, let us build a small gpt. "
    ) * 30  # 重复以增加数据量

    block_size = 32  # 上下文长度(能看到的字符数)
    dataset = CharDataset(sample_text, block_size=block_size)
    dataloader = DataLoader(dataset, batch_size=32, shuffle=True)
    print(f"词表大小: {dataset.vocab_size}, 字符: {''.join(dataset.chars)}")
    print(f"训练样本数: {len(dataset)}")

    # --- 2. 构建模型(小规模,便于 CPU 快速训练)---
    model = make_gpt(
        vocab_size=dataset.vocab_size,
        block_size=block_size,
        n_layer=2,     # 2 层 Transformer 块
        n_head=4,      # 4 个注意力头
        n_embd=64,     # 嵌入维度 64
        dropout=0.1,
    ).to(device)

    # --- 3. 优化器 ---
    optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)

    # --- 4. 训练循环 ---
    epochs = 30
    model.train()
    print("\n开始训练...")
    print("=" * 60)

    for epoch in range(epochs):
        t0 = time.time()
        total_loss = 0.0
        n_batches = 0

        for x, y in dataloader:
            x, y = x.to(device), y.to(device)

            # 前向传播 + 损失
            logits, loss = model(x, y)

            # 反向传播
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()

            total_loss += loss.item()
            n_batches += 1

        avg_loss = total_loss / max(n_batches, 1)
        if (epoch + 1) % 5 == 0 or epoch == 0:
            print(f"Epoch {epoch + 1:3d}/{epochs} | Loss: {avg_loss:.4f} | Time: {time.time() - t0:.2f}s")

    print("=" * 60)
    print("训练完成!")

    # --- 5. 文本生成(自回归采样)---
    print("\n--- 文本生成示例 ---")
    model.eval()

    prompts = ["hello ", "the quick", "deep "]
    for prompt in prompts:
        # 编码 prompt
        ids = [dataset.char2idx.get(c, 0) for c in prompt]
        idx = torch.tensor([ids], dtype=torch.long, device=device)

        # 生成 40 个新字符(温度为 0.8,稍保守)
        generated = model.generate(idx, max_new_tokens=40, temperature=0.8)

        # 解码为文本
        text = "".join(dataset.idx2char.get(i, "?") for i in generated[0].tolist())
        print(f"提示: {prompt!r} -> {text!r}")

    return model, dataset


# ============================================================================
# 第七部分:单元测试
# ============================================================================


def test_components():
    """运行断言测试,验证模型各组件的正确性。"""
    print("\n" + "=" * 60)
    print("运行组件测试...")
    print("=" * 60)

    batch, block_size = 2, 16
    config = GPTConfig(
        vocab_size=50, block_size=block_size,
        n_layer=1, n_head=2, n_embd=32, dropout=0.0,
    )

    # 1. 测试单头注意力(AttentionHead)
    causal_mask = torch.tril(torch.ones(1, 1, block_size, block_size))
    head = AttentionHead(n_embd=32, head_size=16, causal_mask=causal_mask)
    x = torch.randn(batch, block_size, 32)
    out = head(x)
    assert out.shape == (batch, block_size, 16), f"单头注意力输出形状错误: {out.shape}"
    print("[OK] AttentionHead (单头注意力)")

    # 2. 测试多头注意力(MultiHeadAttention)
    mha = MultiHeadAttention(config)
    out = mha(x)
    assert out.shape == x.shape, f"多头注意力输出形状错误: {out.shape} != {x.shape}"
    print("[OK] MultiHeadAttention (多头注意力)")

    # 3. 测试多头拼接维度(核心验证!)
    #    确认每个头输出 head_size 维,拼接后正好是 n_embd 维
    head_outputs = [head(x) for head in mha.heads]
    concat = torch.cat(head_outputs, dim=-1)
    assert concat.shape[-1] == config.n_head * mha.head_size == config.n_embd, \
        f"多头拼接维度错误: {concat.shape[-1]} != {config.n_embd}"
    print(f"[OK] Multi-Head Concatenation ({config.n_head} 个头拼接 = {concat.shape[-1]} 维)")

    # 4. 测试 MLP
    mlp = MLP(config)
    out = mlp(x)
    assert out.shape == x.shape, f"MLP 输出形状错误: {out.shape} != {x.shape}"
    print("[OK] MLP")

    # 5. 测试 Block
    block = Block(config)
    out = block(x)
    assert out.shape == x.shape, f"Block 输出形状错误: {out.shape} != {x.shape}"
    print("[OK] Transformer Block")

    # 6. 测试因果性(关键验证!)
    #    检查第 t 个位置的输出是否不受未来位置影响:
    #    改变输入最后一个 token,前 15 个位置的输出应该完全不变
    mha.eval()
    x1 = torch.randn(1, block_size, 32)
    x2 = x1.clone()
    x2[0, -1, :] = 999.0  # 大幅修改最后一个 token

    with torch.no_grad():
        o1 = mha(x1)
        o2 = mha(x2)
    assert torch.allclose(o1[0, :-1], o2[0, :-1], atol=1e-5), "因果性被破坏: 未来影响了过去"
    print("[OK] Causal Property (未来不会影响过去)")

    # 7. 测试完整 GPT 模型
    gpt = GPT(config)
    idx = torch.randint(0, 50, (batch, block_size))
    logits, loss = gpt(idx, idx)
    assert logits.shape == (batch, block_size, 50), f"模型输出形状错误: {logits.shape}"
    assert loss is not None and loss.ndim == 0, "损失计算错误"
    print("[OK] GPT Model Forward + Loss")

    # 8. 测试推理模式(无 targets)
    logits, loss = gpt(idx)
    assert loss is None, "推理模式下 loss 应为 None"
    print("[OK] GPT Inference Mode (loss=None)")

    # 9. 测试梯度反向传播
    loss = logits.sum()
    loss.backward()
    grad_ok = all(p.grad is not None for p in gpt.parameters() if p.requires_grad)
    assert grad_ok, "梯度反向传播失败"
    print("[OK] Backpropagation")

    # 10. 测试生成
    with torch.no_grad():
        gen = gpt.generate(idx[:1], max_new_tokens=5, temperature=0.8)
    assert gen.shape == (1, block_size + 5), f"生成序列形状错误: {gen.shape}"
    print("[OK] Text Generation")

    print("=" * 60)
    print("所有组件测试通过!")
    print("=" * 60)


# ============================================================================
# 入口
# ============================================================================


if __name__ == "__main__":
    # 1. 运行组件测试
    test_components()

    # 2. 训练字符级 GPT 语言模型
    train_char_gpt()

多头注意力分析

已经拿到完整中间结果(为了可复现,诊断脚本固定了随机种子,数值与直接运行略有不同,但流程完全一致)。下面以测试 2 为例(batch=1, seq=2, d_model=4, num_heads=2, d_k=2),从输入到输出逐步给出每个阶段的具体张量。

第 0 步:输入

  • QK(1, 2, 4),两个 token,每 token 4 维:
token0: [1, 0, 0, 1]
token1: [0, 1, 1, 0]
  • V(1, 2, 4)
token0: [1, 2, 3, 4]
token1: [5, 6, 7, 8]
  • mask(1, 2, 2),全 1(不遮蔽)。

第 1 步:W_q / W_k / W_v 线性投影(形状不变 (1,2,4)→(1,2,4)

  • 每个 Linear 的权重是 (4, 4),对每个 token 做 y = xWᵀ + b,把 4 维向量重新混合。以 Q 为例:
W_q(Q):
token0: [-0.1740,  0.5039, -0.4816, -0.6769]
token1: [ 0.0544,  0.4242, -0.3578, -0.7550]
  • V 投影后:
W_v(V):
token0: [-2.0068,  0.8546,  2.3432, -0.0966]
token1: [-4.4560,  0.8129,  4.6515,  1.1307]

第 2 步:.view(1, -1, 2, 2) 切头(形状 (1,2,4)→(1,2,2,2),语义 (batch, seq, head, d_k)

  • 把每个 token 最后的 4 维从中间切开:前 2 维归头 0,后 2 维归头 1。以 Q 为例:
view 后 (seq, head, d_k):
token0: 头0=[-0.1740,  0.5039]   头1=[-0.4816, -0.6769]
token1: 头0=[ 0.0544,  0.4242]   头1=[-0.3578, -0.7550]
  • V 同理:
token0: 头0=[-2.0068,  0.8546]   头1=[ 2.3432, -0.0966]
token1: 头0=[-4.4560,  0.8129]   头1=[ 4.6515,  1.1307]

第 3 步:.transpose(1, 2) 归拢各头(形状仍 (1,2,2,2),语义变为 (batch, head, seq, d_k)

  • 只是重新排列维度,把"同一头的所有 token"放到一起,让每个头能独立做注意力。Q 变成:
头0:
  token0: [-0.1740,  0.5039]
  token1: [ 0.0544,  0.4242]
头1:
  token0: [-0.4816, -0.6769]
  token1: [-0.3578, -0.7550]
  • K、V 同样重排(下面会用到的 V):
头0: token0=[-2.0068, 0.8546]  token1=[-4.4560, 0.8129]
头1: token0=[ 2.3432, -0.0966] token1=[ 4.6515, 1.1307]

第 4 步:scores = Q·Kᵀ(1,2,2,2)×(1,2,2,2)→(1,2,2,2),语义 (batch, head, seq_q, seq_k)

  • 对每个头独立做 (2×2)·(2×2) 矩阵乘。以头 0 的 score[0][0](token0 对 token0)为例:
    -0.1740×0.2639 + 0.5039×(-0.1146) = -0.1037
  • 完整分数:
头0: [[-0.1037,  0.4678],
     [-0.0343,  0.3979]]
头1: [[ 0.2084,  0.2756],
     [ 0.2935,  0.2897]]
  • 每行代表"某个 query 对两个 key 的相似度"。

第 5 步:/ math.sqrt(2) 缩放(形状不变)

  • 防止点积过大导致 softmax 梯度消失。结果:
头0: [[-0.0733,  0.3308],
     [-0.0242,  0.2814]]
头1: [[ 0.1474,  0.1949],
     [ 0.2075,  0.2049]]

第 6 步:mask 处理(形状不变)

  • mask.unsqueeze(1)(1,2,2) 变成 (1,1,2,2),广播到 4 维后与 scores 对齐。
  • masked_fill(mask==0, -inf):本测试 mask 全 1,没有位置被遮蔽,数值不变。如果某位置是 0,会被填成 -inf,softmax 后概率为 0。

第 7 步:F.softmax(..., dim=-1) 得到注意力权重((1,2,2,2)

  • 对最后一维(每个 query 对所有 key)归一化,每行和为 1。以头 0 第一行为例:
    softmax([-0.0733, 0.3308]) = [0.4003, 0.5997]
  • 完整权重:
头0: [[0.4003, 0.5997],
     [0.4242, 0.5758]]
头1: [[0.4881, 0.5119],
     [0.5007, 0.4993]]
  • 这就是"每个 token 应该用多大比例去看其他 token"。

第 8 步:attention_weights · V(1,2,2,2)→(1,2,2,2),仍是 (batch, head, seq, d_k)

  • 每个头用权重对 V 加权求和。以头 0 token0 为例:
    0.4003×[-2.0068, 0.8546] + 0.5997×[-4.4560, 0.8129] = [-3.4755, 0.8296]
  • 完整结果:
头0: token0=[-3.4755,  0.8296]   token1=[-3.4171,  0.8306]
头1: token0=[ 3.5248,  0.5316]   token1=[ 3.4958,  0.5162]

第 9 步:.transpose(1, 2) 换回((batch, head, seq, d_k)→(batch, seq, head, d_k)

  • 把"头"维度挪回第 2 维,每个 token 的两个头结果重新归到一起:
token0: 头0=[-3.4755, 0.8296]  头1=[3.5248, 0.5316]
token1: 头0=[-3.4171, 0.8306]  头1=[3.4958, 0.5162]
  • 此时内存是非连续的,所以紧跟 .contiguous() 复制成连续内存,为下一步 view 做准备。

第 10 步:.view(1, -1, 4) 合并头((1,2,2,2)→(1,2,4)

  • 把每个 token 的 (头0 2维, 头1 2维) 首尾拼接回 4 维,语义 (batch, seq, d_model)
token0: [-3.4755,  0.8296,  3.5248,  0.5316]
token1: [-3.4171,  0.8306,  3.4958,  0.5162]

第 11 步:W_o 输出投影((1,2,4)→(1,2,4)

  • 用另一个 Linear(4,4) 混合两个头的信息,得到最终输出:
最终输出:
token0: [-1.4836, -1.3123, -2.0893,  1.1161]
token1: [-1.4640, -1.2986, -2.0568,  1.1086]

最终返回

  • output(1, 2, 4),和输入形状一致,可直接喂给下一层。
  • attention_weights(1, 2, 2, 2),即 (batch, heads, seq_q, seq_k),保留每个头各自的注意力分布,方便可视化或做分析。

形状变化总表

  • 输入 (1,2,4) → 投影 (1,2,4) → view (1,2,2,2) → transpose (1,2,2,2) → 注意力计算 (1,2,2,2)(权重同形状)→ transpose+contiguous (1,2,2,2) → view (1,2,4) → W_o (1,2,4)
import torch
import torch.nn as nn
import torch.nn.functional as F
import math

def scaled_dot_product_attention(Q, K, V, mask=None):
    d_k = Q.size(-1)
    scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)

    if mask is not None:
        scores = scores.masked_fill(mask == 0, float('-inf'))

    attention_weights = F.softmax(scores, dim=-1)
    output = torch.matmul(attention_weights, V)

    return output, attention_weights

def test_scaled_dot_product_attention():
    Q = torch.tensor([[[1.0, 0.0], [0.0, 1.0]]])
    K = torch.tensor([[[1.0, 0.0], [0.0, 1.0]]])
    V = torch.tensor([[[1.0, 2.0], [3.0, 4.0]]])
    mask = torch.tensor([[[1, 1], [1, 1]]])

    output, attention_weights = scaled_dot_product_attention(Q, K, V, mask)
    print("Output:", output)
    print("Attention Weights:", attention_weights)

    print("Output shape:", output.shape)
    print("Attention Weights shape:", attention_weights.shape)

class OneHeadAttention(nn.Module):
    def __init__(self, d_model):
        super(OneHeadAttention, self).__init__()
        self.W_q = nn.Linear(d_model, d_model)
        self.W_k = nn.Linear(d_model, d_model)
        self.W_v = nn.Linear(d_model, d_model)

    def forward(self, Q, K, V, mask=None):
        Q = self.W_q(Q)
        K = self.W_k(K)
        V = self.W_v(V)

        output, attention_weights = scaled_dot_product_attention(Q, K, V, mask)
        return output, attention_weights
    
class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, num_heads):
        super(MultiHeadAttention, self).__init__()
        assert d_model % num_heads == 0

        self.d_k = d_model // num_heads
        self.num_heads = num_heads

        self.W_q = nn.Linear(d_model, d_model)
        self.W_k = nn.Linear(d_model, d_model)
        self.W_v = nn.Linear(d_model, d_model)
        self.W_o = nn.Linear(d_model, d_model)

    def forward(self, Q, K, V, mask=None):
        batch_size = Q.size(0)

        Q = self.W_q(Q).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
        K = self.W_k(K).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
        V = self.W_v(V).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)

        if mask is not None:
            mask = mask.unsqueeze(1)

        output, attention_weights = scaled_dot_product_attention(Q, K, V, mask)

        output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.num_heads * self.d_k)
        output = self.W_o(output)
        return output, attention_weights

def test_multi_head_attention():
    d_model = 4
    num_heads = 2
    mha = MultiHeadAttention(d_model, num_heads)

    Q = torch.tensor([[[1.0, 0.0, 0.0, 1.0], [0.0, 1.0, 1.0, 0.0]]])
    K = torch.tensor([[[1.0, 0.0, 0.0, 1.0], [0.0, 1.0, 1.0, 0.0]]])
    V = torch.tensor([[[1.0, 2.0, 3.0, 4.0], [5.0, 6.0, 7.0, 8.0]]])
    mask = torch.tensor([[[1, 1], [1, 1]]])

    output, attention_weights = mha(Q, K, V, mask)
    print("MultiHead Output:", output)
    print("MultiHead Attention Weights:", attention_weights)

    print("MultiHead Output shape:", output.shape)
    print("MultiHead Attention Weights shape:", attention_weights.shape)

if __name__ == "__main__":
    test_scaled_dot_product_attention()
    test_multi_head_attention()
# -*- coding: utf-8 -*-
"""debug_attention.py
逐步打印 MultiHeadAttention 前向传播中每个阶段的张量形状与数值,
方便理解 view / transpose / 注意力计算 / 合并头等每一步的变化。

用法:
    python debug_attention.py        # 运行内置示例(与 trans.py 测试 2 相同输入)
    或在代码中调用 debug_attention(Q, K, V, mask, d_model, num_heads)
"""
import math

import torch
import torch.nn as nn
import torch.nn.functional as F


def _show(name, tensor):
    """打印张量名称、形状和具体数值。"""
    print(f"[{name}] shape={tuple(tensor.shape)}")
    print(tensor)


def debug_attention(Q, K, V, mask=None, d_model=4, num_heads=2, seed=0):
    """按步骤打印多头注意力的中间结果。

    参数:
        Q, K, V  : 形状为 (batch, seq, d_model) 的输入张量
        mask     : 可选,形状为 (batch, seq, seq),0 表示遮蔽
        d_model  : 模型维度
        num_heads: 注意力头数
        seed     : 随机种子,固定后结果可复现
    返回:
        (output, attention_weights)
    """
    assert d_model % num_heads == 0
    d_k = d_model // num_heads
    batch_size = Q.size(0)

    # 固定随机种子,保证线性层初始化一致、结果可复现
    torch.manual_seed(seed)

    # 四个投影层,与 trans.py 中 MultiHeadAttention 的配置一致
    W_q = nn.Linear(d_model, d_model)
    W_k = nn.Linear(d_model, d_model)
    W_v = nn.Linear(d_model, d_model)
    W_o = nn.Linear(d_model, d_model)

    print(f"=== 配置: d_model={d_model}, num_heads={num_heads}, d_k={d_k}, batch={batch_size}")
    print(f"=== W_q.weight 形状: {tuple(W_q.weight.shape)}(按行切成 {num_heads} 块,每块就是每个头的 d_model->d_k 投影)")    # ---- 第 0 步:输入 ----
    _show("第0步 输入 Q", Q)
    _show("第0步 输入 K", K)
    _show("第0步 输入 V", V)
    if mask is not None:
        _show("第0步 输入 mask", mask)

    # ---- 第 1 步:线性投影 (batch, seq, d_model) -> (batch, seq, d_model) ----
    Qp = W_q(Q)
    Kp = W_k(K)
    Vp = W_v(V)
    _show("第1步 W_q(Q) 线性投影后", Qp)
    _show("第1步 W_k(K) 线性投影后", Kp)
    _show("第1步 W_v(V) 线性投影后", Vp)

    # ---- 第 2 步:view 切头 (batch, seq, d_model) -> (batch, seq, num_heads, d_k) ----
    # 每个 token 的 d_model 维向量被均分成 num_heads 份,每份 d_k 维
    Qv = Qp.view(batch_size, -1, num_heads, d_k)
    Kv = Kp.view(batch_size, -1, num_heads, d_k)
    Vv = Vp.view(batch_size, -1, num_heads, d_k)
    _show("第2步 Q view 后 (batch, seq, head, d_k)", Qv)
    _show("第2步 V view 后 (batch, seq, head, d_k)", Vv)

    # ---- 第 3 步:transpose(1,2) (batch, seq, head, d_k) -> (batch, head, seq, d_k) ----
    # 把"头"维度挪到序列长度之前,让每个头独立做注意力
    Qt = Qv.transpose(1, 2)
    Kt = Kv.transpose(1, 2)
    Vt = Vv.transpose(1, 2)
    _show("第3步 Q transpose 后 (batch, head, seq, d_k)", Qt)
    _show("第3步 K transpose 后 (batch, head, seq, d_k)", Kt)
    _show("第3步 V transpose 后 (batch, head, seq, d_k)", Vt)

    # ---- 第 4 步:scores = Q*K^T (batch, head, seq_q, seq_k) ----
    scores = torch.matmul(Qt, Kt.transpose(-2, -1))
    _show("第4步 scores = Q*K^T(未缩放)", scores)

    # ---- 第 5 步:除以 sqrt(d_k) 缩放,防止点积过大 ----
    scores_scaled = scores / math.sqrt(d_k)
    _show("第5步 scores / sqrt(d_k)", scores_scaled)

    # ---- 第 6 步:mask 处理 ----
    if mask is not None:
        mask4 = mask.unsqueeze(1)  # (batch, seq, seq) -> (batch, 1, seq, seq),广播到所有头
        _show("第6步 mask unsqueeze(1) 后 (batch, 1, seq, seq)", mask4)
        scores_scaled = scores_scaled.masked_fill(mask4 == 0, float("-inf"))
        _show("第6步 mask 后 scores(-inf 处 softmax 后为 0)", scores_scaled)    # ---- 第 7 步:softmax(dim=-1) 得到注意力权重,每行和为 1 ----
    attention_weights = F.softmax(scores_scaled, dim=-1)
    _show("第7步 attention_weights (batch, head, seq_q, seq_k)", attention_weights)

    # ---- 第 8 步:权重对 V 加权求和 (batch, head, seq, d_k) ----
    attn_out = torch.matmul(attention_weights, Vt)
    _show("第8步 attention 输出 (batch, head, seq, d_k)", attn_out)

    # ---- 第 9 步:transpose(1,2) 换回 (batch, seq, head, d_k),再 contiguous 保证内存连续 ----
    merged = attn_out.transpose(1, 2).contiguous()
    _show("第9步 transpose 回 (batch, seq, head, d_k) + contiguous", merged)

    # ---- 第 10 步:view 合并头 (batch, seq, head*d_k) = (batch, seq, d_model) ----
    merged2 = merged.view(batch_size, -1, num_heads * d_k)
    _show("第10步 view 合并回 (batch, seq, d_model)", merged2)

    # ---- 第 11 步:W_o 输出投影,混合各头信息 (batch, seq, d_model) ----
    output = W_o(merged2)
    _show("第11步 W_o 输出投影(最终输出)", output)

    print(f"=== 最终: output 形状 {tuple(output.shape)}, attention_weights 形状 {tuple(attention_weights.shape)}")
    return output, attention_weights


def main():
    # 与 trans.py 中 test_multi_head_attention 相同的输入
    Q = torch.tensor([[[1.0, 0.0, 0.0, 1.0], [0.0, 1.0, 1.0, 0.0]]])
    K = torch.tensor([[[1.0, 0.0, 0.0, 1.0], [0.0, 1.0, 1.0, 0.0]]])
    V = torch.tensor([[[1.0, 2.0, 3.0, 4.0], [5.0, 6.0, 7.0, 8.0]]])
    mask = torch.tensor([[[1, 1], [1, 1]]])
    debug_attention(Q, K, V, mask, d_model=4, num_heads=2, seed=0)


if __name__ == "__main__":
    main()

更多推荐