基于Transformer架构的对联生成模型
This commit is contained in:
parent
d90dd3f29e
commit
72b5be1811
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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
|
|
@ -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 token;5.切分目标序列;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
|
||||
|
|
@ -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)
|
||||
Loading…
Reference in New Issue