XingChina's picture
Update v1_rnn_mini/run.py
ad03c96 verified
Raw
History Blame Contribute Delete
2.86 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 numpy as np
import json
import random
# 加载字符映射
with open("checkpoints_mengdie/char2idx.json", "r", encoding="utf-8") as f:
char2idx = json.load(f)
idx2char = {int(v): k for k, v in char2idx.items()}
vocab_size = len(char2idx)
print(f"词表大小: {vocab_size}")
print(f"词表示例: {list(char2idx.items())[:10]}")
# 模型结构(必须与训练时一致!!!)
class TinyCharRNN(nn.Module):
def __init__(self, vocab_size, hidden_size=32):
super().__init__()
self.embedding = nn.Embedding(vocab_size, hidden_size)
self.rnn = nn.RNN(hidden_size, hidden_size, batch_first=True)
self.fc = nn.Linear(hidden_size, vocab_size)
def forward(self, x, hidden=None):
x = self.embedding(x)
out, hidden = self.rnn(x, hidden)
out = self.fc(out)
return out, hidden
# 加载模型
model = TinyCharRNN(vocab_size, hidden_size=32)
model.load_state_dict(torch.load("checkpoints_mengdie/mengdie_final.pth", map_location='cpu'))
model.eval()
print("春梦蝶猫娘模型加载成功!喵~\n")
def generate_response(prompt, length=150, temperature=0.8):
"""根据提示词生成猫娘的回应"""
if not prompt:
prompt = random.choice(list(char2idx.keys()))
# 将prompt转为索引
indices = []
for ch in prompt:
if ch in char2idx:
indices.append(char2idx[ch])
else:
# 如果字符不在词表中,跳过(或者用空格代替)
# 这里跳过,不影响生成
continue
if not indices:
# 如果全部跳过,就用一个默认字符
indices = [char2idx['你']]
input_tensor = torch.tensor([indices])
hidden = None
result = list(prompt) # 保留原始输入
with torch.no_grad():
for _ in range(length):
logits, hidden = model(input_tensor, hidden)
probs = torch.softmax(logits[0, -1] / temperature, dim=0).cpu().numpy()
next_idx = np.random.choice(len(probs), p=probs)
next_char = idx2char[next_idx]
result.append(next_char)
input_tensor = torch.tensor([[next_idx]])
return ''.join(result)
print("你可以开始和春梦蝶聊天了!输入 'q' 退出。")
while True:
user = input("\n你: ")
if user.lower() == 'q':
break
# 使用用户输入的最后几个字符作为生成起点
start = user[-5:] if len(user) >= 5 else user
reply = generate_response(start, length=150, temperature=0.85)
print(f"春梦蝶: {reply}")