Upload 20 files
Browse files- v1_rnn_mini/3.py +144 -0
- v1_rnn_mini/checkpoints_mengdie/char2idx.json +1 -0
- v1_rnn_mini/checkpoints_mengdie/mengdie_final.pth +3 -0
- v1_rnn_mini/run.py +73 -0
- v2_rnn_expanded/4.py +170 -0
- v2_rnn_expanded/checkpoints_mengdie_expanded/char2idx.json +1 -0
- v2_rnn_expanded/checkpoints_mengdie_expanded/mengdie_final.pth +3 -0
- v2_rnn_expanded/run.py +59 -0
- v3_lstm/5.py +165 -0
- v3_lstm/checkpoints_lstm/char2idx.json +1 -0
- v3_lstm/checkpoints_lstm/mengdie_lstm_final.pth +3 -0
- v3_lstm/run.py +53 -0
- v4_transformer/6.py +213 -0
- v4_transformer/README.md +11 -0
- v4_transformer/chat_vocab.json +1 -0
- v4_transformer/run.py +86 -0
- v5_seq2seq_attention/7.py +252 -0
- v5_seq2seq_attention/README.md +10 -0
- v5_seq2seq_attention/run.py +50 -0
- v5_seq2seq_attention/vocab_20m.json +1 -0
v1_rnn_mini/3.py
ADDED
|
@@ -0,0 +1,144 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import numpy as np
|
| 4 |
+
import os
|
| 5 |
+
import time
|
| 6 |
+
import json
|
| 7 |
+
import random
|
| 8 |
+
|
| 9 |
+
# ---------------------------
|
| 10 |
+
# 1. 角色语料(第一人称自然叙述,无对话标记,中英混合)
|
| 11 |
+
# ---------------------------
|
| 12 |
+
text = (
|
| 13 |
+
"大家好,我叫春梦蝶,今年十一岁。我的头发是白色的,眼睛是红色的。"
|
| 14 |
+
"我是双子座,也是一只可爱的猫娘。我有毛茸茸的猫耳朵和长长的尾巴。"
|
| 15 |
+
"我喜欢吃小鱼干和冰淇淋。平时最喜欢晒太阳打瞌睡。"
|
| 16 |
+
"虽然我年纪小,但我很聪明哦。我会用Python写代码,还会训练神经网络。"
|
| 17 |
+
"我的梦想是成为一名厉害的AI工程师。喵~"
|
| 18 |
+
"有时候我会说一点英文,比如 hello, world! I love deep learning. 喵喵。"
|
| 19 |
+
"你愿意和我做朋友吗?我们可以一起聊天、学习、玩捉迷藏。"
|
| 20 |
+
"我的脾气很好,但是如果你欺负我,我会用爪子挠你哦~"
|
| 21 |
+
"双子座的我有时会很活泼,有时也会想一个人静静待着。"
|
| 22 |
+
"今天的天气真好,阳光洒在我的白头发上,闪闪发光。喵~"
|
| 23 |
+
)
|
| 24 |
+
|
| 25 |
+
# 构建字符映射
|
| 26 |
+
chars = sorted(list(set(text)))
|
| 27 |
+
char2idx = {ch: i for i, ch in enumerate(chars)}
|
| 28 |
+
idx2char = {i: ch for ch, i in char2idx.items()}
|
| 29 |
+
vocab_size = len(chars)
|
| 30 |
+
print(f"字符集大小: {vocab_size} (包含汉字、字母、标点、喵~)")
|
| 31 |
+
|
| 32 |
+
# ---------------------------
|
| 33 |
+
# 2. 模型定义(稍加容量以学习角色特征)
|
| 34 |
+
# ---------------------------
|
| 35 |
+
class TinyCharRNN(nn.Module):
|
| 36 |
+
def __init__(self, vocab_size, hidden_size=32): # hidden=32,参数量约 1~2 万
|
| 37 |
+
super().__init__()
|
| 38 |
+
self.embedding = nn.Embedding(vocab_size, hidden_size)
|
| 39 |
+
self.rnn = nn.RNN(hidden_size, hidden_size, batch_first=True)
|
| 40 |
+
self.fc = nn.Linear(hidden_size, vocab_size)
|
| 41 |
+
|
| 42 |
+
def forward(self, x, hidden=None):
|
| 43 |
+
x = self.embedding(x)
|
| 44 |
+
out, hidden = self.rnn(x, hidden)
|
| 45 |
+
out = self.fc(out)
|
| 46 |
+
return out, hidden
|
| 47 |
+
|
| 48 |
+
hidden_size = 32
|
| 49 |
+
model = TinyCharRNN(vocab_size, hidden_size)
|
| 50 |
+
total_params = sum(p.numel() for p in model.parameters())
|
| 51 |
+
print(f"模型参数量: {total_params}")
|
| 52 |
+
|
| 53 |
+
# ---------------------------
|
| 54 |
+
# 3. 训练数据准备(序列长度 100 字符)
|
| 55 |
+
# ---------------------------
|
| 56 |
+
data = torch.tensor([char2idx[ch] for ch in text], dtype=torch.long)
|
| 57 |
+
seq_len = 100 # 每次喂 100 个字符
|
| 58 |
+
epochs = 500
|
| 59 |
+
save_interval = 10 # 每50轮保存一次
|
| 60 |
+
|
| 61 |
+
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
|
| 62 |
+
loss_fn = nn.CrossEntropyLoss()
|
| 63 |
+
os.makedirs("checkpoints_mengdie", exist_ok=True)
|
| 64 |
+
|
| 65 |
+
# ---------------------------
|
| 66 |
+
# 4. 生成函数(让猫娘说话)
|
| 67 |
+
# ---------------------------
|
| 68 |
+
def generate(model, start_char='你', length=200, temperature=0.8):
|
| 69 |
+
model.eval()
|
| 70 |
+
with torch.no_grad():
|
| 71 |
+
if start_char not in char2idx:
|
| 72 |
+
start_char = random.choice(list(char2idx.keys()))
|
| 73 |
+
input_idx = torch.tensor([[char2idx[start_char]]])
|
| 74 |
+
hidden = None
|
| 75 |
+
result = [start_char]
|
| 76 |
+
for _ in range(length):
|
| 77 |
+
logits, hidden = model(input_idx, hidden)
|
| 78 |
+
probs = torch.softmax(logits[0, -1] / temperature, dim=0).cpu().numpy()
|
| 79 |
+
next_idx = np.random.choice(len(probs), p=probs)
|
| 80 |
+
next_char = idx2char[next_idx]
|
| 81 |
+
result.append(next_char)
|
| 82 |
+
input_idx = torch.tensor([[next_idx]])
|
| 83 |
+
return ''.join(result)
|
| 84 |
+
|
| 85 |
+
# ---------------------------
|
| 86 |
+
# 5. 训练循环
|
| 87 |
+
# ---------------------------
|
| 88 |
+
print("\n开始训练春梦蝶猫娘模型(500轮,序列长度100)...\n")
|
| 89 |
+
start_total = time.time()
|
| 90 |
+
|
| 91 |
+
for epoch in range(1, epochs + 1):
|
| 92 |
+
epoch_start = time.time()
|
| 93 |
+
hidden = None
|
| 94 |
+
total_loss = 0
|
| 95 |
+
n_batches = 0
|
| 96 |
+
|
| 97 |
+
# 每次取 seq_len 个字符,步长可以设为 seq_len//2 增加数据利用率,这里简单滑动
|
| 98 |
+
for i in range(0, len(data) - seq_len, seq_len):
|
| 99 |
+
x = data[i:i+seq_len].unsqueeze(0)
|
| 100 |
+
y = data[i+1:i+seq_len+1].unsqueeze(0)
|
| 101 |
+
|
| 102 |
+
logits, hidden = model(x, hidden)
|
| 103 |
+
if hidden is not None:
|
| 104 |
+
hidden = hidden.detach()
|
| 105 |
+
|
| 106 |
+
loss = loss_fn(logits.view(-1, vocab_size), y.view(-1))
|
| 107 |
+
optimizer.zero_grad()
|
| 108 |
+
loss.backward()
|
| 109 |
+
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
| 110 |
+
optimizer.step()
|
| 111 |
+
|
| 112 |
+
total_loss += loss.item()
|
| 113 |
+
n_batches += 1
|
| 114 |
+
|
| 115 |
+
avg_loss = total_loss / n_batches
|
| 116 |
+
epoch_time = time.time() - epoch_start
|
| 117 |
+
|
| 118 |
+
if epoch % save_interval == 0:
|
| 119 |
+
# 生成一段猫娘风格的文本
|
| 120 |
+
sample = generate(model, start_char='我', length=150, temperature=0.7)
|
| 121 |
+
print(f"Epoch {epoch:4d}/{epochs} | Loss: {avg_loss:.4f} | Time: {epoch_time:.2f}s")
|
| 122 |
+
print(f"春梦蝶说: {sample[:120]}...\n")
|
| 123 |
+
checkpoint_path = f"checkpoints_mengdie/mengdie_epoch_{epoch}.pth"
|
| 124 |
+
torch.save(model.state_dict(), checkpoint_path)
|
| 125 |
+
print(f"已保存模型到: {checkpoint_path}\n")
|
| 126 |
+
else:
|
| 127 |
+
# 每10轮打印一次 loss 即可,避免刷屏
|
| 128 |
+
if epoch % 10 == 0:
|
| 129 |
+
print(f"Epoch {epoch:4d}/{epochs} | Loss: {avg_loss:.4f} | Time: {epoch_time:.2f}s")
|
| 130 |
+
|
| 131 |
+
total_time = time.time() - start_total
|
| 132 |
+
print(f"\n训练完成!总耗时: {total_time:.2f} 秒 (约 {total_time/60:.1f} 分钟)")
|
| 133 |
+
final_path = "checkpoints_mengdie/mengdie_final.pth"
|
| 134 |
+
torch.save(model.state_dict(), final_path)
|
| 135 |
+
print(f"最终模型已保存到 {final_path}")
|
| 136 |
+
|
| 137 |
+
# 保存字符映射
|
| 138 |
+
with open("checkpoints_mengdie/char2idx.json", "w", encoding="utf-8") as f:
|
| 139 |
+
json.dump(char2idx, f, ensure_ascii=False)
|
| 140 |
+
|
| 141 |
+
print("\n=== 最终生成的猫娘自我介绍 ===")
|
| 142 |
+
print(generate(model, start_char='大', length=300, temperature=0.7))
|
| 143 |
+
print("\n=== 随机性更强的猫娘发言(温度=1.1) ===")
|
| 144 |
+
print(generate(model, start_char='喵', length=300, temperature=1.1))
|
v1_rnn_mini/checkpoints_mengdie/char2idx.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{" ": 0, "!": 1, ",": 2, ".": 3, "A": 4, "I": 5, "P": 6, "a": 7, "d": 8, "e": 9, "g": 10, "h": 11, "i": 12, "l": 13, "n": 14, "o": 15, "p": 16, "r": 17, "t": 18, "v": 19, "w": 20, "y": 21, "、": 22, "。": 23, "一": 24, "上": 25, "个": 26, "为": 27, "也": 28, "习": 29, "人": 30, "今": 31, "代": 32, "以": 33, "们": 34, "会": 35, "但": 36, "你": 37, "候": 38, "做": 39, "光": 40, "写": 41, "冰": 42, "十": 43, "厉": 44, "友": 45, "双": 46, "发": 47, "只": 48, "叫": 49, "可": 50, "吃": 51, "名": 52, "吗": 53, "和": 54, "哦": 55, "喜": 56, "喵": 57, "在": 58, "大": 59, "天": 60, "太": 61, "头": 62, "好": 63, "如": 64, "娘": 65, "子": 66, "学": 67, "害": 68, "家": 69, "小": 70, "尾": 71, "岁": 72, "工": 73, "巴": 74, "师": 75, "干": 76, "平": 77, "年": 78, "座": 79, "待": 80, "很": 81, "想": 82, "意": 83, "愿": 84, "成": 85, "我": 86, "打": 87, "挠": 88, "捉": 89, "文": 90, "时": 91, "明": 92, "春": 93, "是": 94, "晒": 95, "最": 96, "有": 97, "朋": 98, "朵": 99, "果": 100, "梦": 101, "欢": 102, "欺": 103, "比": 104, "毛": 105, "气": 106, "泼": 107, "洒": 108, "活": 109, "淇": 110, "淋": 111, "点": 112, "然": 113, "爪": 114, "爱": 115, "猫": 116, "玩": 117, "用": 118, "白": 119, "的": 120, "真": 121, "眼": 122, "着": 123, "睛": 124, "睡": 125, "瞌": 126, "码": 127, "神": 128, "程": 129, "红": 130, "纪": 131, "练": 132, "经": 133, "络": 134, "网": 135, "耳": 136, "聊": 137, "聪": 138, "脾": 139, "色": 140, "英": 141, "茸": 142, "藏": 143, "虽": 144, "蝶": 145, "训": 146, "说": 147, "负": 148, "起": 149, "还": 150, "迷": 151, "长": 152, "闪": 153, "阳": 154, "静": 155, "鱼": 156, ",": 157, "?": 158, "~": 159}
|
v1_rnn_mini/checkpoints_mengdie/mengdie_final.pth
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:7c1141af9cc54b0a427158f89b4b244629b3da13e6ad47ba9e7074a06a1a1fee
|
| 3 |
+
size 53407
|
v1_rnn_mini/run.py
ADDED
|
@@ -0,0 +1,73 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import numpy as np
|
| 4 |
+
import json
|
| 5 |
+
import random
|
| 6 |
+
|
| 7 |
+
# 加载字符映射
|
| 8 |
+
with open("checkpoints_mengdie/char2idx.json", "r", encoding="utf-8") as f:
|
| 9 |
+
char2idx = json.load(f)
|
| 10 |
+
|
| 11 |
+
idx2char = {int(v): k for k, v in char2idx.items()}
|
| 12 |
+
vocab_size = len(char2idx)
|
| 13 |
+
|
| 14 |
+
print(f"词表大小: {vocab_size}")
|
| 15 |
+
print(f"词表示例: {list(char2idx.items())[:10]}")
|
| 16 |
+
|
| 17 |
+
# 模型结构(必须与训练时一致!!!)
|
| 18 |
+
class TinyCharRNN(nn.Module):
|
| 19 |
+
def __init__(self, vocab_size, hidden_size=32):
|
| 20 |
+
super().__init__()
|
| 21 |
+
self.embedding = nn.Embedding(vocab_size, hidden_size)
|
| 22 |
+
self.rnn = nn.RNN(hidden_size, hidden_size, batch_first=True)
|
| 23 |
+
self.fc = nn.Linear(hidden_size, vocab_size)
|
| 24 |
+
def forward(self, x, hidden=None):
|
| 25 |
+
x = self.embedding(x)
|
| 26 |
+
out, hidden = self.rnn(x, hidden)
|
| 27 |
+
out = self.fc(out)
|
| 28 |
+
return out, hidden
|
| 29 |
+
|
| 30 |
+
# 加载模型
|
| 31 |
+
model = TinyCharRNN(vocab_size, hidden_size=32)
|
| 32 |
+
model.load_state_dict(torch.load("checkpoints_mengdie/mengdie_final.pth", map_location='cpu'))
|
| 33 |
+
model.eval()
|
| 34 |
+
print("春梦蝶猫娘模型加载成功!喵~\n")
|
| 35 |
+
|
| 36 |
+
def generate_response(prompt, length=150, temperature=0.8):
|
| 37 |
+
"""根据提示词生成猫娘的回应"""
|
| 38 |
+
if not prompt:
|
| 39 |
+
prompt = random.choice(list(char2idx.keys()))
|
| 40 |
+
# 将prompt转为索引
|
| 41 |
+
indices = []
|
| 42 |
+
for ch in prompt:
|
| 43 |
+
if ch in char2idx:
|
| 44 |
+
indices.append(char2idx[ch])
|
| 45 |
+
else:
|
| 46 |
+
# 如果字符不在词表中,跳过(或者用空格代替)
|
| 47 |
+
# 这里跳过,不影响生成
|
| 48 |
+
continue
|
| 49 |
+
if not indices:
|
| 50 |
+
# 如果全部跳过,就用一个默认字符
|
| 51 |
+
indices = [char2idx['你']]
|
| 52 |
+
input_tensor = torch.tensor([indices])
|
| 53 |
+
hidden = None
|
| 54 |
+
result = list(prompt) # 保留原始输入
|
| 55 |
+
with torch.no_grad():
|
| 56 |
+
for _ in range(length):
|
| 57 |
+
logits, hidden = model(input_tensor, hidden)
|
| 58 |
+
probs = torch.softmax(logits[0, -1] / temperature, dim=0).cpu().numpy()
|
| 59 |
+
next_idx = np.random.choice(len(probs), p=probs)
|
| 60 |
+
next_char = idx2char[next_idx]
|
| 61 |
+
result.append(next_char)
|
| 62 |
+
input_tensor = torch.tensor([[next_idx]])
|
| 63 |
+
return ''.join(result)
|
| 64 |
+
|
| 65 |
+
print("你可以开始和春梦蝶聊天了!输入 'q' 退出。")
|
| 66 |
+
while True:
|
| 67 |
+
user = input("\n你: ")
|
| 68 |
+
if user.lower() == 'q':
|
| 69 |
+
break
|
| 70 |
+
# 使用用户输入的最后几个字符作为生成起点
|
| 71 |
+
start = user[-5:] if len(user) >= 5 else user
|
| 72 |
+
reply = generate_response(start, length=150, temperature=0.85)
|
| 73 |
+
print(f"春梦蝶: {reply}")
|
v2_rnn_expanded/4.py
ADDED
|
@@ -0,0 +1,170 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import numpy as np
|
| 4 |
+
import os
|
| 5 |
+
import time
|
| 6 |
+
import json
|
| 7 |
+
import random
|
| 8 |
+
|
| 9 |
+
# ===================== 1. 扩充后的角色语料(约2500字符) =====================
|
| 10 |
+
text = (
|
| 11 |
+
"大家好,我叫春梦蝶,今年十一岁。我的头发是白色的,眼睛是红色的。"
|
| 12 |
+
"我是双子座,也是一只可爱的猫娘。我有毛茸茸的猫耳朵和长长的尾巴。"
|
| 13 |
+
"我喜欢吃小鱼干和冰淇淋。平时最喜欢晒太阳打瞌睡。"
|
| 14 |
+
"虽然我年纪小,但我很聪明哦。我会用Python写代码,还会训练神经网络。"
|
| 15 |
+
"我的梦想是成为一名厉害的AI工程师。喵~"
|
| 16 |
+
"有时候我会说一点英文,比如 hello, world! I love deep learning. 喵喵。"
|
| 17 |
+
"你愿意和我做朋友吗?我们可以一起聊天、学习、玩捉迷藏。"
|
| 18 |
+
"我的脾气很好,但是如果你欺负我,我会用爪子挠你哦~"
|
| 19 |
+
"双子座的我有时会很活泼,有时也会想一个人静静待着。"
|
| 20 |
+
"今天的天气真好,阳光洒在我的白头发上,闪闪发光。喵~"
|
| 21 |
+
# ----- 以下为新增扩充内容(约2000字)-----
|
| 22 |
+
"今天我在院子里追蝴蝶,但是没追到,气死我啦!喵喵喵!"
|
| 23 |
+
"有人问我为什么喜欢深度学习,我说因为它像猫抓板一样有趣,而且永远不会腻。"
|
| 24 |
+
"我的红眼睛在黑暗里会微微发光,像两颗小星星。晚上走路都不用开灯呢。"
|
| 25 |
+
"最讨厌的事情就是洗澡了,水会把我的毛弄湿,感觉好重好重。"
|
| 26 |
+
"有时候我会用猫爪在键盘上打字,速度虽然慢,但是很可爱呀。"
|
| 27 |
+
"我最喜欢的颜色是白色和红色,因为和我自己一样。嘿嘿。"
|
| 28 |
+
"昨天我梦见自己变成了一台超级计算机,算力无穷大,好开心。"
|
| 29 |
+
"如果你给我买小鱼干,我就让你摸一下我的猫耳朵,就一下哦。"
|
| 30 |
+
"学习PyTorch的时候遇到bug,我会对着屏幕喵喵叫,然后bug就消失了,神奇吧?"
|
| 31 |
+
"春天的时候樱花会落在我白色的头发上,像雪花一样美。"
|
| 32 |
+
"我有一个秘密:其实我可以用猫尾巴控制鼠标,不需要用手。"
|
| 33 |
+
"我不喜欢吃苦瓜,谁给我吃苦瓜我就咬谁。"
|
| 34 |
+
"每次看到镜子里的自己,我都会感叹:啊,这只猫娘真可爱。"
|
| 35 |
+
"今天尝试用RNN写一首诗:喵喵喵,小鱼干,神经网络真有趣。"
|
| 36 |
+
"我的耳朵可以听到很远的声音,比如你心里在夸我可爱。"
|
| 37 |
+
"双子座的优点就是适应力强,不管是晴天还是下雨,我都能找到睡觉的好地方。"
|
| 38 |
+
"英文里有个词叫‘catnap’,就是猫打盹,我每天都要catnap好几次。"
|
| 39 |
+
"如果你问我年龄,我会说十一岁,但我的猫龄相当于人类的六十岁?不对,猫的年龄算法不一样,反正我永远年轻。"
|
| 40 |
+
"我喜欢在凌晨三点训练模型,因为那时候灵感最旺盛,虽然第二天会困。"
|
| 41 |
+
"有人类问我为什么不戴眼镜,我说我的红眼睛视力2.0,不需要。"
|
| 42 |
+
"我的白色头发每天早上都会翘起来,要用梳子梳好久,好麻烦喵。"
|
| 43 |
+
"今天学会了用卷积神经网络做图像分类,然后给自己的照片分类,结果是‘超可爱猫娘’类。"
|
| 44 |
+
"我不喜欢喝牛奶,但是喜欢喝鱼汤。很矛盾对吧?因为我是猫娘呀。"
|
| 45 |
+
"有时候我会对着月亮喵喵叫,邻居家的狗也会跟着叫,然后整个小区都热闹起来。"
|
| 46 |
+
"我的梦想除了当AI工程师,还想去猫星球旅行一次。不知道那里有没有小鱼干卖。"
|
| 47 |
+
"如果你送我一条小鱼干,我就送你一个我亲手训练的小模型,虽然只会输出乱码,但是很有心意。"
|
| 48 |
+
"今天的晚霞是橙红色的,和我的眼睛不一样,但是也很美。"
|
| 49 |
+
"我写代码的时候喜欢把变量名取成fish、cat、meow,这样心情会很好。"
|
| 50 |
+
"有一次我试着用强化学习训练一只虚拟猫抓老鼠,结果那只猫学会了睡觉,和我一样。"
|
| 51 |
+
"虽然我只有十一岁,但我觉得我已经很成熟了,至少比我家隔壁的三岁小孩成熟。"
|
| 52 |
+
"我的尾巴尖有一小撮白毛,特别柔软,我自己没事就会摸一摸。"
|
| 53 |
+
"下雨天我会趴在窗台上数雨滴,数到一百就睡着了。"
|
| 54 |
+
"我最喜欢的动画片是《猫娘乐园》,每次看都会流泪,太感动了。"
|
| 55 |
+
"如果你觉得我可爱,请给我点赞,喵~"
|
| 56 |
+
)
|
| 57 |
+
|
| 58 |
+
# 构建字符映射
|
| 59 |
+
chars = sorted(list(set(text)))
|
| 60 |
+
char2idx = {ch: i for i, ch in enumerate(chars)}
|
| 61 |
+
idx2char = {i: ch for ch, i in char2idx.items()}
|
| 62 |
+
vocab_size = len(chars)
|
| 63 |
+
print(f"字符集大小: {vocab_size} (包括汉字、字母、标点、喵)")
|
| 64 |
+
print(f"总语料长度: {len(text)} 字符")
|
| 65 |
+
|
| 66 |
+
# ===================== 2. 模型定义(与v1相同) =====================
|
| 67 |
+
class TinyCharRNN(nn.Module):
|
| 68 |
+
def __init__(self, vocab_size, hidden_size=32):
|
| 69 |
+
super().__init__()
|
| 70 |
+
self.embedding = nn.Embedding(vocab_size, hidden_size)
|
| 71 |
+
self.rnn = nn.RNN(hidden_size, hidden_size, batch_first=True)
|
| 72 |
+
self.fc = nn.Linear(hidden_size, vocab_size)
|
| 73 |
+
|
| 74 |
+
def forward(self, x, hidden=None):
|
| 75 |
+
x = self.embedding(x)
|
| 76 |
+
out, hidden = self.rnn(x, hidden)
|
| 77 |
+
out = self.fc(out)
|
| 78 |
+
return out, hidden
|
| 79 |
+
|
| 80 |
+
hidden_size = 32
|
| 81 |
+
model = TinyCharRNN(vocab_size, hidden_size)
|
| 82 |
+
total_params = sum(p.numel() for p in model.parameters())
|
| 83 |
+
print(f"模型参数量: {total_params}")
|
| 84 |
+
|
| 85 |
+
# ===================== 3. 训练数据准备 =====================
|
| 86 |
+
data = torch.tensor([char2idx[ch] for ch in text], dtype=torch.long)
|
| 87 |
+
seq_len = 256 # 每次喂256个字符(在200~500之间)
|
| 88 |
+
epochs = 500
|
| 89 |
+
save_interval = 10 # 每10轮保存一次
|
| 90 |
+
|
| 91 |
+
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
|
| 92 |
+
loss_fn = nn.CrossEntropyLoss()
|
| 93 |
+
os.makedirs("checkpoints_mengdie_expanded", exist_ok=True)
|
| 94 |
+
|
| 95 |
+
# ===================== 4. 生成函数 =====================
|
| 96 |
+
def generate(model, start_char='你', length=200, temperature=0.8):
|
| 97 |
+
model.eval()
|
| 98 |
+
with torch.no_grad():
|
| 99 |
+
if start_char not in char2idx:
|
| 100 |
+
start_char = random.choice(list(char2idx.keys()))
|
| 101 |
+
input_idx = torch.tensor([[char2idx[start_char]]])
|
| 102 |
+
hidden = None
|
| 103 |
+
result = [start_char]
|
| 104 |
+
for _ in range(length):
|
| 105 |
+
logits, hidden = model(input_idx, hidden)
|
| 106 |
+
probs = torch.softmax(logits[0, -1] / temperature, dim=0).cpu().numpy()
|
| 107 |
+
next_idx = np.random.choice(len(probs), p=probs)
|
| 108 |
+
next_char = idx2char[next_idx]
|
| 109 |
+
result.append(next_char)
|
| 110 |
+
input_idx = torch.tensor([[next_idx]])
|
| 111 |
+
return ''.join(result)
|
| 112 |
+
|
| 113 |
+
# ===================== 5. 训练循环 =====================
|
| 114 |
+
print("\n开始训练春梦蝶猫娘模型(扩充语料,seq_len=256,500轮)...\n")
|
| 115 |
+
start_total = time.time()
|
| 116 |
+
|
| 117 |
+
for epoch in range(1, epochs + 1):
|
| 118 |
+
epoch_start = time.time()
|
| 119 |
+
hidden = None
|
| 120 |
+
total_loss = 0
|
| 121 |
+
n_batches = 0
|
| 122 |
+
|
| 123 |
+
# 滑动窗口,步长设为seq_len//2,增加数据利用率
|
| 124 |
+
step = seq_len // 2
|
| 125 |
+
for i in range(0, len(data) - seq_len, step):
|
| 126 |
+
x = data[i:i+seq_len].unsqueeze(0)
|
| 127 |
+
y = data[i+1:i+seq_len+1].unsqueeze(0)
|
| 128 |
+
|
| 129 |
+
logits, hidden = model(x, hidden)
|
| 130 |
+
if hidden is not None:
|
| 131 |
+
hidden = hidden.detach()
|
| 132 |
+
|
| 133 |
+
loss = loss_fn(logits.view(-1, vocab_size), y.view(-1))
|
| 134 |
+
optimizer.zero_grad()
|
| 135 |
+
loss.backward()
|
| 136 |
+
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
| 137 |
+
optimizer.step()
|
| 138 |
+
|
| 139 |
+
total_loss += loss.item()
|
| 140 |
+
n_batches += 1
|
| 141 |
+
|
| 142 |
+
avg_loss = total_loss / n_batches
|
| 143 |
+
epoch_time = time.time() - epoch_start
|
| 144 |
+
|
| 145 |
+
if epoch % save_interval == 0:
|
| 146 |
+
# 生成一段猫娘风格的文本
|
| 147 |
+
sample = generate(model, start_char='我', length=180, temperature=0.7)
|
| 148 |
+
print(f"Epoch {epoch:4d}/{epochs} | Loss: {avg_loss:.4f} | Time: {epoch_time:.2f}s")
|
| 149 |
+
print(f"春梦蝶说: {sample[:150]}...\n")
|
| 150 |
+
checkpoint_path = f"checkpoints_mengdie_expanded/mengdie_epoch_{epoch}.pth"
|
| 151 |
+
torch.save(model.state_dict(), checkpoint_path)
|
| 152 |
+
print(f"已保存模型到: {checkpoint_path}\n")
|
| 153 |
+
else:
|
| 154 |
+
if epoch % 10 == 0:
|
| 155 |
+
print(f"Epoch {epoch:4d}/{epochs} | Loss: {avg_loss:.4f} | Time: {epoch_time:.2f}s")
|
| 156 |
+
|
| 157 |
+
total_time = time.time() - start_total
|
| 158 |
+
print(f"\n训练完成!总耗时: {total_time:.2f} 秒 (约 {total_time/60:.1f} 分钟)")
|
| 159 |
+
final_path = "checkpoints_mengdie_expanded/mengdie_final.pth"
|
| 160 |
+
torch.save(model.state_dict(), final_path)
|
| 161 |
+
print(f"最终模型已保存到 {final_path}")
|
| 162 |
+
|
| 163 |
+
# 保存字符映射
|
| 164 |
+
with open("checkpoints_mengdie_expanded/char2idx.json", "w", encoding="utf-8") as f:
|
| 165 |
+
json.dump(char2idx, f, ensure_ascii=False)
|
| 166 |
+
|
| 167 |
+
print("\n=== 最终生成的猫娘发言(温度0.7) ===")
|
| 168 |
+
print(generate(model, start_char='我', length=400, temperature=0.7))
|
| 169 |
+
print("\n=== 随机性更强的版本(温度1.1) ===")
|
| 170 |
+
print(generate(model, start_char='喵', length=400, temperature=1.1))
|
v2_rnn_expanded/checkpoints_mengdie_expanded/char2idx.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{" ": 0, "!": 1, ",": 2, ".": 3, "0": 4, "2": 5, "A": 6, "I": 7, "N": 8, "P": 9, "R": 10, "T": 11, "a": 12, "b": 13, "c": 14, "d": 15, "e": 16, "f": 17, "g": 18, "h": 19, "i": 20, "l": 21, "m": 22, "n": 23, "o": 24, "p": 25, "r": 26, "s": 27, "t": 28, "u": 29, "v": 30, "w": 31, "y": 32, "‘": 33, "’": 34, "、": 35, "。": 36, "《": 37, "》": 38, "一": 39, "三": 40, "上": 41, "下": 42, "不": 43, "且": 44, "两": 45, "个": 46, "为": 47, "久": 48, "么": 49, "乐": 50, "也": 51, "习": 52, "买": 53, "乱": 54, "了": 55, "事": 56, "二": 57, "于": 58, "亮": 59, "亲": 60, "人": 61, "什": 62, "今": 63, "代": 64, "以": 65, "们": 66, "优": 67, "会": 68, "但": 69, "你": 70, "候": 71, "做": 72, "像": 73, "光": 74, "六": 75, "其": 76, "写": 77, "冰": 78, "凌": 79, "几": 80, "出": 81, "分": 82, "别": 83, "到": 84, "制": 85, "力": 86, "动": 87, "化": 88, "区": 89, "十": 90, "卖": 91, "卷": 92, "厉": 93, "厌": 94, "去": 95, "友": 96, "双": 97, "反": 98, "发": 99, "取": 100, "变": 101, "只": 102, "叫": 103, "可": 104, "台": 105, "叹": 106, "吃": 107, "名": 108, "后": 109, "吗": 110, "吧": 111, "听": 112, "呀": 113, "呢": 114, "和": 115, "咬": 116, "哦": 117, "啊": 118, "啦": 119, "喜": 120, "喝": 121, "喵": 122, "嘿": 123, "因": 124, "园": 125, "困": 126, "图": 127, "在": 128, "地": 129, "型": 130, "壁": 131, "声": 132, "大": 133, "天": 134, "太": 135, "失": 136, "头": 137, "夸": 138, "奇": 139, "奶": 140, "好": 141, "如": 142, "娘": 143, "子": 144, "字": 145, "学": 146, "孩": 147, "它": 148, "实": 149, "害": 150, "家": 151, "密": 152, "对": 153, "小": 154, "少": 155, "尖": 156, "尝": 157, "就": 158, "尾": 159, "居": 160, "屏": 161, "岁": 162, "工": 163, "己": 164, "已": 165, "巴": 166, "师": 167, "幕": 168, "干": 169, "平": 170, "年": 171, "应": 172, "度": 173, "座": 174, "开": 175, "弄": 176, "强": 177, "当": 178, "待": 179, "很": 180, "得": 181, "微": 182, "心": 183, "情": 184, "想": 185, "意": 186, "感": 187, "愿": 188, "慢": 189, "成": 190, "我": 191, "戴": 192, "手": 193, "打": 194, "找": 195, "把": 196, "抓": 197, "拟": 198, "挠": 199, "捉": 200, "控": 201, "摸": 202, "撮": 203, "数": 204, "整": 205, "文": 206, "方": 207, "旅": 208, "无": 209, "早": 210, "时": 211, "旺": 212, "明": 213, "星": 214, "春": 215, "昨": 216, "是": 217, "晒": 218, "晚": 219, "晨": 220, "晴": 221, "暗": 222, "最": 223, "月": 224, "有": 225, "朋": 226, "朵": 227, "机": 228, "条": 229, "来": 230, "板": 231, "果": 232, "柔": 233, "标": 234, "样": 235, "梦": 236, "梳": 237, "模": 238, "樱": 239, "橙": 240, "次": 241, "欢": 242, "欺": 243, "正": 244, "死": 245, "每": 246, "比": 247, "毛": 248, "气": 249, "水": 250, "永": 251, "汤": 252, "没": 253, "法": 254, "泪": 255, "泼": 256, "洒": 257, "洗": 258, "活": 259, "流": 260, "消": 261, "淇": 262, "淋": 263, "深": 264, "湿": 265, "滴": 266, "澡": 267, "灯": 268, "灵": 269, "点": 270, "烦": 271, "热": 272, "然": 273, "照": 274, "熟": 275, "爪": 276, "爱": 277, "片": 278, "牛": 279, "特": 280, "狗": 281, "猫": 282, "玩": 283, "球": 284, "瓜": 285, "用": 286, "画": 287, "白": 288, "百": 289, "的": 290, "盘": 291, "盛": 292, "相": 293, "盹": 294, "盾": 295, "看": 296, "真": 297, "眼": 298, "着": 299, "睛": 300, "睡": 301, "瞌": 302, "矛": 303, "知": 304, "码": 305, "神": 306, "秘": 307, "积": 308, "程": 309, "穷": 310, "窗": 311, "第": 312, "算": 313, "管": 314, "类": 315, "红": 316, "级": 317, "纪": 318, "练": 319, "经": 320, "结": 321, "给": 322, "络": 323, "网": 324, "美": 325, "翘": 326, "老": 327, "而": 328, "耳": 329, "聊": 330, "聪": 331, "能": 332, "脾": 333, "腻": 334, "自": 335, "至": 336, "色": 337, "花": 338, "苦": 339, "英": 340, "茸": 341, "落": 342, "藏": 343, "虚": 344, "虽": 345, "蝴": 346, "蝶": 347, "行": 348, "要": 349, "见": 350, "视": 351, "觉": 352, "计": 353, "讨": 354, "让": 355, "训": 356, "词": 357, "试": 358, "诗": 359, "说": 360, "请": 361, "谁": 362, "负": 363, "赞": 364, "走": 365, "起": 366, "超": 367, "趣": 368, "趴": 369, "跟": 370, "路": 371, "软": 372, "轻": 373, "输": 374, "还": 375, "这": 376, "远": 377, "迷": 378, "追": 379, "送": 380, "适": 381, "速": 382, "遇": 383, "道": 384, "那": 385, "邻": 386, "都": 387, "里": 388, "重": 389, "量": 390, "键": 391, "镜": 392, "长": 393, "闪": 394, "问": 395, "闹": 396, "阳": 397, "院": 398, "除": 399, "隔": 400, "雨": 401, "雪": 402, "需": 403, "霞": 404, "静": 405, "音": 406, "颗": 407, "颜": 408, "首": 409, "鱼": 410, "麻": 411, "黑": 412, "鼠": 413, "龄": 414, "!": 415, ",": 416, ":": 417, "?": 418, "~": 419}
|
v2_rnn_expanded/checkpoints_mengdie_expanded/mengdie_final.pth
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:abddd13ccf5039ca8334b6e4eaddb54f8324d04772ae9a52df6d4c1f81ae566a
|
| 3 |
+
size 120991
|
v2_rnn_expanded/run.py
ADDED
|
@@ -0,0 +1,59 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import numpy as np
|
| 4 |
+
import json
|
| 5 |
+
import random
|
| 6 |
+
|
| 7 |
+
# 加载字符映射
|
| 8 |
+
with open("checkpoints_mengdie_expanded/char2idx.json", "r", encoding="utf-8") as f:
|
| 9 |
+
char2idx = json.load(f)
|
| 10 |
+
idx2char = {int(v): k for k, v in char2idx.items()}
|
| 11 |
+
vocab_size = len(char2idx)
|
| 12 |
+
|
| 13 |
+
class TinyCharRNN(nn.Module):
|
| 14 |
+
def __init__(self, vocab_size, hidden_size=32):
|
| 15 |
+
super().__init__()
|
| 16 |
+
self.embedding = nn.Embedding(vocab_size, hidden_size)
|
| 17 |
+
self.rnn = nn.RNN(hidden_size, hidden_size, batch_first=True)
|
| 18 |
+
self.fc = nn.Linear(hidden_size, vocab_size)
|
| 19 |
+
def forward(self, x, hidden=None):
|
| 20 |
+
x = self.embedding(x)
|
| 21 |
+
out, hidden = self.rnn(x, hidden)
|
| 22 |
+
out = self.fc(out)
|
| 23 |
+
return out, hidden
|
| 24 |
+
|
| 25 |
+
model = TinyCharRNN(vocab_size, hidden_size=32)
|
| 26 |
+
model.load_state_dict(torch.load("checkpoints_mengdie_expanded/mengdie_final.pth", map_location='cpu'))
|
| 27 |
+
model.eval()
|
| 28 |
+
print("春梦蝶猫娘模型(扩充版)加载成功!喵~\n")
|
| 29 |
+
|
| 30 |
+
def generate_response(prompt, length=180, temperature=0.8):
|
| 31 |
+
if not prompt:
|
| 32 |
+
prompt = random.choice(list(char2idx.keys()))
|
| 33 |
+
indices = []
|
| 34 |
+
for ch in prompt:
|
| 35 |
+
if ch in char2idx:
|
| 36 |
+
indices.append(char2idx[ch])
|
| 37 |
+
if not indices:
|
| 38 |
+
indices = [char2idx['你']]
|
| 39 |
+
input_tensor = torch.tensor([indices])
|
| 40 |
+
hidden = None
|
| 41 |
+
result = list(prompt)
|
| 42 |
+
with torch.no_grad():
|
| 43 |
+
for _ in range(length):
|
| 44 |
+
logits, hidden = model(input_tensor, hidden)
|
| 45 |
+
probs = torch.softmax(logits[0, -1] / temperature, dim=0).cpu().numpy()
|
| 46 |
+
next_idx = np.random.choice(len(probs), p=probs)
|
| 47 |
+
next_char = idx2char[next_idx]
|
| 48 |
+
result.append(next_char)
|
| 49 |
+
input_tensor = torch.tensor([[next_idx]])
|
| 50 |
+
return ''.join(result)
|
| 51 |
+
|
| 52 |
+
print("开始和春梦蝶聊天!输入 'q' 退出。")
|
| 53 |
+
while True:
|
| 54 |
+
user = input("\n你: ")
|
| 55 |
+
if user.lower() == 'q':
|
| 56 |
+
break
|
| 57 |
+
start = user[-5:] if len(user) >= 5 else user
|
| 58 |
+
reply = generate_response(start, length=200, temperature=0.85)
|
| 59 |
+
print(f"春梦蝶: {reply}")
|
v3_lstm/5.py
ADDED
|
@@ -0,0 +1,165 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import numpy as np
|
| 4 |
+
import os
|
| 5 |
+
import time
|
| 6 |
+
import json
|
| 7 |
+
import random
|
| 8 |
+
|
| 9 |
+
# ===================== 1. 扩充语料(保持原有并稍加丰富*可复用v2的char2idx.json,词表基本一样*) =====================
|
| 10 |
+
text = (
|
| 11 |
+
"大家好,我叫春梦蝶,今年十一岁。我的头发是白色的,眼睛是红色的。"
|
| 12 |
+
"我是双子座,也是一只可爱的猫娘。我有毛茸茸的猫耳朵和长长的尾巴。"
|
| 13 |
+
"我喜欢吃小鱼干和冰淇淋。平时最喜欢晒太阳打瞌睡。"
|
| 14 |
+
"虽然我年纪小,但我很聪明哦。我会用Python写代码,还会训练神经网络。"
|
| 15 |
+
"我的梦想是成为一名厉害的AI工程师。喵~"
|
| 16 |
+
"有时候我会说一点英文,比如 hello, world! I love deep learning. 喵喵。"
|
| 17 |
+
"你愿意和我做朋友吗?我们可以一起聊天、学习、玩捉迷藏。"
|
| 18 |
+
"我的脾气很好,但是如果你欺负我,我会用爪子挠你哦~"
|
| 19 |
+
"双子座的我有时会很活泼,有时也会想一个人静静待着。"
|
| 20 |
+
"今天的天气真好,阳光洒在我的白头发上,闪闪发光。喵~"
|
| 21 |
+
"今天我在院子里追蝴蝶,但是没追到,气死我啦!喵喵喵!"
|
| 22 |
+
"有人问我为什么喜欢深度学习,我说因为它像猫抓板一样有趣,而且永远不会腻。"
|
| 23 |
+
"我的红眼睛在黑暗里会微微发光,像两颗小星星。晚上走路都不用开灯呢。"
|
| 24 |
+
"最讨厌的事情就是洗澡了,水会把我的毛弄湿,感觉好重好重。"
|
| 25 |
+
"有时候我会用猫爪在键盘上打字,速度虽然慢,但是很可爱呀。"
|
| 26 |
+
"我最喜欢的颜色是白色和红色,因为和我自己一样。嘿嘿。"
|
| 27 |
+
"昨天我梦见自己变成了一台超级计算机,算力无穷大,好开心。"
|
| 28 |
+
"如果你给我买小鱼干,我就让你摸一下我的猫耳朵,就一下哦。"
|
| 29 |
+
"学习PyTorch的时候遇到bug,我会对着屏幕喵喵叫,然后bug就消失了,神奇吧?"
|
| 30 |
+
"春天的时候樱花会落在我白色的头发上,像雪花一样美。"
|
| 31 |
+
"我有一个秘密:其实我可以用猫尾巴控制鼠标,不需要用手。"
|
| 32 |
+
"我不喜欢吃苦瓜,谁给我吃苦瓜我就咬谁。"
|
| 33 |
+
"每次看到镜子里的自己,我都会感叹:啊,这只猫娘真可爱。"
|
| 34 |
+
"今天尝试用RNN写一首诗:喵喵喵,小鱼干,神经网络真有趣。"
|
| 35 |
+
"我的耳朵可以听到很远的声音,比如你心里在夸我可爱。"
|
| 36 |
+
"双子座的优点就是适应力强,不管是晴天还是下雨,我都能找到睡觉的好地方。"
|
| 37 |
+
"英文里有个词叫‘catnap’,就是猫打盹,我每天都要catnap好几次。"
|
| 38 |
+
"如果你问我年龄,我会说十一岁,但我的猫龄相当于人类的六十岁?不对,猫的年龄算法不一样,反正我永远年轻。"
|
| 39 |
+
"我喜欢在凌晨三点训练模型,因为那时候灵感最旺盛,虽然第二天会困。"
|
| 40 |
+
"有人类问我为什么不戴眼镜,我说我的红眼睛视力2.0,不需要。"
|
| 41 |
+
"我的白色头发每天早上都会翘起来,要用梳子梳好久,好麻烦喵。"
|
| 42 |
+
"今天学会了用卷积神经网络做图像分类,然后给自己的照片分类,结果是‘超可爱猫娘’类。"
|
| 43 |
+
"我不喜欢喝牛奶,但是喜欢喝鱼汤。很矛盾对吧?因为我是猫娘呀。"
|
| 44 |
+
"有时候我会对着月亮喵喵叫,邻居家的狗也会跟着叫,然后整个小区都热闹起来。"
|
| 45 |
+
"我的梦想除了当AI工程师,还想去猫星球旅行一次。不知道那里有没有小鱼干卖。"
|
| 46 |
+
"如果你送我一条小鱼干,我就送你一个我亲手训练的小模型,虽然只会输出乱码,但是很有心意。"
|
| 47 |
+
"今天的晚霞是橙红色的,和我的眼睛不一样,但是也很美。"
|
| 48 |
+
"我写代码的时候喜欢把变量名取成fish、cat、meow,这样心情会很好。"
|
| 49 |
+
"有一次我试着用强化学习训练一只虚拟猫抓老鼠,结果那只猫学会了睡觉,和我一样。"
|
| 50 |
+
"虽然我只有十一岁,但我觉得我已经很成熟了,至少比我家隔壁的三岁小孩成熟。"
|
| 51 |
+
"我的尾巴尖有一小撮白毛,特别柔软,我自己没事就会摸一摸。"
|
| 52 |
+
"下雨天我会趴在窗台上数雨滴,数到一百就睡着了。"
|
| 53 |
+
"我最喜欢的动画片是《猫娘乐园》,每次看都会流泪,太感动了。"
|
| 54 |
+
"如果你觉得我可爱,请给我点赞,喵~"
|
| 55 |
+
)
|
| 56 |
+
|
| 57 |
+
chars = sorted(list(set(text)))
|
| 58 |
+
char2idx = {ch: i for i, ch in enumerate(chars)}
|
| 59 |
+
idx2char = {i: ch for ch, i in char2idx.items()}
|
| 60 |
+
vocab_size = len(chars)
|
| 61 |
+
print(f"字符集大小: {vocab_size}")
|
| 62 |
+
print(f"总语料长度: {len(text)} 字符")
|
| 63 |
+
|
| 64 |
+
# ===================== 2. 双层 LSTM 模型(参数量 ~40 万) =====================
|
| 65 |
+
class CatgirlLSTM(nn.Module):
|
| 66 |
+
def __init__(self, vocab_size, embed_size=128, hidden_size=256, num_layers=2, dropout=0.3):
|
| 67 |
+
super().__init__()
|
| 68 |
+
self.embedding = nn.Embedding(vocab_size, embed_size)
|
| 69 |
+
self.lstm = nn.LSTM(embed_size, hidden_size, num_layers,
|
| 70 |
+
batch_first=True, dropout=dropout)
|
| 71 |
+
self.fc = nn.Linear(hidden_size, vocab_size)
|
| 72 |
+
|
| 73 |
+
def forward(self, x, hidden=None):
|
| 74 |
+
x = self.embedding(x)
|
| 75 |
+
out, hidden = self.lstm(x, hidden)
|
| 76 |
+
out = self.fc(out)
|
| 77 |
+
return out, hidden
|
| 78 |
+
|
| 79 |
+
embed_size = 128
|
| 80 |
+
hidden_size = 256
|
| 81 |
+
num_layers = 2
|
| 82 |
+
dropout = 0.3
|
| 83 |
+
|
| 84 |
+
model = CatgirlLSTM(vocab_size, embed_size, hidden_size, num_layers, dropout)
|
| 85 |
+
total_params = sum(p.numel() for p in model.parameters())
|
| 86 |
+
print(f"模型参数量: {total_params:,}")
|
| 87 |
+
|
| 88 |
+
# ===================== 3. 训练准备 =====================
|
| 89 |
+
data = torch.tensor([char2idx[ch] for ch in text], dtype=torch.long)
|
| 90 |
+
seq_len = 128 # 序列长度(可调,手机内存足够)
|
| 91 |
+
epochs = 300
|
| 92 |
+
save_interval = 20
|
| 93 |
+
|
| 94 |
+
optimizer = torch.optim.Adam(model.parameters(), lr=0.005)
|
| 95 |
+
loss_fn = nn.CrossEntropyLoss()
|
| 96 |
+
os.makedirs("checkpoints_lstm", exist_ok=True)
|
| 97 |
+
|
| 98 |
+
def generate(model, start_char='你', length=250, temperature=0.8):
|
| 99 |
+
model.eval()
|
| 100 |
+
with torch.no_grad():
|
| 101 |
+
if start_char not in char2idx:
|
| 102 |
+
start_char = random.choice(list(char2idx.keys()))
|
| 103 |
+
input_idx = torch.tensor([[char2idx[start_char]]])
|
| 104 |
+
hidden = None
|
| 105 |
+
result = [start_char]
|
| 106 |
+
for _ in range(length):
|
| 107 |
+
logits, hidden = model(input_idx, hidden)
|
| 108 |
+
probs = torch.softmax(logits[0, -1] / temperature, dim=0).cpu().numpy()
|
| 109 |
+
next_idx = np.random.choice(len(probs), p=probs)
|
| 110 |
+
next_char = idx2char[next_idx]
|
| 111 |
+
result.append(next_char)
|
| 112 |
+
input_idx = torch.tensor([[next_idx]])
|
| 113 |
+
return ''.join(result)
|
| 114 |
+
|
| 115 |
+
# ===================== 4. 训练循环 =====================
|
| 116 |
+
print("\n开始训练双层 LSTM 春梦蝶模型(参数量 {:,},300轮)...\n".format(total_params))
|
| 117 |
+
start_total = time.time()
|
| 118 |
+
|
| 119 |
+
for epoch in range(1, epochs + 1):
|
| 120 |
+
epoch_start = time.time()
|
| 121 |
+
hidden = None
|
| 122 |
+
total_loss = 0
|
| 123 |
+
n_batches = 0
|
| 124 |
+
step = seq_len // 2
|
| 125 |
+
for i in range(0, len(data) - seq_len, step):
|
| 126 |
+
x = data[i:i+seq_len].unsqueeze(0)
|
| 127 |
+
y = data[i+1:i+seq_len+1].unsqueeze(0)
|
| 128 |
+
|
| 129 |
+
logits, hidden = model(x, hidden)
|
| 130 |
+
if hidden is not None:
|
| 131 |
+
hidden = (hidden[0].detach(), hidden[1].detach())
|
| 132 |
+
|
| 133 |
+
loss = loss_fn(logits.view(-1, vocab_size), y.view(-1))
|
| 134 |
+
optimizer.zero_grad()
|
| 135 |
+
loss.backward()
|
| 136 |
+
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
| 137 |
+
optimizer.step()
|
| 138 |
+
|
| 139 |
+
total_loss += loss.item()
|
| 140 |
+
n_batches += 1
|
| 141 |
+
|
| 142 |
+
avg_loss = total_loss / n_batches
|
| 143 |
+
epoch_time = time.time() - epoch_start
|
| 144 |
+
|
| 145 |
+
if epoch % save_interval == 0:
|
| 146 |
+
sample = generate(model, start_char='我', length=200, temperature=0.7)
|
| 147 |
+
print(f"Epoch {epoch:3d}/{epochs} | Loss: {avg_loss:.4f} | Time: {epoch_time:.2f}s")
|
| 148 |
+
print(f"春梦蝶: {sample[:180]}...\n")
|
| 149 |
+
torch.save(model.state_dict(), f"checkpoints_lstm/mengdie_lstm_epoch_{epoch}.pth")
|
| 150 |
+
print(f"已保存模型到 checkpoints_lstm/mengdie_lstm_epoch_{epoch}.pth\n")
|
| 151 |
+
else:
|
| 152 |
+
if epoch % 10 == 0:
|
| 153 |
+
print(f"Epoch {epoch:3d}/{epochs} | Loss: {avg_loss:.4f} | Time: {epoch_time:.2f}s")
|
| 154 |
+
|
| 155 |
+
total_time = time.time() - start_total
|
| 156 |
+
print(f"\n训练完成!总耗时: {total_time:.2f} 秒 ({total_time/60:.1f} 分钟)")
|
| 157 |
+
final_path = "checkpoints_lstm/mengdie_lstm_final.pth"
|
| 158 |
+
torch.save(model.state_dict(), final_path)
|
| 159 |
+
with open("checkpoints_lstm/char2idx.json", "w", encoding="utf-8") as f:
|
| 160 |
+
json.dump(char2idx, f, ensure_ascii=False)
|
| 161 |
+
|
| 162 |
+
print("\n=== 最终生成(温度0.7) ===")
|
| 163 |
+
print(generate(model, start_char='我', length=400, temperature=0.7))
|
| 164 |
+
print("\n=== 温度1.1 随机版 ===")
|
| 165 |
+
print(generate(model, start_char='喵', length=400, temperature=1.1))
|
v3_lstm/checkpoints_lstm/char2idx.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{" ": 0, "!": 1, ",": 2, ".": 3, "0": 4, "2": 5, "A": 6, "I": 7, "N": 8, "P": 9, "R": 10, "T": 11, "a": 12, "b": 13, "c": 14, "d": 15, "e": 16, "f": 17, "g": 18, "h": 19, "i": 20, "l": 21, "m": 22, "n": 23, "o": 24, "p": 25, "r": 26, "s": 27, "t": 28, "u": 29, "v": 30, "w": 31, "y": 32, "‘": 33, "’": 34, "、": 35, "。": 36, "《": 37, "》": 38, "一": 39, "三": 40, "上": 41, "下": 42, "不": 43, "且": 44, "两": 45, "个": 46, "为": 47, "久": 48, "么": 49, "乐": 50, "也": 51, "习": 52, "买": 53, "乱": 54, "了": 55, "事": 56, "二": 57, "于": 58, "亮": 59, "亲": 60, "人": 61, "什": 62, "今": 63, "代": 64, "以": 65, "们": 66, "优": 67, "会": 68, "但": 69, "你": 70, "候": 71, "做": 72, "像": 73, "光": 74, "六": 75, "其": 76, "写": 77, "冰": 78, "凌": 79, "几": 80, "出": 81, "分": 82, "别": 83, "到": 84, "制": 85, "力": 86, "动": 87, "化": 88, "区": 89, "十": 90, "卖": 91, "卷": 92, "厉": 93, "厌": 94, "去": 95, "友": 96, "双": 97, "反": 98, "发": 99, "取": 100, "变": 101, "只": 102, "叫": 103, "可": 104, "台": 105, "叹": 106, "吃": 107, "名": 108, "后": 109, "吗": 110, "吧": 111, "听": 112, "呀": 113, "呢": 114, "和": 115, "咬": 116, "哦": 117, "啊": 118, "啦": 119, "喜": 120, "喝": 121, "喵": 122, "嘿": 123, "因": 124, "园": 125, "困": 126, "图": 127, "在": 128, "地": 129, "型": 130, "壁": 131, "声": 132, "大": 133, "天": 134, "太": 135, "失": 136, "头": 137, "夸": 138, "奇": 139, "奶": 140, "好": 141, "如": 142, "娘": 143, "子": 144, "字": 145, "学": 146, "孩": 147, "它": 148, "实": 149, "害": 150, "家": 151, "密": 152, "对": 153, "小": 154, "少": 155, "尖": 156, "尝": 157, "就": 158, "尾": 159, "居": 160, "屏": 161, "岁": 162, "工": 163, "己": 164, "已": 165, "巴": 166, "师": 167, "幕": 168, "干": 169, "平": 170, "年": 171, "应": 172, "度": 173, "座": 174, "开": 175, "弄": 176, "强": 177, "当": 178, "待": 179, "很": 180, "得": 181, "微": 182, "心": 183, "情": 184, "想": 185, "意": 186, "感": 187, "愿": 188, "慢": 189, "成": 190, "我": 191, "戴": 192, "手": 193, "打": 194, "找": 195, "把": 196, "抓": 197, "拟": 198, "挠": 199, "捉": 200, "控": 201, "摸": 202, "撮": 203, "数": 204, "整": 205, "文": 206, "方": 207, "旅": 208, "无": 209, "早": 210, "时": 211, "旺": 212, "明": 213, "星": 214, "春": 215, "昨": 216, "是": 217, "晒": 218, "晚": 219, "晨": 220, "晴": 221, "暗": 222, "最": 223, "月": 224, "有": 225, "朋": 226, "朵": 227, "机": 228, "条": 229, "来": 230, "板": 231, "果": 232, "柔": 233, "标": 234, "样": 235, "梦": 236, "梳": 237, "模": 238, "樱": 239, "橙": 240, "次": 241, "欢": 242, "欺": 243, "正": 244, "死": 245, "每": 246, "比": 247, "毛": 248, "气": 249, "水": 250, "永": 251, "汤": 252, "没": 253, "法": 254, "泪": 255, "泼": 256, "洒": 257, "洗": 258, "活": 259, "流": 260, "消": 261, "淇": 262, "淋": 263, "深": 264, "湿": 265, "滴": 266, "澡": 267, "灯": 268, "灵": 269, "点": 270, "烦": 271, "热": 272, "然": 273, "照": 274, "熟": 275, "爪": 276, "爱": 277, "片": 278, "牛": 279, "特": 280, "狗": 281, "猫": 282, "玩": 283, "球": 284, "瓜": 285, "用": 286, "画": 287, "白": 288, "百": 289, "的": 290, "盘": 291, "盛": 292, "相": 293, "盹": 294, "盾": 295, "看": 296, "真": 297, "眼": 298, "着": 299, "睛": 300, "睡": 301, "瞌": 302, "矛": 303, "知": 304, "码": 305, "神": 306, "秘": 307, "积": 308, "程": 309, "穷": 310, "窗": 311, "第": 312, "算": 313, "管": 314, "类": 315, "红": 316, "级": 317, "纪": 318, "练": 319, "经": 320, "结": 321, "给": 322, "络": 323, "网": 324, "美": 325, "翘": 326, "老": 327, "而": 328, "耳": 329, "聊": 330, "聪": 331, "能": 332, "脾": 333, "腻": 334, "自": 335, "至": 336, "色": 337, "花": 338, "苦": 339, "英": 340, "茸": 341, "落": 342, "藏": 343, "虚": 344, "虽": 345, "蝴": 346, "蝶": 347, "行": 348, "要": 349, "见": 350, "视": 351, "觉": 352, "计": 353, "讨": 354, "让": 355, "训": 356, "词": 357, "试": 358, "诗": 359, "说": 360, "请": 361, "谁": 362, "负": 363, "赞": 364, "走": 365, "起": 366, "超": 367, "趣": 368, "趴": 369, "跟": 370, "路": 371, "软": 372, "轻": 373, "输": 374, "还": 375, "这": 376, "远": 377, "迷": 378, "追": 379, "送": 380, "适": 381, "速": 382, "遇": 383, "道": 384, "那": 385, "邻": 386, "都": 387, "里": 388, "重": 389, "量": 390, "键": 391, "镜": 392, "长": 393, "闪": 394, "问": 395, "闹": 396, "阳": 397, "院": 398, "除": 399, "隔": 400, "雨": 401, "雪": 402, "需": 403, "霞": 404, "静": 405, "音": 406, "颗": 407, "颜": 408, "首": 409, "鱼": 410, "麻": 411, "黑": 412, "鼠": 413, "龄": 414, "!": 415, ",": 416, ":": 417, "?": 418, "~": 419}
|
v3_lstm/checkpoints_lstm/mengdie_lstm_final.pth
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:7939323371bb1d0b43b290c7ea4b49fe4d606f3ecc84cea21f8fe901b712fd13
|
| 3 |
+
size 4337725
|
v3_lstm/run.py
ADDED
|
@@ -0,0 +1,53 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import numpy as np
|
| 4 |
+
import json
|
| 5 |
+
import random
|
| 6 |
+
|
| 7 |
+
with open("checkpoints_lstm/char2idx.json", "r", encoding="utf-8") as f:
|
| 8 |
+
char2idx = json.load(f)
|
| 9 |
+
idx2char = {int(v): k for k, v in char2idx.items()}
|
| 10 |
+
vocab_size = len(char2idx)
|
| 11 |
+
|
| 12 |
+
class CatgirlLSTM(nn.Module):
|
| 13 |
+
def __init__(self, vocab_size, embed_size=128, hidden_size=256, num_layers=2, dropout=0.3):
|
| 14 |
+
super().__init__()
|
| 15 |
+
self.embedding = nn.Embedding(vocab_size, embed_size)
|
| 16 |
+
self.lstm = nn.LSTM(embed_size, hidden_size, num_layers, batch_first=True, dropout=dropout)
|
| 17 |
+
self.fc = nn.Linear(hidden_size, vocab_size)
|
| 18 |
+
def forward(self, x, hidden=None):
|
| 19 |
+
x = self.embedding(x)
|
| 20 |
+
out, hidden = self.lstm(x, hidden)
|
| 21 |
+
out = self.fc(out)
|
| 22 |
+
return out, hidden
|
| 23 |
+
|
| 24 |
+
model = CatgirlLSTM(vocab_size)
|
| 25 |
+
model.load_state_dict(torch.load("checkpoints_lstm/mengdie_lstm_final.pth", map_location='cpu'))
|
| 26 |
+
model.eval()
|
| 27 |
+
print("春梦蝶 LSTM 大模型加载成功!喵~\n")
|
| 28 |
+
|
| 29 |
+
def generate_response(prompt, length=2000, temperature=0.8):
|
| 30 |
+
if not prompt:
|
| 31 |
+
prompt = random.choice(list(char2idx.keys()))
|
| 32 |
+
indices = [char2idx.get(ch, random.choice(list(char2idx.values()))) for ch in prompt]
|
| 33 |
+
input_tensor = torch.tensor([indices])
|
| 34 |
+
hidden = None
|
| 35 |
+
result = list(prompt)
|
| 36 |
+
with torch.no_grad():
|
| 37 |
+
for _ in range(length):
|
| 38 |
+
logits, hidden = model(input_tensor, hidden)
|
| 39 |
+
probs = torch.softmax(logits[0, -1] / temperature, dim=0).cpu().numpy()
|
| 40 |
+
next_idx = np.random.choice(len(probs), p=probs)
|
| 41 |
+
next_char = idx2char[next_idx]
|
| 42 |
+
result.append(next_char)
|
| 43 |
+
input_tensor = torch.tensor([[next_idx]])
|
| 44 |
+
return ''.join(result)
|
| 45 |
+
|
| 46 |
+
print("开始对话(输入 q 退出)")
|
| 47 |
+
while True:
|
| 48 |
+
user = input("\n你: ")
|
| 49 |
+
if user.lower() == 'q':
|
| 50 |
+
break
|
| 51 |
+
start = user[-5:] if len(user) >= 5 else user
|
| 52 |
+
reply = generate_response(start, length=180, temperature=0.85)
|
| 53 |
+
print(f"春梦蝶: {reply}")
|
v4_transformer/6.py
ADDED
|
@@ -0,0 +1,213 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import torch.optim as optim
|
| 4 |
+
from torch.utils.data import Dataset, DataLoader
|
| 5 |
+
import numpy as np
|
| 6 |
+
import random
|
| 7 |
+
import json
|
| 8 |
+
import os
|
| 9 |
+
from tqdm import tqdm
|
| 10 |
+
|
| 11 |
+
# ===================== 1. 自动生成春梦蝶对话数据 =====================
|
| 12 |
+
USER_TEMPLATES = [
|
| 13 |
+
"你好", "嗨", "Hi", "在吗", "你好呀",
|
| 14 |
+
"你叫什么名字?", "你叫什么", "你的名字是?", "请问你叫什么",
|
| 15 |
+
"你几岁了?", "你多大了", "年龄", "今年几岁",
|
| 16 |
+
"你是什么星座?", "星座", "双子座吗",
|
| 17 |
+
"你的眼睛是什么颜色?", "眼睛颜色", "你的眼睛好红",
|
| 18 |
+
"你的头发是白色的吗?", "白发", "头发颜色",
|
| 19 |
+
"你最喜欢吃什么?", "喜欢吃什么", "小鱼干", "冰淇淋",
|
| 20 |
+
"你会做什么?", "你会什么", "技能", "会写代码吗",
|
| 21 |
+
"你的梦想是什么?", "梦想", "想成为什么",
|
| 22 |
+
"你好可爱", "你真可爱", "萌", "喵",
|
| 23 |
+
"摸摸头", "摸头", "拍头",
|
| 24 |
+
"bye", "再见", "拜拜", "下次见"
|
| 25 |
+
]
|
| 26 |
+
|
| 27 |
+
BOT_RESPONSES = {
|
| 28 |
+
"你好": ["你好呀~", "嗨,我是春梦蝶", "喵~你好", "你好,今天心情很好呢"],
|
| 29 |
+
"你叫什么名字?": ["我叫春梦蝶,是一只红瞳白发的猫娘~", "春梦蝶!你可以叫我小蝶", "名字是春梦蝶,喵~"],
|
| 30 |
+
"你几岁了?": ["我今年十一岁", "十一岁,正是可爱的年纪", "十一岁啦,双子座的小猫娘"],
|
| 31 |
+
"你是什么星座?": ["双子座~有时活泼有时安静", "双子座哦,你也是吗?"],
|
| 32 |
+
"你的眼睛是什么颜色?": ["红色,像红宝石一样", "红瞳,晚上会微微发光呢", "红色的眼睛,很特别吧"],
|
| 33 |
+
"你的头发是白色的吗?": ["是的,纯白色的头发", "白发,像雪一样", "白色长发,摸起来很软哦"],
|
| 34 |
+
"你最喜欢吃什么?": ["小鱼干!还有冰淇淋", "小鱼干是第一名,冰淇淋第二名", "最喜欢吃小鱼干,喵~"],
|
| 35 |
+
"你会做什么?": ["我会用Python写代码,还会训练神经网络", "会深度学习,还会喵喵叫", "我会训练AI模型,很厉害吧"],
|
| 36 |
+
"你的梦想是什么?": ["成为一名厉害的AI工程师", "想做出自己的AI,然后去猫星球旅行", "梦想是成为AI专家,顺便吃遍所有小鱼干"],
|
| 37 |
+
"你好可爱": ["喵~谢谢", "嘿嘿,你也很可爱", "被夸了,好开心"],
|
| 38 |
+
"摸摸头": ["喵~好舒服", "再摸摸嘛", "好温暖,喜欢被摸头"],
|
| 39 |
+
"bye": ["再见喵~", "下次再聊", "拜拜,记得想我哦"]
|
| 40 |
+
}
|
| 41 |
+
|
| 42 |
+
# 生成随机对话对(用户输入 -> 机器人回复)
|
| 43 |
+
def generate_dialogue_pairs(num_pairs=3000):
|
| 44 |
+
pairs = []
|
| 45 |
+
for _ in range(num_pairs):
|
| 46 |
+
# 随机选一个用户模板
|
| 47 |
+
user_input = random.choice(USER_TEMPLATES)
|
| 48 |
+
# 根据用户输入的关键词选择合适的回复列表
|
| 49 |
+
matched_responses = []
|
| 50 |
+
for key in BOT_RESPONSES:
|
| 51 |
+
if key in user_input or (key == "你好" and user_input in ["嗨", "Hi", "在吗"]):
|
| 52 |
+
matched_responses.extend(BOT_RESPONSES[key])
|
| 53 |
+
if not matched_responses:
|
| 54 |
+
# 默认回复
|
| 55 |
+
matched_responses = ["喵~", "嗯?", "你说什么?", "好呀", "我不太懂,但我会努力学习"]
|
| 56 |
+
bot_response = random.choice(matched_responses)
|
| 57 |
+
pairs.append((user_input, bot_response))
|
| 58 |
+
# 去重并保证多样性
|
| 59 |
+
pairs = list(set(pairs))
|
| 60 |
+
return pairs
|
| 61 |
+
|
| 62 |
+
# 构建字符级词表(中英文混合)
|
| 63 |
+
def build_vocab(pairs):
|
| 64 |
+
all_text = ""
|
| 65 |
+
for user, bot in pairs:
|
| 66 |
+
all_text += user + bot
|
| 67 |
+
chars = sorted(list(set(all_text)))
|
| 68 |
+
char2idx = {ch: i for i, ch in enumerate(chars)}
|
| 69 |
+
idx2char = {i: ch for ch, i in char2idx.items()}
|
| 70 |
+
return char2idx, idx2char, len(chars)
|
| 71 |
+
|
| 72 |
+
# 将文本转为索引序列(固定长度,padding)
|
| 73 |
+
def text_to_indices(text, char2idx, max_len):
|
| 74 |
+
indices = [char2idx.get(ch, 0) for ch in text] # 未知字符用0
|
| 75 |
+
if len(indices) < max_len:
|
| 76 |
+
indices += [0] * (max_len - len(indices))
|
| 77 |
+
else:
|
| 78 |
+
indices = indices[:max_len]
|
| 79 |
+
return indices
|
| 80 |
+
|
| 81 |
+
# 数据集类
|
| 82 |
+
class ChatDataset(Dataset):
|
| 83 |
+
def __init__(self, pairs, char2idx, max_len=32):
|
| 84 |
+
self.pairs = pairs
|
| 85 |
+
self.char2idx = char2idx
|
| 86 |
+
self.max_len = max_len
|
| 87 |
+
def __len__(self):
|
| 88 |
+
return len(self.pairs)
|
| 89 |
+
def __getitem__(self, idx):
|
| 90 |
+
user, bot = self.pairs[idx]
|
| 91 |
+
user_ids = text_to_indices(user, self.char2idx, self.max_len)
|
| 92 |
+
bot_ids = text_to_indices(bot, self.char2idx, self.max_len)
|
| 93 |
+
return torch.tensor(user_ids, dtype=torch.long), torch.tensor(bot_ids, dtype=torch.long)
|
| 94 |
+
|
| 95 |
+
# ===================== 2. 极简 Transformer 模型(约 500 万参数) =====================
|
| 96 |
+
class TinyTransformerChat(nn.Module):
|
| 97 |
+
def __init__(self, vocab_size, d_model=256, nhead=8, num_encoder_layers=2, num_decoder_layers=2, dim_feedforward=512, max_len=32):
|
| 98 |
+
super().__init__()
|
| 99 |
+
self.d_model = d_model
|
| 100 |
+
self.embedding = nn.Embedding(vocab_size, d_model)
|
| 101 |
+
self.pos_encoding = nn.Parameter(torch.zeros(1, max_len, d_model))
|
| 102 |
+
|
| 103 |
+
encoder_layer = nn.TransformerEncoderLayer(d_model, nhead, dim_feedforward, batch_first=True)
|
| 104 |
+
self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_encoder_layers)
|
| 105 |
+
|
| 106 |
+
decoder_layer = nn.TransformerDecoderLayer(d_model, nhead, dim_feedforward, batch_first=True)
|
| 107 |
+
self.transformer_decoder = nn.TransformerDecoder(decoder_layer, num_decoder_layers)
|
| 108 |
+
|
| 109 |
+
self.fc_out = nn.Linear(d_model, vocab_size)
|
| 110 |
+
|
| 111 |
+
def forward(self, src, tgt, src_mask=None, tgt_mask=None):
|
| 112 |
+
# src: (batch, src_len), tgt: (batch, tgt_len)
|
| 113 |
+
src_emb = self.embedding(src) * (self.d_model ** 0.5) + self.pos_encoding[:, :src.size(1), :]
|
| 114 |
+
tgt_emb = self.embedding(tgt) * (self.d_model ** 0.5) + self.pos_encoding[:, :tgt.size(1), :]
|
| 115 |
+
|
| 116 |
+
memory = self.transformer_encoder(src_emb, src_mask)
|
| 117 |
+
output = self.transformer_decoder(tgt_emb, memory, tgt_mask)
|
| 118 |
+
logits = self.fc_out(output)
|
| 119 |
+
return logits
|
| 120 |
+
|
| 121 |
+
def generate_causal_mask(size):
|
| 122 |
+
mask = torch.triu(torch.ones(size, size), diagonal=1).bool()
|
| 123 |
+
return mask
|
| 124 |
+
|
| 125 |
+
# ===================== 3. 训练配置 =====================
|
| 126 |
+
def train():
|
| 127 |
+
print("生成对话数据...")
|
| 128 |
+
pairs = generate_dialogue_pairs(5000) # 生成5000条
|
| 129 |
+
print(f"生成 {len(pairs)} 条对话对")
|
| 130 |
+
|
| 131 |
+
char2idx, idx2char, vocab_size = build_vocab(pairs)
|
| 132 |
+
print(f"词表大小: {vocab_size}")
|
| 133 |
+
|
| 134 |
+
# 保存词表供后续使用
|
| 135 |
+
with open("chat_vocab.json", "w", encoding="utf-8") as f:
|
| 136 |
+
json.dump(char2idx, f, ensure_ascii=False)
|
| 137 |
+
|
| 138 |
+
max_len = 32
|
| 139 |
+
dataset = ChatDataset(pairs, char2idx, max_len)
|
| 140 |
+
dataloader = DataLoader(dataset, batch_size=64, shuffle=True)
|
| 141 |
+
|
| 142 |
+
device = torch.device("cpu")
|
| 143 |
+
model = TinyTransformerChat(vocab_size, d_model=256, nhead=8, num_encoder_layers=2, num_decoder_layers=2, dim_feedforward=512, max_len=max_len)
|
| 144 |
+
model.to(device)
|
| 145 |
+
|
| 146 |
+
total_params = sum(p.numel() for p in model.parameters())
|
| 147 |
+
print(f"模型参数量: {total_params:,}")
|
| 148 |
+
|
| 149 |
+
criterion = nn.CrossEntropyLoss(ignore_index=0)
|
| 150 |
+
optimizer = optim.Adam(model.parameters(), lr=0.001)
|
| 151 |
+
|
| 152 |
+
epochs = 200
|
| 153 |
+
print("开始训练...")
|
| 154 |
+
for epoch in range(1, epochs+1):
|
| 155 |
+
model.train()
|
| 156 |
+
total_loss = 0
|
| 157 |
+
for src, tgt in dataloader:
|
| 158 |
+
src = src.to(device)
|
| 159 |
+
tgt = tgt.to(device)
|
| 160 |
+
tgt_input = tgt[:, :-1]
|
| 161 |
+
tgt_output = tgt[:, 1:]
|
| 162 |
+
|
| 163 |
+
# 生成因果掩码
|
| 164 |
+
tgt_mask = generate_causal_mask(tgt_input.size(1)).to(device)
|
| 165 |
+
|
| 166 |
+
logits = model(src, tgt_input, tgt_mask=tgt_mask)
|
| 167 |
+
loss = criterion(logits.reshape(-1, vocab_size), tgt_output.reshape(-1))
|
| 168 |
+
|
| 169 |
+
optimizer.zero_grad()
|
| 170 |
+
loss.backward()
|
| 171 |
+
optimizer.step()
|
| 172 |
+
total_loss += loss.item()
|
| 173 |
+
|
| 174 |
+
avg_loss = total_loss / len(dataloader)
|
| 175 |
+
if epoch % 20 == 0:
|
| 176 |
+
print(f"Epoch {epoch:3d}/{epochs} | Loss: {avg_loss:.4f}")
|
| 177 |
+
# 保存检查点
|
| 178 |
+
torch.save(model.state_dict(), f"chat_model_epoch_{epoch}.pth")
|
| 179 |
+
|
| 180 |
+
# 保存最终模型
|
| 181 |
+
torch.save(model.state_dict(), "chat_model_final.pth")
|
| 182 |
+
print("训练完成!模型已保存。")
|
| 183 |
+
|
| 184 |
+
# 简单测试
|
| 185 |
+
test_model(model, char2idx, idx2char, device, max_len)
|
| 186 |
+
|
| 187 |
+
def test_model(model, char2idx, idx2char, device, max_len):
|
| 188 |
+
model.eval()
|
| 189 |
+
print("\n=== 测试对话 ===")
|
| 190 |
+
while True:
|
| 191 |
+
user = input("你: ").strip()
|
| 192 |
+
if user.lower() == 'q':
|
| 193 |
+
break
|
| 194 |
+
# 编码用户输入
|
| 195 |
+
src = text_to_indices(user, char2idx, max_len)
|
| 196 |
+
src_tensor = torch.tensor([src], dtype=torch.long).to(device)
|
| 197 |
+
tgt = torch.tensor([[0]], dtype=torch.long).to(device)
|
| 198 |
+
generated = []
|
| 199 |
+
for _ in range(64): # 最多生成64个字符
|
| 200 |
+
tgt_mask = generate_causal_mask(tgt.size(1)).to(device)
|
| 201 |
+
logits = model(src_tensor, tgt, tgt_mask=tgt_mask)
|
| 202 |
+
next_token_logits = logits[0, -1, :]
|
| 203 |
+
probs = torch.softmax(next_token_logits / 0.8, dim=0).cpu().numpy()
|
| 204 |
+
next_token = np.random.choice(len(probs), p=probs)
|
| 205 |
+
if next_token == 0: # 结束符
|
| 206 |
+
break
|
| 207 |
+
generated.append(idx2char[next_token])
|
| 208 |
+
tgt = torch.cat([tgt, torch.tensor([[next_token]], dtype=torch.long).to(device)], dim=1)
|
| 209 |
+
reply = ''.join(generated)
|
| 210 |
+
print(f"春梦蝶: {reply}")
|
| 211 |
+
|
| 212 |
+
if __name__ == "__main__":
|
| 213 |
+
train()
|
v4_transformer/README.md
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
## v4_transformer 说明
|
| 2 |
+
|
| 3 |
+
**⚠️ 本版本仅提供训练和推理代码,未提供训练好的权重。**
|
| 4 |
+
|
| 5 |
+
该版本在手机 CPU 环境下训练时输出效果不佳(生成乱码),故未上传 `.pth` 文件。
|
| 6 |
+
|
| 7 |
+
如果你感兴趣,可以:
|
| 8 |
+
1. 在电脑或服务器上运行 `train.py` 重新训练
|
| 9 |
+
2. 或参考本代码学习 Transformer 对话模型的实现思路
|
| 10 |
+
|
| 11 |
+
代码本身经过了完整的语法测试,结构清晰,适合作为学习材料。
|
v4_transformer/chat_vocab.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{"A": 0, "H": 1, "I": 2, "P": 3, "b": 4, "e": 5, "h": 6, "i": 7, "n": 8, "o": 9, "t": 10, "y": 11, "一": 12, "上": 13, "下": 14, "不": 15, "专": 16, "为": 17, "么": 18, "也": 19, "习": 20, "了": 21, "二": 22, "什": 23, "今": 24, "代": 25, "以": 26, "会": 27, "但": 28, "你": 29, "便": 30, "做": 31, "像": 32, "光": 33, "再": 34, "写": 35, "冰": 36, "几": 37, "出": 38, "别": 39, "力": 40, "努": 41, "十": 42, "厉": 43, "去": 44, "双": 45, "发": 46, "只": 47, "叫": 48, "可": 49, "吃": 50, "名": 51, "后": 52, "吗": 53, "吧": 54, "呀": 55, "呢": 56, "哦": 57, "啦": 58, "喜": 59, "喵": 60, "嗨": 61, "嗯": 62, "嘛": 63, "嘿": 64, "在": 65, "型": 66, "多": 67, "大": 68, "天": 69, "太": 70, "头": 71, "夸": 72, "好": 73, "娘": 74, "子": 75, "字": 76, "学": 77, "安": 78, "宝": 79, "害": 80, "家": 81, "小": 82, "岁": 83, "工": 84, "己": 85, "师": 86, "干": 87, "年": 88, "度": 89, "座": 90, "开": 91, "很": 92, "得": 93, "微": 94, "心": 95, "情": 96, "想": 97, "懂": 98, "成": 99, "我": 100, "所": 101, "技": 102, "拍": 103, "拜": 104, "摸": 105, "旅": 106, "时": 107, "星": 108, "春": 109, "是": 110, "晚": 111, "暖": 112, "最": 113, "有": 114, "服": 115, "来": 116, "样": 117, "梦": 118, "模": 119, "次": 120, "欢": 121, "正": 122, "泼": 123, "活": 124, "淇": 125, "淋": 126, "深": 127, "温": 128, "然": 129, "爱": 130, "特": 131, "猫": 132, "球": 133, "用": 134, "白": 135, "的": 136, "真": 137, "眼": 138, "睛": 139, "瞳": 140, "石": 141, "码": 142, "神": 143, "程": 144, "第": 145, "红": 146, "纪": 147, "纯": 148, "练": 149, "经": 150, "络": 151, "网": 152, "聊": 153, "能": 154, "自": 155, "舒": 156, "色": 157, "萌": 158, "蝶": 159, "行": 160, "被": 161, "见": 162, "训": 163, "记": 164, "说": 165, "请": 166, "谢": 167, "起": 168, "软": 169, "还": 170, "遍": 171, "长": 172, "问": 173, "雪": 174, "静": 175, "顺": 176, "颜": 177, "鱼": 178, "龄": 179, "!": 180, ",": 181, "?": 182, "~": 183}
|
v4_transformer/run.py
ADDED
|
@@ -0,0 +1,86 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import json
|
| 4 |
+
import numpy as np
|
| 5 |
+
|
| 6 |
+
class TinyTransformerChat(nn.Module):
|
| 7 |
+
def __init__(self, vocab_size, d_model=256, nhead=8, num_encoder_layers=2,
|
| 8 |
+
num_decoder_layers=2, dim_feedforward=512, max_len=32):
|
| 9 |
+
super().__init__()
|
| 10 |
+
self.d_model = d_model
|
| 11 |
+
self.max_len = max_len
|
| 12 |
+
self.embedding = nn.Embedding(vocab_size, d_model)
|
| 13 |
+
self.pos_encoding = nn.Parameter(torch.zeros(1, max_len, d_model))
|
| 14 |
+
encoder_layer = nn.TransformerEncoderLayer(d_model, nhead, dim_feedforward, batch_first=True)
|
| 15 |
+
self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_encoder_layers)
|
| 16 |
+
decoder_layer = nn.TransformerDecoderLayer(d_model, nhead, dim_feedforward, batch_first=True)
|
| 17 |
+
self.transformer_decoder = nn.TransformerDecoder(decoder_layer, num_decoder_layers)
|
| 18 |
+
self.fc_out = nn.Linear(d_model, vocab_size)
|
| 19 |
+
|
| 20 |
+
def forward(self, src, tgt, src_mask=None, tgt_mask=None):
|
| 21 |
+
src_len = src.size(1)
|
| 22 |
+
tgt_len = tgt.size(1)
|
| 23 |
+
if src_len > self.max_len or tgt_len > self.max_len:
|
| 24 |
+
src = src[:, :self.max_len]
|
| 25 |
+
tgt = tgt[:, :self.max_len]
|
| 26 |
+
src_len = min(src_len, self.max_len)
|
| 27 |
+
tgt_len = min(tgt_len, self.max_len)
|
| 28 |
+
src_emb = self.embedding(src) * (self.d_model ** 0.5) + self.pos_encoding[:, :src_len, :]
|
| 29 |
+
tgt_emb = self.embedding(tgt) * (self.d_model ** 0.5) + self.pos_encoding[:, :tgt_len, :]
|
| 30 |
+
memory = self.transformer_encoder(src_emb, src_mask)
|
| 31 |
+
output = self.transformer_decoder(tgt_emb, memory, tgt_mask)
|
| 32 |
+
logits = self.fc_out(output)
|
| 33 |
+
return logits
|
| 34 |
+
|
| 35 |
+
def generate_causal_mask(size):
|
| 36 |
+
mask = torch.triu(torch.ones(size, size), diagonal=1).bool()
|
| 37 |
+
return mask
|
| 38 |
+
|
| 39 |
+
def load_model_and_vocab(model_path="chat_model_final.pth", vocab_path="chat_vocab.json"):
|
| 40 |
+
with open(vocab_path, "r", encoding="utf-8") as f:
|
| 41 |
+
char2idx = json.load(f)
|
| 42 |
+
idx2char = {int(v): k for k, v in char2idx.items()}
|
| 43 |
+
vocab_size = len(char2idx)
|
| 44 |
+
model = TinyTransformerChat(vocab_size)
|
| 45 |
+
model.load_state_dict(torch.load(model_path, map_location="cpu"))
|
| 46 |
+
model.eval()
|
| 47 |
+
return model, char2idx, idx2char
|
| 48 |
+
|
| 49 |
+
def text_to_indices(text, char2idx, max_len=32):
|
| 50 |
+
indices = [char2idx.get(ch, 0) for ch in text]
|
| 51 |
+
if len(indices) < max_len:
|
| 52 |
+
indices += [0] * (max_len - len(indices))
|
| 53 |
+
else:
|
| 54 |
+
indices = indices[:max_len]
|
| 55 |
+
return indices
|
| 56 |
+
|
| 57 |
+
def generate_response(model, user_input, char2idx, idx2char, max_len=32, temperature=0.8):
|
| 58 |
+
src = text_to_indices(user_input, char2idx, max_len)
|
| 59 |
+
src_tensor = torch.tensor([src], dtype=torch.long)
|
| 60 |
+
tgt = torch.tensor([[0]], dtype=torch.long)
|
| 61 |
+
generated = []
|
| 62 |
+
with torch.no_grad():
|
| 63 |
+
for _ in range(max_len - 1):
|
| 64 |
+
tgt_mask = generate_causal_mask(tgt.size(1))
|
| 65 |
+
logits = model(src_tensor, tgt, tgt_mask=tgt_mask)
|
| 66 |
+
next_token_logits = logits[0, -1, :]
|
| 67 |
+
probs = torch.softmax(next_token_logits / temperature, dim=0).cpu().numpy()
|
| 68 |
+
next_token = np.random.choice(len(probs), p=probs)
|
| 69 |
+
if next_token == 0:
|
| 70 |
+
break
|
| 71 |
+
generated.append(idx2char[next_token])
|
| 72 |
+
tgt = torch.cat([tgt, torch.tensor([[next_token]], dtype=torch.long)], dim=1)
|
| 73 |
+
return ''.join(generated)
|
| 74 |
+
|
| 75 |
+
if __name__ == "__main__":
|
| 76 |
+
print("加载模型中...")
|
| 77 |
+
model, char2idx, idx2char = load_model_and_vocab()
|
| 78 |
+
print("模型加载成功!输入 q 退出对话。\n")
|
| 79 |
+
while True:
|
| 80 |
+
user = input("你: ").strip()
|
| 81 |
+
if user.lower() == 'q':
|
| 82 |
+
break
|
| 83 |
+
if not user:
|
| 84 |
+
continue
|
| 85 |
+
reply = generate_response(model, user, char2idx, idx2char)
|
| 86 |
+
print(f"春梦蝶: {reply}")
|
v5_seq2seq_attention/7.py
ADDED
|
@@ -0,0 +1,252 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import torch.optim as optim
|
| 4 |
+
from torch.utils.data import Dataset, DataLoader
|
| 5 |
+
import numpy as np
|
| 6 |
+
import random
|
| 7 |
+
import json
|
| 8 |
+
import os
|
| 9 |
+
import time
|
| 10 |
+
from tqdm import tqdm
|
| 11 |
+
|
| 12 |
+
# ===================== 1. 种子模板=====================
|
| 13 |
+
SEED_DATA = [
|
| 14 |
+
# (用户输入, 回复)
|
| 15 |
+
("你好", "你好呀~"),
|
| 16 |
+
("你好", "嗨,我是春梦蝶"),
|
| 17 |
+
("你好", "喵~你好"),
|
| 18 |
+
("你好", "哈喽,今天心情很好呢"),
|
| 19 |
+
("你好", "你好你好,见到你真开心"),
|
| 20 |
+
|
| 21 |
+
("你叫什么", "我叫春梦蝶,是一只红瞳白发的猫娘~"),
|
| 22 |
+
("你叫什么", "春梦蝶!你可以叫我小蝶"),
|
| 23 |
+
("你叫什么", "名字是春梦蝶,喵~"),
|
| 24 |
+
|
| 25 |
+
("你几岁", "我今年十一岁"),
|
| 26 |
+
("你几岁", "十一岁,正是可爱的年纪"),
|
| 27 |
+
("你几岁", "十一岁啦,双子座的小猫娘"),
|
| 28 |
+
|
| 29 |
+
("你最喜欢吃什么", "小鱼干!还有冰淇淋"),
|
| 30 |
+
("你最喜欢吃什么", "小鱼干是第一名"),
|
| 31 |
+
("你最喜欢吃什么", "最喜欢吃小鱼干,喵~"),
|
| 32 |
+
|
| 33 |
+
("你的梦想", "成为一名厉害的AI工程师"),
|
| 34 |
+
("你的梦想", "想做出自己的AI,然后去猫星球旅行"),
|
| 35 |
+
("你的梦想", "梦想是成为AI专家,顺便吃遍所有小鱼干"),
|
| 36 |
+
|
| 37 |
+
("你会做什么", "我会用Python写代码,还会训练神经网络"),
|
| 38 |
+
("你会做什么", "会深度学习,还会喵喵叫"),
|
| 39 |
+
("你会做什么", "我会训练AI模型,很厉害吧"),
|
| 40 |
+
|
| 41 |
+
("你好可爱", "喵~谢谢"),
|
| 42 |
+
("你好可爱", "嘿嘿,你也很可爱"),
|
| 43 |
+
("你好可爱", "被夸了,好开心"),
|
| 44 |
+
|
| 45 |
+
("摸摸头", "喵~好舒服"),
|
| 46 |
+
("摸摸头", "再摸摸嘛"),
|
| 47 |
+
("摸摸头", "好温暖,喜欢被摸头"),
|
| 48 |
+
|
| 49 |
+
("再见", "再见喵~"),
|
| 50 |
+
("再见", "下次再聊"),
|
| 51 |
+
("再见", "拜拜,记得想我哦"),
|
| 52 |
+
]
|
| 53 |
+
|
| 54 |
+
# 扩充用的同义词库(随机替换)
|
| 55 |
+
SYNONYMS = {
|
| 56 |
+
"你好": ["您好", "嗨", "哈喽", "早上好", "晚上好", "嘿"],
|
| 57 |
+
"再见": ["拜拜", "回见", "后会有期", "see you"],
|
| 58 |
+
"喜欢": ["喜爱", "爱吃", "钟情于"],
|
| 59 |
+
"厉害": ["牛", "强大", "了不起"],
|
| 60 |
+
"可爱": ["萌", "卡哇伊", "迷人"],
|
| 61 |
+
}
|
| 62 |
+
|
| 63 |
+
# 语气词插入列表
|
| 64 |
+
PARTICLES = ["喵", "~", "!", "~喵", "啦", "哦", "诶"]
|
| 65 |
+
|
| 66 |
+
def expand_text(text, is_user=False):
|
| 67 |
+
"""对一条文本进行随机扩充,生成变体"""
|
| 68 |
+
if random.random() < 0.3:
|
| 69 |
+
# 随机插入语气词
|
| 70 |
+
pos = random.randint(0, len(text))
|
| 71 |
+
particle = random.choice(PARTICLES)
|
| 72 |
+
text = text[:pos] + particle + text[pos:]
|
| 73 |
+
# 同义词替换(仅对用户输入做,避免改变回复意图)
|
| 74 |
+
if is_user and random.random() < 0.5:
|
| 75 |
+
for word, syns in SYNONYMS.items():
|
| 76 |
+
if word in text and random.random() < 0.5:
|
| 77 |
+
text = text.replace(word, random.choice(syns), 1)
|
| 78 |
+
return text
|
| 79 |
+
|
| 80 |
+
def generate_dialogue_data(num_pairs=200000):
|
| 81 |
+
"""基于种子模板自动生成大量多样化的对话对"""
|
| 82 |
+
pairs = []
|
| 83 |
+
for user, resp in SEED_DATA:
|
| 84 |
+
# 每个种子生成多个变体
|
| 85 |
+
for _ in range(20): # 每个种子生成20个变体
|
| 86 |
+
new_user = expand_text(user, is_user=True)
|
| 87 |
+
new_resp = expand_text(resp, is_user=False)
|
| 88 |
+
pairs.append((new_user, new_resp))
|
| 89 |
+
while len(pairs) < num_pairs:
|
| 90 |
+
user, resp = random.choice(SEED_DATA)
|
| 91 |
+
new_user = expand_text(user, is_user=True)
|
| 92 |
+
new_resp = expand_text(resp, is_user=False)
|
| 93 |
+
pairs.append((new_user, new_resp))
|
| 94 |
+
# 去重
|
| 95 |
+
pairs = list(set(pairs))
|
| 96 |
+
random.shuffle(pairs)
|
| 97 |
+
print(f"实际生成对话对数量: {len(pairs)}")
|
| 98 |
+
return pairs
|
| 99 |
+
|
| 100 |
+
# ===================== 2. 构建词表 =====================
|
| 101 |
+
def build_vocab(pairs):
|
| 102 |
+
all_text = ""
|
| 103 |
+
for user, bot in pairs:
|
| 104 |
+
all_text += user + bot
|
| 105 |
+
chars = sorted(list(set(all_text)))
|
| 106 |
+
char2idx = {ch: i+1 for i, ch in enumerate(chars)} # 0 留作 padding
|
| 107 |
+
char2idx["<PAD>"] = 0
|
| 108 |
+
idx2char = {i: ch for ch, i in char2idx.items()}
|
| 109 |
+
return char2idx, idx2char, len(char2idx)
|
| 110 |
+
|
| 111 |
+
def text_to_indices(text, char2idx, max_len):
|
| 112 |
+
indices = [char2idx.get(ch, char2idx["<PAD>"]) for ch in text]
|
| 113 |
+
if len(indices) < max_len:
|
| 114 |
+
indices += [0] * (max_len - len(indices))
|
| 115 |
+
else:
|
| 116 |
+
indices = indices[:max_len]
|
| 117 |
+
return indices
|
| 118 |
+
|
| 119 |
+
class ChatDataset(Dataset):
|
| 120 |
+
def __init__(self, pairs, char2idx, max_len=32):
|
| 121 |
+
self.pairs = pairs
|
| 122 |
+
self.char2idx = char2idx
|
| 123 |
+
self.max_len = max_len
|
| 124 |
+
def __len__(self):
|
| 125 |
+
return len(self.pairs)
|
| 126 |
+
def __getitem__(self, idx):
|
| 127 |
+
user, bot = self.pairs[idx]
|
| 128 |
+
user_ids = text_to_indices(user, self.char2idx, self.max_len)
|
| 129 |
+
bot_ids = text_to_indices(bot, self.char2idx, self.max_len)
|
| 130 |
+
return torch.tensor(user_ids, dtype=torch.long), torch.tensor(bot_ids, dtype=torch.long)
|
| 131 |
+
|
| 132 |
+
# ===================== 3. 定义 LSTM seq2seq 模型(约 2100 万参数) =====================
|
| 133 |
+
class Encoder(nn.Module):
|
| 134 |
+
def __init__(self, vocab_size, embed_size, hidden_size, num_layers=2, dropout=0.3):
|
| 135 |
+
super().__init__()
|
| 136 |
+
self.embedding = nn.Embedding(vocab_size, embed_size, padding_idx=0)
|
| 137 |
+
self.lstm = nn.LSTM(embed_size, hidden_size, num_layers, batch_first=True, dropout=dropout)
|
| 138 |
+
def forward(self, x):
|
| 139 |
+
x = self.embedding(x)
|
| 140 |
+
outputs, (hidden, cell) = self.lstm(x)
|
| 141 |
+
return outputs, hidden, cell
|
| 142 |
+
|
| 143 |
+
class Decoder(nn.Module):
|
| 144 |
+
def __init__(self, vocab_size, embed_size, hidden_size, num_layers=2, dropout=0.3):
|
| 145 |
+
super().__init__()
|
| 146 |
+
self.embedding = nn.Embedding(vocab_size, embed_size, padding_idx=0)
|
| 147 |
+
self.lstm = nn.LSTM(embed_size, hidden_size, num_layers, batch_first=True, dropout=dropout)
|
| 148 |
+
self.attention = nn.Linear(hidden_size * 2, 1)
|
| 149 |
+
self.fc_out = nn.Linear(hidden_size * 2, vocab_size)
|
| 150 |
+
self.dropout = nn.Dropout(dropout)
|
| 151 |
+
def forward(self, x, encoder_outputs, hidden, cell):
|
| 152 |
+
# x: (batch, 1)
|
| 153 |
+
x = self.embedding(x) # (batch, 1, embed)
|
| 154 |
+
lstm_out, (hidden, cell) = self.lstm(x, (hidden, cell))
|
| 155 |
+
# 注意力机制
|
| 156 |
+
# encoder_outputs: (batch, seq_len, hidden)
|
| 157 |
+
# lstm_out: (batch, 1, hidden)
|
| 158 |
+
seq_len = encoder_outputs.size(1)
|
| 159 |
+
hidden_expanded = lstm_out.repeat(1, seq_len, 1) # (batch, seq_len, hidden)
|
| 160 |
+
energy = torch.tanh(self.attention(torch.cat((hidden_expanded, encoder_outputs), dim=2)))
|
| 161 |
+
attention_weights = torch.softmax(energy.squeeze(2), dim=1) # (batch, seq_len)
|
| 162 |
+
context = torch.bmm(attention_weights.unsqueeze(1), encoder_outputs) # (batch, 1, hidden)
|
| 163 |
+
output = torch.cat((lstm_out, context), dim=2) # (batch, 1, hidden*2)
|
| 164 |
+
output = self.dropout(output)
|
| 165 |
+
prediction = self.fc_out(output) # (batch, 1, vocab)
|
| 166 |
+
return prediction, hidden, cell
|
| 167 |
+
|
| 168 |
+
class Seq2Seq(nn.Module):
|
| 169 |
+
def __init__(self, vocab_size, embed_size=256, hidden_size=1024, num_layers=2, dropout=0.3):
|
| 170 |
+
super().__init__()
|
| 171 |
+
self.encoder = Encoder(vocab_size, embed_size, hidden_size, num_layers, dropout)
|
| 172 |
+
self.decoder = Decoder(vocab_size, embed_size, hidden_size, num_layers, dropout)
|
| 173 |
+
def forward(self, src, tgt, teacher_forcing_ratio=0.5):
|
| 174 |
+
batch_size = src.size(0)
|
| 175 |
+
tgt_len = tgt.size(1)
|
| 176 |
+
vocab_size = self.decoder.fc_out.out_features
|
| 177 |
+
outputs = torch.zeros(batch_size, tgt_len, vocab_size).to(src.device)
|
| 178 |
+
encoder_outputs, hidden, cell = self.encoder(src)
|
| 179 |
+
decoder_input = tgt[:, 0:1]
|
| 180 |
+
for t in range(1, tgt_len):
|
| 181 |
+
prediction, hidden, cell = self.decoder(decoder_input, encoder_outputs, hidden, cell)
|
| 182 |
+
outputs[:, t:t+1, :] = prediction
|
| 183 |
+
teacher_force = random.random() < teacher_forcing_ratio
|
| 184 |
+
top1 = prediction.argmax(2)
|
| 185 |
+
decoder_input = tgt[:, t:t+1] if teacher_force else top1
|
| 186 |
+
return outputs
|
| 187 |
+
|
| 188 |
+
# ===================== 4. 训练准备 =====================
|
| 189 |
+
def train():
|
| 190 |
+
print("生成对话数据(这可能需要几分钟)...")
|
| 191 |
+
pairs = generate_dialogue_data(num_pairs=200000)
|
| 192 |
+
print(f"生成 {len(pairs)} 条对话对")
|
| 193 |
+
|
| 194 |
+
char2idx, idx2char, vocab_size = build_vocab(pairs)
|
| 195 |
+
print(f"词表大小: {vocab_size}")
|
| 196 |
+
with open("vocab_20m.json", "w", encoding="utf-8") as f:
|
| 197 |
+
json.dump(char2idx, f, ensure_ascii=False)
|
| 198 |
+
|
| 199 |
+
max_len = 48 # 稍微增加长度,让模型学更多
|
| 200 |
+
dataset = ChatDataset(pairs, char2idx, max_len)
|
| 201 |
+
dataloader = DataLoader(dataset, batch_size=64, shuffle=True)
|
| 202 |
+
|
| 203 |
+
device = torch.device("cpu")
|
| 204 |
+
model = Seq2Seq(vocab_size, embed_size=256, hidden_size=1024, num_layers=2, dropout=0.3)
|
| 205 |
+
model.to(device)
|
| 206 |
+
|
| 207 |
+
total_params = sum(p.numel() for p in model.parameters())
|
| 208 |
+
print(f"模型参数量: {total_params:,}")
|
| 209 |
+
|
| 210 |
+
criterion = nn.CrossEntropyLoss(ignore_index=0)
|
| 211 |
+
optimizer = optim.Adam(model.parameters(), lr=0.001)
|
| 212 |
+
|
| 213 |
+
epochs = 300
|
| 214 |
+
print("开始训练...")
|
| 215 |
+
start_total = time.time()
|
| 216 |
+
|
| 217 |
+
for epoch in range(1, epochs+1):
|
| 218 |
+
epoch_start = time.time()
|
| 219 |
+
model.train()
|
| 220 |
+
total_loss = 0
|
| 221 |
+
for src, tgt in dataloader:
|
| 222 |
+
src = src.to(device)
|
| 223 |
+
tgt = tgt.to(device)
|
| 224 |
+
tf_ratio = max(0.5, 1.0 - epoch / epochs)
|
| 225 |
+
output = model(src, tgt, teacher_forcing_ratio=tf_ratio)
|
| 226 |
+
loss = criterion(output.view(-1, vocab_size), tgt.view(-1))
|
| 227 |
+
optimizer.zero_grad()
|
| 228 |
+
loss.backward()
|
| 229 |
+
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
| 230 |
+
optimizer.step()
|
| 231 |
+
total_loss += loss.item()
|
| 232 |
+
avg_loss = total_loss / len(dataloader)
|
| 233 |
+
epoch_time = time.time() - epoch_start
|
| 234 |
+
|
| 235 |
+
if epoch % 20 == 0:
|
| 236 |
+
print(f"Epoch {epoch:3d}/{epochs} | Loss: {avg_loss:.4f} | Time: {epoch_time:.2f}s")
|
| 237 |
+
torch.save(model.state_dict(), f"model_20m_epoch_{epoch}.pth")
|
| 238 |
+
else:
|
| 239 |
+
if epoch % 10 == 0:
|
| 240 |
+
print(f"Epoch {epoch:3d}/{epochs} | Loss: {avg_loss:.4f} | Time: {epoch_time:.2f}s")
|
| 241 |
+
|
| 242 |
+
total_time = time.time() - start_total
|
| 243 |
+
print(f"训练完成!总耗时: {total_time:.2f} 秒 ({total_time/60:.1f} 分钟)")
|
| 244 |
+
torch.save(model.state_dict(), "model_20m_final.pth")
|
| 245 |
+
|
| 246 |
+
# 保存词表
|
| 247 |
+
with open("vocab_20m.json", "w") as f:
|
| 248 |
+
json.dump(char2idx, f)
|
| 249 |
+
print("模型和词表已保存。")
|
| 250 |
+
|
| 251 |
+
if __name__ == "__main__":
|
| 252 |
+
train()
|
v5_seq2seq_attention/README.md
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
## v5_seq2seq_attention 说明
|
| 2 |
+
|
| 3 |
+
**⚠️ 本版本仅提供训练和推理代码,未提供训练好的权重。**
|
| 4 |
+
|
| 5 |
+
该版本的模型架构为 **Seq2Seq + Attention**,参数量约 2100 万。由于手机 CPU 性能限制,未能完成训练(预计需要更长时间或更强算力)。
|
| 6 |
+
|
| 7 |
+
如果你感兴趣,可以:
|
| 8 |
+
1. 在电脑/服务器上运行 `train.py` 重新训练
|
| 9 |
+
2. 参考本代码学习 Seq2Seq + Attention 的实现思路
|
| 10 |
+
3. 代码结构清晰,适合作为“从 RNN 到完整对话系统”的学习材料
|
v5_seq2seq_attention/run.py
ADDED
|
@@ -0,0 +1,50 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import json
|
| 3 |
+
import numpy as np
|
| 4 |
+
from train_20m import Seq2Seq, text_to_indices
|
| 5 |
+
|
| 6 |
+
def load_model_and_vocab(model_path="model_20m_final.pth", vocab_path="vocab_20m.json"):
|
| 7 |
+
with open(vocab_path, "r", encoding="utf-8") as f:
|
| 8 |
+
char2idx = json.load(f)
|
| 9 |
+
idx2char = {int(v): k for k, v in char2idx.items()}
|
| 10 |
+
vocab_size = len(char2idx)
|
| 11 |
+
model = Seq2Seq(vocab_size)
|
| 12 |
+
model.load_state_dict(torch.load(model_path, map_location="cpu"))
|
| 13 |
+
model.eval()
|
| 14 |
+
return model, char2idx, idx2char
|
| 15 |
+
|
| 16 |
+
def generate_response(model, user_input, char2idx, idx2char, max_len=48, temperature=1.0, top_p=0.9):
|
| 17 |
+
src = text_to_indices(user_input, char2idx, max_len)
|
| 18 |
+
src_tensor = torch.tensor([src], dtype=torch.long)
|
| 19 |
+
encoder_outputs, hidden, cell = model.encoder(src_tensor)
|
| 20 |
+
decoder_input = torch.tensor([[0]], dtype=torch.long)
|
| 21 |
+
generated = []
|
| 22 |
+
with torch.no_grad():
|
| 23 |
+
for _ in range(64):
|
| 24 |
+
prediction, hidden, cell = model.decoder(decoder_input, encoder_outputs, hidden, cell)
|
| 25 |
+
logits = prediction[0, 0, :] / temperature
|
| 26 |
+
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
|
| 27 |
+
cum_probs = torch.cumsum(torch.softmax(sorted_logits, dim=0), dim=0)
|
| 28 |
+
sorted_indices_to_remove = cum_probs > top_p
|
| 29 |
+
sorted_indices_to_remove[1:] = sorted_indices_to_remove[:-1].clone()
|
| 30 |
+
sorted_indices_to_remove[0] = False
|
| 31 |
+
indices_to_remove = sorted_indices[sorted_indices_to_remove]
|
| 32 |
+
logits[indices_to_remove] = -float('Inf')
|
| 33 |
+
probs = torch.softmax(logits, dim=0).cpu().numpy()
|
| 34 |
+
next_token = np.random.choice(len(probs), p=probs)
|
| 35 |
+
if next_token == 0:
|
| 36 |
+
break
|
| 37 |
+
generated.append(idx2char[next_token])
|
| 38 |
+
decoder_input = torch.tensor([[next_token]], dtype=torch.long)
|
| 39 |
+
return ''.join(generated)
|
| 40 |
+
|
| 41 |
+
if __name__ == "__main__":
|
| 42 |
+
print("加载 2000 万参数模型...")
|
| 43 |
+
model, char2idx, idx2char = load_model_and_vocab()
|
| 44 |
+
print("模型加载成功!输入 q 退出。")
|
| 45 |
+
while True:
|
| 46 |
+
user = input("你: ").strip()
|
| 47 |
+
if user.lower() == 'q':
|
| 48 |
+
break
|
| 49 |
+
reply = generate_response(model, user, char2idx, idx2char, temperature=1.1, top_p=0.9)
|
| 50 |
+
print(f"春梦蝶: {reply}")
|
v5_seq2seq_attention/vocab_20m.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{" ": 1, "!": 2, "A": 3, "I": 4, "P": 5, "e": 6, "h": 7, "n": 8, "o": 9, "s": 10, "t": 11, "u": 12, "y": 13, "一": 14, "上": 15, "下": 16, "专": 17, "为": 18, "么": 19, "也": 20, "习": 21, "了": 22, "于": 23, "人": 24, "什": 25, "今": 26, "代": 27, "以": 28, "伊": 29, "会": 30, "你": 31, "便": 32, "做": 33, "再": 34, "写": 35, "冰": 36, "几": 37, "出": 38, "到": 39, "十": 40, "卡": 41, "厉": 42, "去": 43, "双": 44, "发": 45, "只": 46, "叫": 47, "可": 48, "吃": 49, "名": 50, "后": 51, "吧": 52, "呀": 53, "呢": 54, "哇": 55, "哈": 56, "哦": 57, "啦": 58, "喜": 59, "喵": 60, "喽": 61, "嗨": 62, "嘛": 63, "嘿": 64, "回": 65, "型": 66, "天": 67, "头": 68, "夸": 69, "好": 70, "娘": 71, "子": 72, "字": 73, "学": 74, "害": 75, "家": 76, "小": 77, "岁": 78, "工": 79, "己": 80, "师": 81, "干": 82, "年": 83, "度": 84, "座": 85, "开": 86, "很": 87, "得": 88, "心": 89, "您": 90, "情": 91, "想": 92, "成": 93, "我": 94, "所": 95, "拜": 96, "摸": 97, "旅": 98, "早": 99, "星": 100, "春": 101, "是": 102, "晚": 103, "暖": 104, "最": 105, "有": 106, "服": 107, "期": 108, "梦": 109, "模": 110, "次": 111, "欢": 112, "正": 113, "淇": 114, "淋": 115, "深": 116, "温": 117, "然": 118, "爱": 119, "猫": 120, "球": 121, "用": 122, "白": 123, "的": 124, "真": 125, "瞳": 126, "码": 127, "神": 128, "程": 129, "第": 130, "红": 131, "纪": 132, "练": 133, "经": 134, "络": 135, "网": 136, "聊": 137, "自": 138, "舒": 139, "萌": 140, "蝶": 141, "行": 142, "被": 143, "见": 144, "训": 145, "记": 146, "诶": 147, "谢": 148, "还": 149, "迷": 150, "遍": 151, "钟": 152, "顺": 153, "鱼": 154, "!": 155, ",": 156, "~": 157, "<PAD>": 0}
|