XingChina's picture
Update v4_transformer/run.py
f3071cf verified
Raw
History Blame Contribute Delete
3.95 kB
#!/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}")