week5

目标:

训练基于transformer的单向语言模型,并完成文本生成。

示意图:

内容:

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

# ====================================
# 1. 模型超参数
# ====================================
torch.manual_seed(42)

batch_size = 32
block_size = 64        # 最大上下文长度
n_embd = 128           # 词向量维度
n_head = 4             # 注意力头数
n_layer = 3            # Transformer 层数
dropout = 0.1
learning_rate = 3e-4
max_iters = 1500
eval_interval = 300

device = (
    "cuda" if torch.cuda.is_available()
    else "mps" if torch.backends.mps.is_available()
    else "cpu"
)

print("当前设备:", device)


# ====================================
# 2. 准备训练数据
# ====================================
# 建议准备自己的中文语料 data.txt。
# 没有文件时使用演示语料,仅用于验证训练流程。
if os.path.exists("data.txt"):
    with open("data.txt", "r", encoding="utf-8") as f:
        text = f.read()
else:
    sentences = [
        "今天天气很好,我们一起出去散步。",
        "人工智能正在改变我们的生活。",
        "深度学习是机器学习的重要分支。",
        "自然语言处理可以帮助计算机理解文本。",
        "神经网络可以通过训练学习数据中的规律。",
        "Transformer使用注意力机制处理序列信息。",
        "语言模型的任务是预测下一个字符。",
        "机器学习需要大量数据进行训练。",
        "学习编程需要不断地练习和思考。",
        "使用PyTorch可以方便地构建深度学习模型。",
        "我们正在学习如何训练一个语言模型。",
        "文本生成是自然语言处理的重要应用。",
    ]
    text = "\n".join(sentences * 200)

# 字符级 Tokenizer
chars = sorted(list(set(text)))
vocab_size = len(chars)

stoi = {ch: i for i, ch in enumerate(chars)}
itos = {i: ch for ch, i in stoi.items()}

def encode(s):
    return [stoi[c] for c in s]

def decode(ids):
    return "".join(itos[int(i)] for i in ids)

data = torch.tensor(encode(text), dtype=torch.long)

# 训练集与验证集
n = int(0.9 * len(data))
train_data = data[:n]
val_data = data[n:]

assert min(len(train_data), len(val_data)) > block_size

print("词表大小:", vocab_size)
print("训练字符数:", len(train_data))


# ====================================
# 3. 构建训练批次
# ====================================
def get_batch(split):
    source = train_data if split == "train" else val_data

    ix = torch.randint(
        0,
        len(source) - block_size,
        (batch_size,)
    )

    x = torch.stack([
        source[i:i + block_size] for i in ix
    ])

    # y 相对于 x 向右移动一个位置
    y = torch.stack([
        source[i + 1:i + block_size + 1] for i in ix
    ])

    return x.to(device), y.to(device)


# ====================================
# 4. 多头因果自注意力
# ====================================
class CausalSelfAttention(nn.Module):
    def __init__(self):
        super().__init__()

        assert n_embd % n_head == 0

        self.num_heads = n_head
        self.head_dim = n_embd // n_head

        self.qkv = nn.Linear(n_embd, 3 * n_embd)
        self.proj = nn.Linear(n_embd, n_embd)

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

        # 下三角因果 Mask
        self.register_buffer(
            "mask",
            torch.tril(
                torch.ones(block_size, block_size)
            ).view(1, 1, block_size, block_size)
        )

    def forward(self, x):
        B, T, C = x.shape

        # 一次线性映射生成 Q、K、V
        q, k, v = self.qkv(x).split(C, dim=2)

        # [B, T, C] -> [B, Head, T, HeadDim]
        q = q.view(B, T, self.num_heads,
                   self.head_dim).transpose(1, 2)
        k = k.view(B, T, self.num_heads,
                   self.head_dim).transpose(1, 2)
        v = v.view(B, T, self.num_heads,
                   self.head_dim).transpose(1, 2)

        # Scaled Dot-Product Attention
        scores = q @ k.transpose(-2, -1)
        scores = scores / math.sqrt(self.head_dim)

        # 遮挡未来 Token
        scores = scores.masked_fill(
            self.mask[:, :, :T, :T] == 0,
            float("-inf")
        )

        weights = F.softmax(scores, dim=-1)
        weights = self.attn_dropout(weights)

        out = weights @ v

        # 拼接所有注意力头
        out = out.transpose(1, 2).contiguous()
        out = out.view(B, T, C)

        return self.resid_dropout(self.proj(out))


# ====================================
# 5. 前馈神经网络
# ====================================
class FeedForward(nn.Module):
    def __init__(self):
        super().__init__()

        self.net = nn.Sequential(
            nn.Linear(n_embd, 4 * n_embd),
            nn.GELU(),
            nn.Linear(4 * n_embd, n_embd),
            nn.Dropout(dropout)
        )

    def forward(self, x):
        return self.net(x)


# ====================================
# 6. Transformer Block
# ====================================
class TransformerBlock(nn.Module):
    def __init__(self):
        super().__init__()

        self.ln1 = nn.LayerNorm(n_embd)
        self.attn = CausalSelfAttention()

        self.ln2 = nn.LayerNorm(n_embd)
        self.ffn = FeedForward()

    def forward(self, x):
        # Pre-Norm + 残差连接
        x = x + self.attn(self.ln1(x))
        x = x + self.ffn(self.ln2(x))
        return x


