| |
| |
| |
| |
|
|
| 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 |
|
|
| |
| 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) |
|
|
| |
| def text_to_indices(text, char2idx, max_len): |
| indices = [char2idx.get(ch, 0) for ch in text] |
| 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) |
|
|
| |
| 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_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 |
|
|
| |
| def train(): |
| print("生成对话数据...") |
| pairs = generate_dialogue_pairs(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): |
| 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() |