基于Transformer架构的对联生成模型

This commit is contained in:
huaian_zhou 2023-12-06 15:33:39 +08:00
parent d90dd3f29e
commit 72b5be1811
10 changed files with 1549769 additions and 1 deletions

View File

@ -7,7 +7,6 @@ from load_dataset import tokenizer
if __name__ == "__main__":
cls_dict = {1: "World", 2: "Sports", 3: "Business", 4: "Sci/Tech"}
# classes[2, 4, 3, 1]
seqs = [
# "3",
# "Fears for T N pension after talks",

101
trans_couplet/config.py Normal file
View File

@ -0,0 +1,101 @@
import os
import sys
import torch
import logging
from datetime import datetime
def logger_init(
log_filename="monitor", log_level=logging.DEBUG, log_dir="./log/", only_file=False
):
"""
:param log_filename: 日志文件名
:param log_level: 日志等级
:param log_dir: 日志目录
:parma only_file: 是否只保存到日志文件中
"""
# 指定日志文件路径
if not os.path.exists(log_dir):
os.makedirs(log_dir)
log_filepath = os.path.join(
log_dir, log_filename + "_" + str(datetime.now())[:10] + ".txt"
)
# 指定日志格式
formatter = "[%(asctime)s] - %(levelname)s: %(message)s"
# 只保存到日志文件中
if only_file:
logging.basicConfig(
filename=log_filepath,
level=log_level,
format=formatter,
datefmt="%Y-%m-%d %H:%M:%S",
)
# 保存到日志文件并输出到终端
else:
logging.basicConfig(
level=log_level,
format=formatter,
datefmt="%Y-%m-%d %H:%M:%S",
handlers=[
logging.FileHandler(log_filepath),
logging.StreamHandler(sys.stdout),
],
)
class Config:
"""
基于Transformer架构的对联生成模型配置信息
"""
def __init__(self):
# 数据集相关配置
self.task_dir = os.path.dirname(os.path.abspath(__file__))
self.data_dir = os.path.join(self.task_dir, "data")
self.train_filepaths = [
os.path.join(self.data_dir, "couplet", "train", "in.txt"),
os.path.join(self.data_dir, "couplet", "train", "out.txt"),
]
self.test_filepaths = [
os.path.join(self.data_dir, "couplet", "test", "in.txt"),
os.path.join(self.data_dir, "couplet", "test", "out.txt"),
]
self.min_freq = 1
# 模型相关配置
self.batch_size = 256
self.d_model = 512
self.nhead = 8
self.num_encoder_layers = 6
self.num_decoder_layers = 6
self.dim_feedforward = 1024
self.dropout = 0.1
self.beta1 = 0.9
self.beta2 = 0.98
self.epsilon = 10e-9
self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
self.epochs = 50
self.info_print_steps = 30 # 打印信息间隔steps
self.model_eval_epochs = 2 # 验证模型epochs
self.model_save_dir = os.path.join(self.task_dir, "cache")
if not os.path.exists(self.model_save_dir):
os.makedirs(self.model_save_dir)
# 日志相关
self.log_dir = os.path.join(self.model_save_dir, "log")
logger_init(
log_filename="log_train",
log_level=logging.INFO,
log_dir=self.log_dir,
)
# logging.info("### 将当前配置打印到日志文件中 ")
# for key, value in self.__dict__.items():
# logging.info(f"### {key} = {value}")
if __name__ == "__main__":
config = Config()
for key, value in config.__dict__.items():
print(f"### {key} = {value}")
print("=" * 20)

90
trans_couplet/couplet.py Normal file
View File

@ -0,0 +1,90 @@
import torch
from config import Config
from couplet_model import CoupletModel
from load_dataset import tokenizer
def inference(model, src, vocab, config: Config):
model.eval()
with torch.no_grad():
# 分词、转换为索引序列
ids = vocab.lookup_indices(tokenizer(src))
# 变换源输入序列形状为[src_len, 1]其中1表示batch_size
src = torch.tensor(ids, dtype=torch.long).reshape(-1, 1).to(config.device)
# 初始化目标输入序列第一个token为<bos> [tgt_len, 1]
tgt = torch.ones(1, 1).type_as(src.data).fill_(vocab["<bos>"])
# 对源输入序列执行一次Encode获得memory [src_len, 1, embed_dim]
memory = model.infer_encoder(src)
# 对目标输入序列循环执行Decode操作每次预测下一个token
max_len = src.shape[0] # 执行Decode操作最大次数
for i in range(max_len):
out = model.infer_decoder(tgt, memory) # [tgt_len, 1, embed_dim]
out = out.transpose(0, 1) # [1, tgt_len, embed_dim]
# 对最后一个token的对应输出进行分类 [1, vocab_size]
prob = model.classification(out[:, -1, :])
# 选择概率最大的类别即当前时刻预测的token
_, next_tok_id = torch.max(prob, dim=1)
next_tok_id = next_tok_id.item()
# 将当前时刻预测的token与目标输入序列拼接作为下一个时刻的目标输入序列再去预测下一个词
tgt = torch.cat(
[tgt, torch.ones(1, 1).type_as(src.data).fill_(next_tok_id)], dim=0
)
# 若当前时刻预测token为结束<eos>,则跳出循环结束预测
if next_tok_id == vocab["<eos>"]:
break
tgt = tgt.flatten()
# 把索引序列转换为token序列、拼接成句子并删除<bos>和<eos>
return (
"".join(vocab.lookup_tokens(tgt.tolist()))
.replace("<bos>", "")
.replace("<eos>", "")
)
def do_couplet(srcs, config: Config):
"""
根据上联对出下联
"""
# 读取词表
vocab = torch.load(config.model_save_dir + "/vocab.pt")
# 创建模型
couplet_model = CoupletModel(
vocab_size=len(vocab),
d_model=config.d_model,
nhead=config.nhead,
num_encoder_layers=config.num_encoder_layers,
num_decoder_layers=config.num_decoder_layers,
dim_feedforward=config.dim_feedforward,
dropout=config.dropout,
).to(config.device)
# 加载模型权重
loaded_params = torch.load(config.model_save_dir + "/model.pt")
couplet_model.load_state_dict(loaded_params)
# 执行推理——一次生成一条下联
results = []
for src in srcs:
result = inference(couplet_model, src, vocab, config)
results.append(result)
return results
if __name__ == "__main__":
srcs = ["晚风摇树树还挺", "忽忽几晨昏,离别间之,疾病间之,不及终年同静好", "风声、雨声、读书声,声声入耳"]
srcs = [" ".join(src) for src in srcs]
tgts = ["晨露润花花更红", "茕茕小儿女,孱羸若此,娇憨若此,更烦二老费精神", "家事、国事、天下事,事事关心"]
config = Config()
results = do_couplet(srcs, config)
for src, tgt, result in zip(srcs, tgts, results):
print(f"上联:{''.join(src.split())}")
print(f"A.I.{result}")
print(f"下联:{tgt}")
print("=" * 20)

View File

@ -0,0 +1,159 @@
import sys, os
sys.path.append(os.getcwd())
import torch
import torch.nn as nn
from model.my_transformer import Transformer
from model.my_embedding import PositionalEncoding, TokenEmbedding
class CoupletModel(nn.Module):
"""
基于Transformer架构的对联生成模型
"""
def __init__(
self,
vocab_size,
d_model=512,
nhead=8,
num_encoder_layers=6,
num_decoder_layers=6,
dim_feedforward=2048,
dropout=0.1,
):
"""
:param vocab_size: 词表大小
"""
super(CoupletModel, self).__init__()
# 注意对联生成模型中Encoder和Decoder使用同一个词表所以也使用同一个Token Embedding层
self.token_embedding = TokenEmbedding(vocab_size, d_model)
self.pos_encoding = PositionalEncoding(d_model=d_model, dropout=dropout)
self.transformer = Transformer(
d_model=d_model,
nhead=nhead,
num_encoder_layers=num_encoder_layers,
num_decoder_layers=num_decoder_layers,
dim_feedforward=dim_feedforward,
dropout=dropout,
)
self.classification = nn.Linear(d_model, vocab_size)
self._init_parameters()
def forward(
self,
src=None,
tgt=None,
src_mask=None,
tgt_mask=None,
memory_mask=None,
src_key_padding_mask=None,
tgt_key_padding_mask=None,
memory_key_padding_mask=None,
):
"""
:param src: Encoder的输入序列 # [src_len, batch_size]
:param tgt: Decoder的输入序列 # [tgt_len, batch_size]
:return: # [tgt_len, batch_size, vocab_size]
"""
# Token Embedding和Positional Encoding
src_embed = self.token_embedding(src)
src_embed = self.pos_encoding(src_embed) # [src_len, batch_size, embed_dim]
tgt_embed = self.token_embedding(tgt)
tgt_embed = self.pos_encoding(tgt_embed) # [tgt_len, batch_size, embed_dim]
outputs = self.transformer(
src=src_embed,
tgt=tgt_embed,
src_mask=src_mask,
tgt_mask=tgt_mask,
memory_mask=memory_mask,
src_key_padding_mask=src_key_padding_mask,
tgt_key_padding_mask=tgt_key_padding_mask,
memory_key_padding_mask=memory_key_padding_mask,
) # [tgt_len, batch_size, embed_dim]
logits = self.classification(outputs) # [tgt_len, batch_size, vocab_size]
return logits
def infer_encoder(self, src):
"""
仅执行Encoder部分在inference阶段调用
"""
src_embed = self.token_embedding(src)
src_embed = self.pos_encoding(src_embed)
# 注意推理时Encoder输入序列没有进行填充所以src_key_padding_mask为None
memory = self.transformer.encoder(src_embed)
return memory
def infer_decoder(self, tgt, memory):
"""
仅执行Decoder部分在inference阶段调用
"""
tgt_embed = self.token_embedding(tgt)
tgt_embed = self.pos_encoding(tgt_embed)
# 注意推理时Decoder输入序列没有进行填充且按时间步逐次输入所以tgt_key_padding_mask和tgt_mask为None
# 同理Encoder输入序列也没有填充所以memory_key_padding_mask为None
outputs = self.transformer.decoder(tgt_embed, memory=memory)
return outputs
def _init_parameters(self):
"""
初始化模型参数
"""
for param in self.parameters():
if param.dim() > 1:
nn.init.xavier_uniform_(param)
if __name__ == "__main__":
batch_size = 2
src_len = 7
tgt_len = 8
d_model = 32
nhead = 4
# [src_len, batch_size]
src = torch.tensor([[4, 3, 2, 6, 0, 0, 0], [5, 7, 8, 2, 4, 0, 0]]).transpose(0, 1)
src_key_padding_mask = torch.tensor(
[
[False, False, False, False, True, True, True],
[False, False, False, False, False, True, True],
]
)
# [tgt_len, batch_size]
tgt = torch.tensor([[1, 3, 3, 5, 4, 3, 0, 0], [1, 6, 8, 2, 9, 1, 0, 0]]).transpose(
0, 1
)
tgt_key_padding_mask = torch.tensor(
[
[False, False, False, False, False, False, True, True],
[False, False, False, False, False, False, True, True],
]
)
trans_model = CoupletModel(
vocab_size=16,
d_model=d_model,
nhead=nhead,
num_encoder_layers=6,
num_decoder_layers=6,
dim_feedforward=256,
dropout=0.1,
)
# 生成Attention mask
tgt_mask = trans_model.transformer.generate_attn_mask(tgt_len)
logits = trans_model(
src=src,
tgt=tgt,
tgt_mask=tgt_mask,
src_key_padding_mask=src_key_padding_mask,
tgt_key_padding_mask=tgt_key_padding_mask,
memory_key_padding_mask=src_key_padding_mask,
)
print(logits.shape) # [tgt_len, batch_size, tgt_vocab_size] [8, 2, 16]

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

View File

@ -0,0 +1,224 @@
import logging
from collections import Counter
import torch
from torchtext.vocab import vocab
from torch.nn.utils.rnn import pad_sequence
from torch.utils.data import DataLoader
from tqdm import tqdm
# 数据预处理和加载数据集1.分词2.构建词表3.转换为索引序列4.填充和添加bos、eos token5.切分目标序列6.生成掩码7.构造DataLoader
def tokenizer(string: str):
"""
分词器
"""
return string.split()
def build_vocab(tokenizer, filepaths, min_freq=1, specials=None):
"""
构建词表
:param tokenizer: 分词器
:param filepaths: 文件路径
:param min_freq: 词表中token最小频率构建词表时排除频率小于该值的token
:param specials: 特殊token
"""
# 特殊token
if specials is None:
specials = ["<unk>", "<pad>", "<bos>", "<eos>"]
counter = Counter() # 计数器
# 分词并统计所有token的频率
# 注意对联生成任务中Encoder和Decoder用相同的词表由上下联文本数据集共同生成
# 统计上联数据集token
with open(filepaths[0], encoding="utf8") as f:
for string in f:
counter.update(tokenizer(string))
# 统计下联数据集token
with open(filepaths[1], encoding="utf8") as f:
for string in f:
counter.update(tokenizer(string))
# 生成词表并设置默认idx应对查询OOV token
vocabulary = vocab(counter, min_freq=min_freq, specials=specials)
vocabulary.set_default_index(vocabulary["<unk>"])
return vocabulary
class LoadDataset:
"""
加载对联数据集
"""
def __init__(self, train_filepaths=None, tokenizer=None, batch_size=2, min_freq=1):
"""
:param train_filepaths: 训练文件路径
:param tokenizer: 分词器
:param batch_size: 批次大小
:param min_freq: 词表中token最小频率
"""
# 构建词表
self.tokenizer = tokenizer
self.specials = ["<unk>", "<pad>", "<bos>", "<eos>"]
self.vocab = build_vocab(
self.tokenizer,
filepaths=train_filepaths,
min_freq=min_freq,
specials=self.specials,
)
self.batch_size = batch_size
self.PAD_IDX = self.vocab["<pad>"]
self.BOS_IDX = self.vocab["<bos>"]
self.EOS_IDX = self.vocab["<eos>"]
def token_to_idx(self, filepaths):
"""
将token序列转换为词表中的索引序列
"""
# 上下联数据迭代器
raw_in_iter = iter(open(filepaths[0], encoding="utf8"))
raw_out_iter = iter(open(filepaths[1], encoding="utf8"))
# data列表中每一个元素是一对上下联索引序列
data = []
# 分词并将token转换为词表中的索引
logging.info(f"### 正在将数据集 {filepaths} 转换成Token索引序列")
for raw_in, raw_out in tqdm(zip(raw_in_iter, raw_out_iter), ncols=80):
in_tensor = torch.tensor(
[self.vocab[token] for token in self.tokenizer(raw_in.rstrip("\n"))],
dtype=torch.long,
)
out_tensor = torch.tensor(
self.vocab.lookup_indices(self.tokenizer(raw_out.rstrip("\n"))),
dtype=torch.long,
)
data.append((in_tensor, out_tensor))
return data
def generate_batch(self, data_batch):
"""
对每个batch中的样本进行处理的函数将作为一个参数传入DataLoader的构造函数
:param data_batch: 一个batch的数据
:return:
"""
# 分别存放一个batch上下联数据
in_batch, out_batch = [], []
# 遍历一个batch处理每一个样本
for in_item, out_item in data_batch:
# 编码器输入序列不做处理
in_batch.append(in_item)
# 解码器输入序列添加<bos>和<eos>两个特殊token
out = torch.cat(
[torch.tensor([self.BOS_IDX]), out_item, torch.tensor([self.EOS_IDX])],
dim=0,
)
out_batch.append(out)
# 以batch内最长序列为标准进行填充
# [in_len, batch_size]
in_batch = pad_sequence(in_batch, padding_value=self.PAD_IDX)
# [out_len, batch_size]
out_batch = pad_sequence(out_batch, padding_value=self.PAD_IDX)
return in_batch, out_batch
def data_loader(self, train_filepaths, test_filepaths):
"""
生成训练测试数据集
"""
train_data = self.token_to_idx(train_filepaths)
test_data = self.token_to_idx(test_filepaths)
# DataLoader取一个批次数据调用collate_fn处理然后返回collate_fn的返回值
train_loader = DataLoader(
train_data,
batch_size=self.batch_size,
shuffle=True,
collate_fn=self.generate_batch,
)
test_loader = DataLoader(
test_data,
batch_size=self.batch_size,
shuffle=True,
collate_fn=self.generate_batch,
)
return train_loader, test_loader
def generate_attn_mask(self, sz, device):
"""
生成注意力掩码矩阵
"""
mask = torch.tril(torch.ones((sz, sz), device=device)) # tril取矩阵下三角包括对角线
mask = mask.masked_fill(mask == 0, float("-inf")).masked_fill(
mask == 1, float(0.0)
)
return mask
def create_mask(self, src, tgt, device="cpu"):
"""
生成掩码
"""
src_len = src.shape[0]
tgt_len = tgt.shape[0]
# Encoder的注意力掩码全0
src_mask = torch.zeros((src_len, src_len), device=device)
# Decoder的注意力掩码 # [tgt_len, tgt_len]
tgt_mask = self.generate_attn_mask(tgt_len, device)
# Encoder输入序列的Padding掩码 # [batch_size, src_len]
src_padding_mask = (src == self.PAD_IDX).transpose(0, 1)
# Decoder输入序列的Padding掩码 # [batch_size, tgt_len]
tgt_padding_mask = (tgt == self.PAD_IDX).transpose(0, 1)
return src_mask, tgt_mask, src_padding_mask, tgt_padding_mask
if __name__ == "__main__":
import os
data_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "data")
filepaths = [
os.path.join(data_dir, "couplet", "test", "in.txt"),
os.path.join(data_dir, "couplet", "test", "out.txt"),
]
dataset = LoadDataset(
filepaths,
tokenizer=tokenizer,
batch_size=128,
min_freq=1,
)
train_loader, test_loader = dataset.data_loader(
filepaths,
filepaths,
)
print("pad idx: ", dataset.PAD_IDX)
for src, tgt in train_loader:
tgt_input = tgt[:-1, :] # Decoder输入即目标输入
tgt_output = tgt[1:, :] # Decoder输出即目标输出
(
src_mask,
tgt_mask,
src_padding_mask,
tgt_padding_mask,
) = dataset.create_mask(src, tgt_input)
print("src shape: ", src.shape) # [src_len, batch_size]
print(src.transpose(0, 1)[:3])
print("src input shape: ", src.shape)
print("src_padding_mask shape (batch_size, src_len): ", src_padding_mask.shape)
print("tgt shape: ", tgt.shape)
print("tgt input shape: ", tgt_input.shape) # [tgt_len, batch_size]
print(tgt_input.transpose(0, 1)[:3])
print("tgt output shape: ", tgt_output.shape) # [tgt_len, batch_size]
print(tgt_output.transpose(0, 1)[:3])
print("tgt_mask shape (tgt_len, tgt_len): ", tgt_mask.shape)
print("tgt_padding_mask shape (batch_size, tgt_len): ", tgt_padding_mask.shape)
break

