| |
| |
| |
| |
|
|
| import torch |
| import torch.nn as nn |
| import numpy as np |
| import json |
| import random |
|
|
| |
| with open("checkpoints_mengdie_expanded/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) |
|
|
| 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_expanded/mengdie_final.pth", map_location='cpu')) |
| model.eval() |
| print("春梦蝶猫娘模型(扩充版)加载成功!喵~\n") |
|
|
| def generate_response(prompt, length=180, temperature=0.8): |
| if not prompt: |
| prompt = random.choice(list(char2idx.keys())) |
| indices = [] |
| for ch in prompt: |
| if ch in char2idx: |
| indices.append(char2idx[ch]) |
| 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=200, temperature=0.85) |
| print(f"春梦蝶: {reply}") |