# ====================================
# 7. GPT 单向语言模型
# ====================================
class MiniGPT(nn.Module):
    def __init__(self):
        super().__init__()

        self.token_embedding = nn.Embedding(
            vocab_size, n_embd
        )

        self.position_embedding = nn.Embedding(
            block_size, n_embd
        )

        self.blocks = nn.Sequential(
            *[TransformerBlock() for _ in range(n_layer)]
        )

        self.ln_f = nn.LayerNorm(n_embd)
        self.lm_head = nn.Linear(n_embd, vocab_size)

    def forward(self, idx, targets=None):
        B, T = idx.shape

        assert T <= block_size

        token_emb = self.token_embedding(idx)

        positions = torch.arange(
            T, device=idx.device
        )
        pos_emb = self.position_embedding(positions)

        x = token_emb + pos_emb
        x = self.blocks(x)
        x = self.ln_f(x)

        logits = self.lm_head(x)

        loss = None

        if targets is not None:
            B, T, V = logits.shape

            loss = F.cross_entropy(
                logits.reshape(B * T, V),
                targets.reshape(B * T)
            )

        return logits, loss

    @torch.no_grad()
    def generate(
        self,
        idx,
        max_new_tokens=100,
        temperature=1.0,
        top_k=None
    ):
        assert temperature > 0

        self.eval()

        for _ in range(max_new_tokens):
            # 只使用最后 block_size 个 Token
            idx_cond = idx[:, -block_size:]

            logits, _ = self(idx_cond)

            # 获取最后一个位置的预测
            logits = logits[:, -1, :]

            # 温度调节
            logits = logits / temperature

            # Top-K 采样
            if top_k is not None:
                k = min(top_k, logits.size(-1))
                values, _ = torch.topk(logits, k)

                logits[logits < values[:, [-1]]] = -float(
                    "inf"
                )

            probs = F.softmax(logits, dim=-1)

            next_token = torch.multinomial(
                probs,
                num_samples=1
            )

            idx = torch.cat(
                [idx, next_token],
                dim=1
            )

        return idx


# ====================================
# 8. 训练与验证
# ====================================
model = MiniGPT().to(device)

print(
    "模型参数量:",
    sum(p.numel() for p in model.parameters())
)

optimizer = torch.optim.AdamW(
    model.parameters(),
    lr=learning_rate
)

@torch.no_grad()
def estimate_loss():
    model.eval()
    results = {}

    for split in ["train", "val"]:
        losses = torch.zeros(20)

        for i in range(20):
            xb, yb = get_batch(split)
            _, loss = model(xb, yb)
            losses[i] = loss.item()

        results[split] = losses.mean().item()

    model.train()
    return results


for step in range(max_iters):
    if step % eval_interval == 0 or step == max_iters - 1:
        losses = estimate_loss()

        print(
            f"step {step:4d} | "
            f"train loss: {losses['train']:.4f} | "
            f"val loss: {losses['val']:.4f}"
        )

    xb, yb = get_batch("train")

    logits, loss = model(xb, yb)

    optimizer.zero_grad(set_to_none=True)
    loss.backward()

    # 梯度裁剪
    torch.nn.utils.clip_grad_norm_(
        model.parameters(), 1.0
    )

    optimizer.step()


# ====================================
# 9. 保存模型
# ====================================
torch.save({
    "model_state_dict": model.state_dict(),
    "stoi": stoi,
    "itos": itos,
    "block_size": block_size,
    "n_embd": n_embd,
    "n_head": n_head,
    "n_layer": n_layer,
}, "mini_gpt.pth")

print("模型已保存:mini_gpt.pth")


# ====================================
# 10. 文本生成
# ====================================
prompt = "人工智能"

# 字符级词表无法识别训练集外的字符
unknown = set(prompt) - set(stoi)
if unknown:
    raise ValueError(f"提示词包含词表外字符: {unknown}")

context = torch.tensor(
    [encode(prompt)],
    dtype=torch.long,
    device=device
)

generated = model.generate(
    context,
    max_new_tokens=100,
    temperature=0.8,
    top_k=10
)

print("\n========== 生成结果 ==========")
print(decode(generated[0].tolist()))
相关推荐
孙启超2 小时前
【FDE开发指南】第 1 课:认识 FDE —— 从一次生产事故说起
人工智能·ai·职场技能
燐妤3 小时前
LangGraph-复习总览
python·ai·面试·agent·学习方法·langgraph
VIP_CQCRE3 小时前
让 Claude 实时联网搜索:Ace Data Cloud Serp MCP 接入指南
ai·claude·搜索·mcp·acedatacloud
西索斯coding3 小时前
grok-4.7 调用一直 429 怎么办?不是额度用完——是 RPM 和 TPM 双桶限流,附响应头读取代码和退避策略
大数据·人工智能·机器学习·ai
东方芷兰4 小时前
Agent 技术摘要 06 —— Harness、原生视觉、Jev、mmproj、dsh
人工智能·笔记·python·ai·langchain·ai编程
小贺儿开发4 小时前
Unity 文物新生 会讲故事的文物展墙
人工智能·科技·unity·ai·视频·互动·演示
最强小杰5 小时前
claude-opus-5.5 API 接入教程:anthropic-version 头、tools 定义、流式事件格式三个常见问题全部梳理,收藏备用
java·服务器·网络·ai
进击的雷神5 小时前
论文里的架构图是张死 PNG?Edit-Banana 把它变回可编辑的 DrawIO
ai·开源·drawio
zhanghaha13146 小时前
AI Agent_7 AI 提示词工程(Prompt Engineering)
ai