简介
聊天机器人是AI最常见的应用之一。本文将使用PyTorch实现一个基于Transformer的简易聊天机器人,它能学习对话模式并生成回复。
环境准备
安装依赖
pip install torch transformers datasets
硬件要求
- 推荐GPU(如NVIDIA RTX 3060以上)
- 至少8GB显存
数据准备
数据集选择
我们使用DailyDialog数据集,包含日常对话。
from datasets import load_dataset
dataset = load_dataset("daily_dialog")
数据预处理
- 标记化:使用BERT tokenizer
- 构建输入输出对:将对话历史作为输入,下一句作为输出
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased")
def preprocess(examples):
# 处理逻辑
return encodings
模型构建
使用预训练模型
我们基于DialoGPT-small进行微调,它是GPT-2的对话版本。
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained("microsoft/DialoGPT-small")
自定义模型(可选)
如果想从头训练,可以使用Transformer模块:
import torch.nn as nn
class ChatbotTransformer(nn.Module):
def __init__(self, vocab_size, d_model, nhead, num_layers):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.transformer = nn.Transformer(d_model, nhead, num_layers)
self.fc = nn.Linear(d_model, vocab_size)
def forward(self, src, tgt):
# 前向传播
pass
训练
设置训练参数
from transformers import TrainingArguments, Trainer
training_args = TrainingArguments(
output_dir="./results",
num_train_epochs=3,
per_device_train_batch_size=4,
save_steps=500,
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset,
)
trainer.train()
损失函数
使用交叉熵损失,忽略填充标记。
推理与部署
生成回复
def generate_response(model, tokenizer, input_text):
inputs = tokenizer.encode(input_text, return_tensors="pt")
outputs = model.generate(inputs, max_length=100, pad_token_id=tokenizer.eos_token_id)
response = tokenizer.decode(outputs[0], skip_special_tokens=True)
return response
部署为Web服务
使用Flask或FastAPI:
from flask import Flask, request, jsonify
app = Flask(__name__)
@app.route("/chat", methods=["POST"])
def chat():
data = request.json
response = generate_response(model, tokenizer, data["message"])
return jsonify({"response": response})
优化建议
- 数据增强:添加同义词替换
- 模型蒸馏:使用更小的模型加快推理
- 多轮对话:维护对话历史
总结
通过本教程,你已经学会用PyTorch构建一个基础聊天机器人。实际应用中还需考虑上下文管理、情感识别等。继续探索吧!