Код: Выделить всё
class ResnetEncoder(nn.Module):
def __init__(self, d_model):
super().__init__()
self.resnet = resnet50(weights=ResNet50_Weights.IMAGENET1K_V2)
self.resnet = nn.Sequential(*list(self.resnet.children())[:-1])
self.fc = nn.Linear(2048, d_model)
# freeze param
for param in self.resnet.parameters():
param.requires_grad = False
def forward(self, x):
"""
:param x: [n, c, h, w]
:return: [n, d_model]
"""
# [n, 2048, 1, 1]
x = self.resnet(x)
# [n, 2048]
x = torch.flatten(x, 1)
# [n, d_model]
x = self.fc(x)
return x
class TransformerDecoder(nn.Module):
def __init__(self, d_model, vocab_size, pe_dropout, nhead, num_layers):
super().__init__()
self.embed = TokenEmbedding(vocab_size, d_model)
self.pe = PositionalEncoding(d_model, dropout=pe_dropout)
decoder_layer = nn.TransformerDecoderLayer(d_model=d_model, nhead=nhead)
self.transformer_decoder = nn.TransformerDecoder(decoder_layer, num_layers)
self.fc = nn.Linear(d_model, vocab_size)
def forward(self, memory: Tensor, tgt: Tensor, tgt_pad_mask: Tensor):
"""
:param memory: [n, d_model]
:param tgt: [seq_len, n]
:param tgt_pad_mask: [n, seq_len]
:return:
"""
# [1, n, d_model]
memory = memory.unsqueeze(0)
seq_len = tgt.size(0)
# [seq_len, seq_len]
tgt_mask = torch.triu(torch.ones(seq_len, seq_len), diagonal=1).to(tgt.device)
# [seq_len, n, d_model]
tgt = self.embed(tgt)
# [seq_len, n, d_model]
tgt = self.pe(tgt)
# [seq_len, n, d_model]
tgt = self.transformer_decoder(tgt, memory, tgt_mask, tgt_key_padding_mask=tgt_pad_mask)
# [seq_len, n, vocab_size]
logits = self.fc(tgt)
return logits
class Resnet2Transformer(nn.Module):
def __init__(self, d_model, vocab, pe_dropout, nhead, num_layers):
super().__init__()
self.encoder = ResnetEncoder(d_model)
self.decoder = TransformerDecoder(d_model, len(vocab), pe_dropout, nhead, num_layers)
self.pad_idx = vocab.stoi['']
self.vocab = vocab
self.show_trainable_params()
def forward(self, imgs: Tensor, captions: Tensor):
"""
:param imgs: [n, c, h, w]
:param captions: [seq_len, n]
:return:
"""
memory = self.encoder(imgs)
# [n, seq_len]
tgt_pad_mask = captions.permute(1, 0) == self.pad_idx
tgt_pad_mask = tgt_pad_mask.float()
logits = self.decoder(memory, captions, tgt_pad_mask)
return logits
def show_trainable_params(self):
encoder_params = sum(p.numel() for p in self.encoder.parameters() if p.requires_grad)
decoder_params = sum(p.numel() for p in self.decoder.parameters() if p.requires_grad)
print(f'encoder_params : {encoder_params / 1e6:.2f}M')
print(f'decoder_params : {decoder_params / 1e6:.2f}M')
@staticmethod
def get_transform():
return ResNet50_Weights.IMAGENET1K_V2.transforms()
def inference(self, img, maxlen=30):
caption = [self.vocab.stoi['']]
with torch.no_grad():
self.eval()
memory = self.encoder(img)
for i in range(maxlen):
# [seq_len, 1]
tgt = torch.tensor(caption, device=img.device).unsqueeze(1)
tgt_pad_mask = tgt.permute(1, 0) == self.pad_idx
tgt_pad_mask = tgt_pad_mask.float()
logits = self.decoder(memory, tgt, tgt_pad_mask)
predict = logits[-1].argmax(1).item()
caption.append(predict)
if self.vocab.itos[predict] == '':
break
return [self.vocab.itos[i] for i in caption]
Код: Выделить всё
class EarlyStopping:
def __init__(self, checkpoint_path, patience=5, verbose=True, best_loss=float('inf'), vocab=None):
"""
:param patience: 当验证集损失没有改善时,允许的最大 epoch 数
:param verbose: 是否打印早停信息
"""
self.patience = patience
self.verbose = verbose
self.counter = 0
self.early_stop = False
self.best_loss = best_loss
self.vocab = vocab
self.checkpoint_path = checkpoint_path
def __call__(self, val_loss, model, optimizer, train_log):
if self.best_loss is None:
self.best_loss = val_loss
self.save_checkpoint(val_loss, model, optimizer, train_log)
elif val_loss >= self.best_loss:
self.counter += 1
if self.verbose:
print(f"Bad Loss [{self.counter}/{self.patience}]")
if self.counter >= self.patience:
self.early_stop = True
else:
self.save_checkpoint(val_loss, model, optimizer, train_log)
def save_checkpoint(self, val_loss, model, optimizer, train_log):
if self.verbose:
print('-' * 20)
print(f"[step : {train_log[-1]['step']}] Loss : {self.best_loss:.2f} ==> {val_loss:.2f}")
print("Save CheckPoint")
print('-' * 20)
save_checkpoint({
"state_dict": model.state_dict(),
"optimizer": optimizer.state_dict(),
"train_log": train_log,
"best_loss": val_loss,
"vocab": self.vocab
}, self.checkpoint_path)
self.best_loss = val_loss
self.counter = 0
def get_dataset(root_folder, annotation_file, mini_batch, val_split, test, test_count, transform):
print(f"device : {device}")
torch.backends.cudnn.benchmark = True
train_loader, val_loader, dataset = get_loader(
root_folder=root_folder,
annotation_file=annotation_file,
transform=transform,
num_workers=2,
batch_size=mini_batch,
val_split=val_split,
test=test,
test_count=test_count
)
return train_loader, val_loader, dataset
def forward(model, imgs, captions, criterion):
"""
:param model:
:param imgs: [n, c, h, w]
:param captions:[seq_len, n]
:param criterion:
:return:
"""
# 希望模型能够预测出
# [sequence_length, n, vocab_size]
outputs = model(imgs, captions[:-1])
# [sequence_length * n, vocab_size]
logits = outputs.reshape(-1, outputs.shape[2])
# [sequence_length * n]
label = captions[1:].reshape(-1)
# criterion(logits, targets)
loss = criterion(logits, label)
return loss
def train_step(train_log_, train_loader, model, optimizer, criterion, step, epoch, num_epochs, batch_size):
# 将模型设置为训练模式,影响模型的 dropout, normalize 等行为
model.train()
t = tqdm(train_loader)
for imgs, captions in t:
# [n, c, h, w]
imgs = imgs.to(device)
# [seq_len, n]
# caption : [, ...,, , ..., ] = [1, ..., 2, 0, ..., 0]
captions = captions.to(device)
loss = forward(model, imgs, captions, criterion)
t.set_description(f'[epoch : {epoch}/{num_epochs}] train loss ==> {loss.item():.2f}')
if len(train_log_) > 0 and train_log_[-1]['step'] == step:
train_log_[-1]['loss'] = loss.item()
else:
train_log_.append({
'step': step,
'loss': loss.item()
})
step += batch_size
optimizer.zero_grad()
loss.backward()
optimizer.step()
return step
def validate(model, val_loader, criterion, epoch, num_epochs):
model.eval()
val_loss = 0.0
with torch.no_grad():
t = tqdm(val_loader)
for imgs, captions in t:
imgs, captions = imgs.to(device), captions.to(device)
loss = forward(model, imgs, captions, criterion)
val_loss += loss.item()
t.set_description(f'[epoch : {epoch}/{num_epochs}] validate loss ==> {loss.item():.2f}')
val_loss /= len(val_loader)
return val_loss
def train(
root_folder, annotation_file,
num_epochs, lr, patience, verbose, batch_size, val_split,
transform,
load_model, checkpoint_path, load_checkpoint_path,
test, test_count,
step=0
):
train_loader, val_loader, dataset = get_dataset(root_folder, annotation_file, batch_size, val_split, test,
test_count, transform)
# model = TransformerCaptioner(device, vocab=dataset.vocab).to(device)
model = Resnet2Transformer(512, dataset.vocab, 0.1, 4, 3).to(device)
# 忽略 PAD 的损失,因为 caption 中 PAD 要么在 END 之后,要么没有
criterion = nn.CrossEntropyLoss(ignore_index=dataset.vocab.stoi[""])
optimizer = optim.Adam(model.parameters(), lr=lr)
# train
best_loss = float('inf')
train_log_ = []
if load_model:
step, best_loss, train_log_ = load_checkpoint(torch.load(load_checkpoint_path), model, optimizer)
early_stopping = EarlyStopping(checkpoint_path, patience, verbose, best_loss, dataset.vocab)
for epoch in range(num_epochs):
step = train_step(
train_log_,
train_loader,
model,
optimizer,
criterion,
step,
epoch + 1,
num_epochs,
batch_size=batch_size
)
val_loss = validate(model, val_loader, criterion, epoch + 1, num_epochs)
early_stopping(val_loss, model, optimizer, train_log_)
if early_stopping.early_stop:
print("Early stopping")
break
return train_log_
Код: Выделить всё
spacy_eng = English()
class Vocabulary:
def __init__(self, freq_threshold):
"""
只记录词频 > freq_threshold 的 token
:param freq_threshold:
"""
# : unknown
self.itos = {0: "", 1: "", 2: "", 3: ""}
self.stoi = {"": 0, "": 1, "": 2, "": 3}
self.freq_threshold = freq_threshold
def __len__(self):
return len(self.itos)
@staticmethod
def tokenizer_eng(text):
"""
对输入的英文文本进行分词,并将每个单词转换为小写形式
:param text:
:return:
"""
return [tok.text.lower() for tok in spacy_eng.tokenizer(text)]
def build_vocabulary(self, sentence_list):
frequencies = {}
idx = len(self.itos)
for sentence in sentence_list:
for word in self.tokenizer_eng(sentence):
if word not in frequencies:
frequencies[word] = 1
else:
frequencies[word] += 1
if frequencies[word] == self.freq_threshold:
self.stoi[word] = idx
self.itos[idx] = word
idx += 1
def text_to_idx(self, text):
tokenized_text = self.tokenizer_eng(text)
return [
self.stoi[token] if token in self.stoi else self.stoi[""]
for token in tokenized_text
]
class FlickrDataset(Dataset):
def __init__(self, root_dir, captions_file, transform=None, freq_threshold=5):
self.root_dir = root_dir
self.df = pd.read_csv(captions_file)
self.transform = transform
# Get img, caption columns
self.imgs = self.df["image"]
self.captions = self.df["caption"]
# Initialize vocabulary and build vocab
self.vocab = Vocabulary(freq_threshold)
self.vocab.build_vocabulary(self.captions.tolist())
def __len__(self):
return len(self.df)
def __getitem__(self, index):
caption = self.captions[index]
img_id = self.imgs[index]
img = Image.open(os.path.join(self.root_dir, img_id)).convert("RGB")
if self.transform is not None:
img = self.transform(img)
caption_index = [self.vocab.stoi[""]]
caption_index.extend(self.vocab.text_to_idx(caption))
caption_index.append(self.vocab.stoi[""])
return img, torch.tensor(caption_index)
class MyCollate:
def __init__(self, pad_idx):
"""
将 Dataset 中一个 batch 的 item 进行处理
:param pad_idx:
"""
self.pad_idx = pad_idx
def __call__(self, batch):
"""
:param batch: 一个列表,每个元素是 (data, label) 这样的 tuple,本项目中是 (image, caption),caption 也叫 target
imgs[i] : [1, c, h, w]
target[i] : [sequence_length]
:return: imgs : [n, c, h, w], targets : [sequence_length, n]
"""
imgs = [item[0].unsqueeze(0) for item in batch]
# [n, c, h, w]
imgs = torch.cat(imgs, dim=0)
# [n, sequence_length]
targets = [item[1] for item in batch]
"""
batch_first=False : [sequence_length, batch_size]
RNN 处理序列问题按照时间步,若维度为 [n, sequence_length] 则需要多一个转置的步骤
"""
targets = pad_sequence(targets, batch_first=False, padding_value=self.pad_idx)
return imgs, targets
def get_loader(
root_folder,
annotation_file,
transform,
batch_size=32,
num_workers=8,
shuffle=True,
pin_memory=False,
val_split=0.2,
random_seed=42,
test=True,
test_count=(1, 1),
) -> (DataLoader, DataLoader, Dataset):
dataset = FlickrDataset(root_folder, annotation_file, transform=transform)
pad_idx = dataset.vocab.stoi[""]
train_size = int((1 - val_split) * len(dataset))
val_size = len(dataset) - train_size
train_dataset, val_dataset = random_split(
dataset, [train_size, val_size],
generator=torch.Generator().manual_seed(random_seed)
)
if test:
train_dataset = [train_dataset[i] for i in range(test_count[0] * batch_size)]
val_dataset = [val_dataset[i] for i in range(test_count[1] * batch_size)]
train_loader = DataLoader(
dataset=train_dataset,
batch_size=batch_size,
num_workers=num_workers,
shuffle=shuffle,
# 是否将数据加载到内存的固定位置
pin_memory=pin_memory,
# 自定义如何将不同的样本合并成一个 batch
collate_fn=MyCollate(pad_idx=pad_idx),
)
val_loader = DataLoader(
dataset=val_dataset,
batch_size=batch_size,
num_workers=num_workers,
shuffle=False, # 验证集通常不需要打乱
pin_memory=pin_memory,
collate_fn=MyCollate(pad_idx=pad_idx),
)
return train_loader, val_loader, dataset
Я уже пробовал модель RNN, и она отлично сработала для подписи. Но после того, как перешёл на новую модель, надпись очень плохая. Я хочу знать, что-то не так с моим определением модели или методом обучения?
Подробнее здесь: https://stackoverflow.com/questions/790 ... verfitting