File size: 3,949 Bytes
f3071cf
 
 
 
 
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
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
#!/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 json
import numpy as np

class TinyTransformerChat(nn.Module):
    def __init__(self, vocab_size, d_model=256, nhead=8, num_encoder_layers=2,
                 num_decoder_layers=2, dim_feedforward=512, max_len=32):
        super().__init__()
        self.d_model = d_model
        self.max_len = max_len
        self.embedding = nn.Embedding(vocab_size, d_model)
        self.pos_encoding = nn.Parameter(torch.zeros(1, max_len, d_model))
        encoder_layer = nn.TransformerEncoderLayer(d_model, nhead, dim_feedforward, batch_first=True)
        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_encoder_layers)
        decoder_layer = nn.TransformerDecoderLayer(d_model, nhead, dim_feedforward, batch_first=True)
        self.transformer_decoder = nn.TransformerDecoder(decoder_layer, num_decoder_layers)
        self.fc_out = nn.Linear(d_model, vocab_size)

    def forward(self, src, tgt, src_mask=None, tgt_mask=None):
        src_len = src.size(1)
        tgt_len = tgt.size(1)
        if src_len > self.max_len or tgt_len > self.max_len:
            src = src[:, :self.max_len]
            tgt = tgt[:, :self.max_len]
            src_len = min(src_len, self.max_len)
            tgt_len = min(tgt_len, self.max_len)
        src_emb = self.embedding(src) * (self.d_model ** 0.5) + self.pos_encoding[:, :src_len, :]
        tgt_emb = self.embedding(tgt) * (self.d_model ** 0.5) + self.pos_encoding[:, :tgt_len, :]
        memory = self.transformer_encoder(src_emb, src_mask)
        output = self.transformer_decoder(tgt_emb, memory, tgt_mask)
        logits = self.fc_out(output)
        return logits

def generate_causal_mask(size):
    mask = torch.triu(torch.ones(size, size), diagonal=1).bool()
    return mask

def load_model_and_vocab(model_path="chat_model_final.pth", vocab_path="chat_vocab.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 = TinyTransformerChat(vocab_size)
    model.load_state_dict(torch.load(model_path, map_location="cpu"))
    model.eval()
    return model, char2idx, idx2char

def text_to_indices(text, char2idx, max_len=32):
    indices = [char2idx.get(ch, 0) for ch in text]
    if len(indices) < max_len:
        indices += [0] * (max_len - len(indices))
    else:
        indices = indices[:max_len]
    return indices

def generate_response(model, user_input, char2idx, idx2char, max_len=32, temperature=0.8):
    src = text_to_indices(user_input, char2idx, max_len)
    src_tensor = torch.tensor([src], dtype=torch.long)
    tgt = torch.tensor([[0]], dtype=torch.long)
    generated = []
    with torch.no_grad():
        for _ in range(max_len - 1):
            tgt_mask = generate_causal_mask(tgt.size(1))
            logits = model(src_tensor, tgt, tgt_mask=tgt_mask)
            next_token_logits = logits[0, -1, :]
            probs = torch.softmax(next_token_logits / temperature, dim=0).cpu().numpy()
            next_token = np.random.choice(len(probs), p=probs)
            if next_token == 0:
                break
            generated.append(idx2char[next_token])
            tgt = torch.cat([tgt, torch.tensor([[next_token]], dtype=torch.long)], dim=1)
    return ''.join(generated)

if __name__ == "__main__":
    print("加载模型中...")
    model, char2idx, idx2char = load_model_and_vocab()
    print("模型加载成功!输入 q 退出对话。\n")
    while True:
        user = input("你: ").strip()
        if user.lower() == 'q':
            break
        if not user:
            continue
        reply = generate_response(model, user, char2idx, idx2char)
        print(f"春梦蝶: {reply}")