File size: 2,479 Bytes
e13f18f
 
 
 
 
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
#!/usr/bin/env python3
# Copyright (c) 2026 XingChina
# SPDX-License-Identifier: BSD-3-Clause
# 本代码采用 BSD 3-Clause 许可证,详见项目根目录的 LICENSE 文件。

import torch
import json
import numpy as np
from train_20m import Seq2Seq, text_to_indices

def load_model_and_vocab(model_path="model_20m_final.pth", vocab_path="vocab_20m.json"):
    with open(vocab_path, "r", encoding="utf-8") as f:
        char2idx = json.load(f)
    idx2char = {int(v): k for k, v in char2idx.items()}
    vocab_size = len(char2idx)
    model = Seq2Seq(vocab_size)
    model.load_state_dict(torch.load(model_path, map_location="cpu"))
    model.eval()
    return model, char2idx, idx2char

def generate_response(model, user_input, char2idx, idx2char, max_len=48, temperature=1.0, top_p=0.9):
    src = text_to_indices(user_input, char2idx, max_len)
    src_tensor = torch.tensor([src], dtype=torch.long)
    encoder_outputs, hidden, cell = model.encoder(src_tensor)
    decoder_input = torch.tensor([[0]], dtype=torch.long)
    generated = []
    with torch.no_grad():
        for _ in range(64):
            prediction, hidden, cell = model.decoder(decoder_input, encoder_outputs, hidden, cell)
            logits = prediction[0, 0, :] / temperature
            sorted_logits, sorted_indices = torch.sort(logits, descending=True)
            cum_probs = torch.cumsum(torch.softmax(sorted_logits, dim=0), dim=0)
            sorted_indices_to_remove = cum_probs > top_p
            sorted_indices_to_remove[1:] = sorted_indices_to_remove[:-1].clone()
            sorted_indices_to_remove[0] = False
            indices_to_remove = sorted_indices[sorted_indices_to_remove]
            logits[indices_to_remove] = -float('Inf')
            probs = torch.softmax(logits, dim=0).cpu().numpy()
            next_token = np.random.choice(len(probs), p=probs)
            if next_token == 0:
                break
            generated.append(idx2char[next_token])
            decoder_input = torch.tensor([[next_token]], dtype=torch.long)
    return ''.join(generated)

if __name__ == "__main__":
    print("加载 2000 万参数模型...")
    model, char2idx, idx2char = load_model_and_vocab()
    print("模型加载成功!输入 q 退出。")
    while True:
        user = input("你: ").strip()
        if user.lower() == 'q':
            break
        reply = generate_response(model, user, char2idx, idx2char, temperature=1.1, top_p=0.9)
        print(f"春梦蝶: {reply}")