| |
| |
| |
| |
|
|
| 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}") |