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