207
trans_couplet/train.py Normal file
View File

@ -0,0 +1,207 @@
import os
import time
import logging
from typing import Any
import torch
import torch.nn as nn
from copy import deepcopy
from config import Config
from couplet_model import CoupletModel
from load_dataset import LoadDataset, tokenizer
class CustomScheduler(object):
"""
更新优化器中的学习率
"""
def __init__(self, d_model, warmup_steps=4000, optimizer=None):
super(CustomScheduler, self).__init__()
self.d_model = torch.tensor(d_model, dtype=torch.float32)
self.warmup_steps = warmup_steps
self.steps = 1.0
self.optimizer = optimizer
def step(self):
arg1 = self.steps ** (-0.5)
arg2 = self.steps * (self.warmup_steps ** (-1.5))
self.steps += 1.0
lr = (self.d_model**-0.5) * min(arg1, arg2) # 计算新的学习率
# 更新优化器中的学习率
for param in self.optimizer.param_groups:
param["lr"] = lr
def accuracy(logits, y_true, PAD_IDX):
"""
计算预测token的准确率
:param logits: 模型输出logits # [tgt_len, batch_size, vocab_size]
:param y_true: 真实索引值 # [tgt_len, batch_size]
:param PAD_IDX: padding token <pad>对应的索引值
"""
# [tgt_len, batch_size, vocab_size]->[batch_size * tgt_len, ]
y_pred = logits.transpose(0, 1).argmax(axis=2).reshape(-1)
# [tgt_len, batch_size]->[batch_size * tgt_len, ]
y_true = y_true.transpose(0, 1).reshape(-1)
# padding token mask其中padding token所在位置为False非padding token位置为True
mask = torch.logical_not(y_true.eq(PAD_IDX))
acc = y_pred.eq(y_true) # 比较预测值与真实值
acc = acc.logical_and(mask) # 筛掉padding token
# 计算预测准确率
total = mask.sum().item()
correct = acc.sum().item()
return float(correct) / total, total, correct
def train_model(config: Config):
# 1.加载、划分数据集
logging.info("############ 载入数据集 ############")
dataset = LoadDataset(
train_filepaths=config.train_filepaths,
tokenizer=tokenizer,
batch_size=config.batch_size,
min_freq=config.min_freq,
)
logging.info("############ 划分数据集 ############")
train_loader, test_loader = dataset.data_loader(
config.train_filepaths,
config.test_filepaths,
)
# 保存词表
vocab_save_path = os.path.join(config.model_save_dir, "vocab.pt")
torch.save(dataset.vocab, vocab_save_path)
# 2.创建模型对象
logging.info("############ 初始化模型 ############")
couplet_model = CoupletModel(
vocab_size=len(dataset.vocab),
d_model=config.d_model,
nhead=config.nhead,
num_encoder_layers=config.num_encoder_layers,
num_decoder_layers=config.num_decoder_layers,
dim_feedforward=config.dim_feedforward,
dropout=config.dropout,
)
model_save_path = os.path.join(config.model_save_dir, "model.pt")
# 加载已有模型(参数)
if os.path.exists(model_save_path):
loaded_params = torch.load(model_save_path)
couplet_model.load_state_dict(loaded_params)
logging.info("####成功载入已有模型,进行追加训练......")
couplet_model = couplet_model.to(config.device)
# 3.定义损失函数和优化器
# 注意对padding token不计算损失等效于给tgt input seq执行row padding mask操作
loss_fn = nn.CrossEntropyLoss(ignore_index=dataset.PAD_IDX)
optimizer = torch.optim.Adam(
couplet_model.parameters(),
lr=0.0,
betas=(config.beta1, config.beta2),
eps=config.epsilon,
)
# 学习率更新器
lr_scheduler = CustomScheduler(config.d_model, optimizer=optimizer)
# 4.开始训练
couplet_model.train()
best_eval_acc = 0.0 # 最佳验证准确率
for epoch in range(config.epochs):
losses = 0.0 # 记录一个epoch的损失
start_time = time.time()
for idx, (src, tgt) in enumerate(train_loader):
src = src.to(config.device) # 源输入序列 [src_len, batch_size]
tgt = tgt.to(config.device)
tgt_input = tgt[:-1, :] # 目标输入序列 [tgt_len, batch_size]
tgt_output = tgt[1:, :] # 目标输出序列 [tgt_len, batch_size]
# 生成掩码
(
src_mask,
tgt_mask,
src_padding_mask,
tgt_padding_mask,
) = dataset.create_mask(src, tgt_input, config.device)
# feed forward
logits = couplet_model(
src=src,
tgt=tgt_input,
src_mask=src_mask,
tgt_mask=tgt_mask,
src_key_padding_mask=src_padding_mask,
tgt_key_padding_mask=tgt_padding_mask,
memory_key_padding_mask=src_padding_mask,
) # [tgt_len, batch_size, vocab_size]
optimizer.zero_grad() # 清空权重(上个批次计算的)的梯度值
# 计算损失
# [tgt_len * batch_size, vocab_size] with [tgt_len * batch_size, ]
# 注意计算CrossEntropyLoss时Pytorch会先自动将tgt_out索引序列进行one-hot编码
# 然后对logits执行softmax得到每个类别的概率最后通过交叉熵公式计算
loss = loss_fn(logits.reshape(-1, logits.shape[-1]), tgt_output.reshape(-1))
loss.backward() # 反向传播,计算权重的梯度值
lr_scheduler.step() # 更新学习率
optimizer.step() # 更新权重值
losses += loss.item()
acc, _, _ = accuracy(logits, tgt_output, dataset.PAD_IDX) # 计算准确率
if (idx + 1) % config.info_print_steps == 0:
logging.info(
f"Epoch: {epoch} Batch: [{idx}/{len(train_loader)}] Train loss: {loss.item():.3f} Train acc: {acc:.5f}"
)
end_time = time.time()
train_loss = losses / len(train_loader)
logging.info(
f"Epoch: {epoch} Train loss: {train_loss:.3f} Epoch time: {(end_time - start_time):.3f}s"
)
if (epoch + 1) % config.model_eval_epochs == 0:
eval_acc = evaluate(couplet_model, dataset, test_loader, config)
logging.info(f"Eval acc: {eval_acc:.3f}")
# 保存验证准确率最好的模型
if eval_acc > best_eval_acc:
best_eval_acc = eval_acc
state_dict = deepcopy(couplet_model.state_dict())
torch.save(state_dict, model_save_path)
logging.info(f"Best eval acc: {best_eval_acc:.3f}")
def evaluate(model, dataset: LoadDataset, val_loader, config: Config):
model.eval() # 设置验证模式
correct, total = 0, 0
with torch.no_grad():
for idx, (src, tgt) in enumerate(val_loader):
src = src.to(config.device)
tgt = tgt.to(config.device)
tgt_input = tgt[:-1, :]
tgt_output = tgt[1:, :]
(
src_mask,
tgt_mask,
src_padding_mask,
tgt_padding_mask,
) = dataset.create_mask(src, tgt_input, config.device)
logits = model(
src=src,
tgt=tgt_input,
src_mask=src_mask,
tgt_mask=tgt_mask,
src_key_padding_mask=src_padding_mask,
tgt_key_padding_mask=tgt_padding_mask,
memory_key_padding_mask=src_padding_mask,
)
_, t, c = accuracy(logits, tgt_output, dataset.PAD_IDX)
total += t
correct += c
model.train() # 验证结束,重新设置为训练模式
return float(correct) / total
if __name__ == "__main__":
config = Config()
train_model(config)