File size: 2,265 Bytes
7bf1c66
 
 
 
 
efd483e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
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
#!/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}")