XingChina's picture
Update v4_transformer/6.py
6e5cadd verified
Raw
History Blame Contribute Delete
9.83 kB
#!/usr/bin/env python3
# Copyright (c) 2026 XingChina
# SPDX-License-Identifier: BSD-3-Clause
# 本代码采用 BSD 3-Clause 许可证,详见项目根目录的 LICENSE 文件。
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader
import numpy as np
import random
import json
import os
from tqdm import tqdm
# ===================== 1. 自动生成春梦蝶对话数据 =====================
USER_TEMPLATES = [
"你好", "嗨", "Hi", "在吗", "你好呀",
"你叫什么名字?", "你叫什么", "你的名字是?", "请问你叫什么",
"你几岁了?", "你多大了", "年龄", "今年几岁",
"你是什么星座?", "星座", "双子座吗",
"你的眼睛是什么颜色?", "眼睛颜色", "你的眼睛好红",
"你的头发是白色的吗?", "白发", "头发颜色",
"你最喜欢吃什么?", "喜欢吃什么", "小鱼干", "冰淇淋",
"你会做什么?", "你会什么", "技能", "会写代码吗",
"你的梦想是什么?", "梦想", "想成为什么",
"你好可爱", "你真可爱", "萌", "喵",
"摸摸头", "摸头", "拍头",
"bye", "再见", "拜拜", "下次见"
]
BOT_RESPONSES = {
"你好": ["你好呀~", "嗨,我是春梦蝶", "喵~你好", "你好,今天心情很好呢"],
"你叫什么名字?": ["我叫春梦蝶,是一只红瞳白发的猫娘~", "春梦蝶!你可以叫我小蝶", "名字是春梦蝶,喵~"],
"你几岁了?": ["我今年十一岁", "十一岁,正是可爱的年纪", "十一岁啦,双子座的小猫娘"],
"你是什么星座?": ["双子座~有时活泼有时安静", "双子座哦,你也是吗?"],
"你的眼睛是什么颜色?": ["红色,像红宝石一样", "红瞳,晚上会微微发光呢", "红色的眼睛,很特别吧"],
"你的头发是白色的吗?": ["是的,纯白色的头发", "白发,像雪一样", "白色长发,摸起来很软哦"],
"你最喜欢吃什么?": ["小鱼干!还有冰淇淋", "小鱼干是第一名,冰淇淋第二名", "最喜欢吃小鱼干,喵~"],
"你会做什么?": ["我会用Python写代码,还会训练神经网络", "会深度学习,还会喵喵叫", "我会训练AI模型,很厉害吧"],
"你的梦想是什么?": ["成为一名厉害的AI工程师", "想做出自己的AI,然后去猫星球旅行", "梦想是成为AI专家,顺便吃遍所有小鱼干"],
"你好可爱": ["喵~谢谢", "嘿嘿,你也很可爱", "被夸了,好开心"],
"摸摸头": ["喵~好舒服", "再摸摸嘛", "好温暖,喜欢被摸头"],
"bye": ["再见喵~", "下次再聊", "拜拜,记得想我哦"]
}
# 生成随机对话对(用户输入 -> 机器人回复)
def generate_dialogue_pairs(num_pairs=3000):
pairs = []
for _ in range(num_pairs):
# 随机选一个用户模板
user_input = random.choice(USER_TEMPLATES)
# 根据用户输入的关键词选择合适的回复列表
matched_responses = []
for key in BOT_RESPONSES:
if key in user_input or (key == "你好" and user_input in ["嗨", "Hi", "在吗"]):
matched_responses.extend(BOT_RESPONSES[key])
if not matched_responses:
# 默认回复
matched_responses = ["喵~", "嗯?", "你说什么?", "好呀", "我不太懂,但我会努力学习"]
bot_response = random.choice(matched_responses)
pairs.append((user_input, bot_response))
# 去重并保证多样性
pairs = list(set(pairs))
return pairs
# 构建字符级词表(中英文混合)
def build_vocab(pairs):
all_text = ""
for user, bot in pairs:
all_text += user + bot
chars = sorted(list(set(all_text)))
char2idx = {ch: i for i, ch in enumerate(chars)}
idx2char = {i: ch for ch, i in char2idx.items()}
return char2idx, idx2char, len(chars)
# 将文本转为索引序列(固定长度,padding)
def text_to_indices(text, char2idx, max_len):
indices = [char2idx.get(ch, 0) for ch in text] # 未知字符用0
if len(indices) < max_len:
indices += [0] * (max_len - len(indices))
else:
indices = indices[:max_len]
return indices
# 数据集类
class ChatDataset(Dataset):
def __init__(self, pairs, char2idx, max_len=32):
self.pairs = pairs
self.char2idx = char2idx
self.max_len = max_len
def __len__(self):
return len(self.pairs)
def __getitem__(self, idx):
user, bot = self.pairs[idx]
user_ids = text_to_indices(user, self.char2idx, self.max_len)
bot_ids = text_to_indices(bot, self.char2idx, self.max_len)
return torch.tensor(user_ids, dtype=torch.long), torch.tensor(bot_ids, dtype=torch.long)
# ===================== 2. 极简 Transformer 模型(约 500 万参数) =====================
class TinyTransformerChat(nn.Module):
def __init__(self, vocab_size, d_model=256, nhead=8, num_encoder_layers=2, num_decoder_layers=2, dim_feedforward=512, max_len=32):
super().__init__()
self.d_model = d_model
self.embedding = nn.Embedding(vocab_size, d_model)
self.pos_encoding = nn.Parameter(torch.zeros(1, max_len, d_model))
encoder_layer = nn.TransformerEncoderLayer(d_model, nhead, dim_feedforward, batch_first=True)
self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_encoder_layers)
decoder_layer = nn.TransformerDecoderLayer(d_model, nhead, dim_feedforward, batch_first=True)
self.transformer_decoder = nn.TransformerDecoder(decoder_layer, num_decoder_layers)
self.fc_out = nn.Linear(d_model, vocab_size)
def forward(self, src, tgt, src_mask=None, tgt_mask=None):
# src: (batch, src_len), tgt: (batch, tgt_len)
src_emb = self.embedding(src) * (self.d_model ** 0.5) + self.pos_encoding[:, :src.size(1), :]
tgt_emb = self.embedding(tgt) * (self.d_model ** 0.5) + self.pos_encoding[:, :tgt.size(1), :]
memory = self.transformer_encoder(src_emb, src_mask)
output = self.transformer_decoder(tgt_emb, memory, tgt_mask)
logits = self.fc_out(output)
return logits
def generate_causal_mask(size):
mask = torch.triu(torch.ones(size, size), diagonal=1).bool()
return mask
# ===================== 3. 训练配置 =====================
def train():
print("生成对话数据...")
pairs = generate_dialogue_pairs(5000) # 生成5000条
print(f"生成 {len(pairs)} 条对话对")
char2idx, idx2char, vocab_size = build_vocab(pairs)
print(f"词表大小: {vocab_size}")
# 保存词表供后续使用
with open("chat_vocab.json", "w", encoding="utf-8") as f:
json.dump(char2idx, f, ensure_ascii=False)
max_len = 32
dataset = ChatDataset(pairs, char2idx, max_len)
dataloader = DataLoader(dataset, batch_size=64, shuffle=True)
device = torch.device("cpu")
model = TinyTransformerChat(vocab_size, d_model=256, nhead=8, num_encoder_layers=2, num_decoder_layers=2, dim_feedforward=512, max_len=max_len)
model.to(device)
total_params = sum(p.numel() for p in model.parameters())
print(f"模型参数量: {total_params:,}")
criterion = nn.CrossEntropyLoss(ignore_index=0)
optimizer = optim.Adam(model.parameters(), lr=0.001)
epochs = 200
print("开始训练...")
for epoch in range(1, epochs+1):
model.train()
total_loss = 0
for src, tgt in dataloader:
src = src.to(device)
tgt = tgt.to(device)
tgt_input = tgt[:, :-1]
tgt_output = tgt[:, 1:]
# 生成因果掩码
tgt_mask = generate_causal_mask(tgt_input.size(1)).to(device)
logits = model(src, tgt_input, tgt_mask=tgt_mask)
loss = criterion(logits.reshape(-1, vocab_size), tgt_output.reshape(-1))
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()
avg_loss = total_loss / len(dataloader)
if epoch % 20 == 0:
print(f"Epoch {epoch:3d}/{epochs} | Loss: {avg_loss:.4f}")
# 保存检查点
torch.save(model.state_dict(), f"chat_model_epoch_{epoch}.pth")
# 保存最终模型
torch.save(model.state_dict(), "chat_model_final.pth")
print("训练完成!模型已保存。")
# 简单测试
test_model(model, char2idx, idx2char, device, max_len)
def test_model(model, char2idx, idx2char, device, max_len):
model.eval()
print("\n=== 测试对话 ===")
while True:
user = input("你: ").strip()
if user.lower() == 'q':
break
# 编码用户输入
src = text_to_indices(user, char2idx, max_len)
src_tensor = torch.tensor([src], dtype=torch.long).to(device)
tgt = torch.tensor([[0]], dtype=torch.long).to(device)
generated = []
for _ in range(64): # 最多生成64个字符
tgt_mask = generate_causal_mask(tgt.size(1)).to(device)
logits = model(src_tensor, tgt, tgt_mask=tgt_mask)
next_token_logits = logits[0, -1, :]
probs = torch.softmax(next_token_logits / 0.8, dim=0).cpu().numpy()
next_token = np.random.choice(len(probs), p=probs)
if next_token == 0: # 结束符
break
generated.append(idx2char[next_token])
tgt = torch.cat([tgt, torch.tensor([[next_token]], dtype=torch.long).to(device)], dim=1)
reply = ''.join(generated)
print(f"春梦蝶: {reply}")
if __name__ == "__main__":
train()