#!/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_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}")