This guide demonstrates the standard end-to-end workflow for training a model using bert4torch, including tokenizer setup, dataset definition, model architecture construction, compilation, and training with callbacks.
from bert4torch.tokenizers import Tokenizer
from bert4torch.models import build_transformer_model, BaseModel
from bert4torch.snippets import ListDataset
from bert4torch.callbacks import Callback, Logger, Tensorboard, AdversarialTraining
import torch.nn as nn
import torch
import torch.optim as optim
from torch.utils.data import DataLoader
# 1. Setup Tokenizer
tokenizer = Tokenizer(dict_path, do_lower_case=True)
# 2. Define Dataset
class MyDataset(ListDataset):
@staticmethod
def load_data(filenames):
D = []
return D
def collate_fn(batch):
batch_token_ids, batch_segment_ids, batch_labels = [], [], []
return [batch_token_ids, batch_segment_ids], batch_labels.flatten()
train_dataloader = DataLoader(MyDataset('file_path'), batch_size=batch_size, shuffle=True, collate_fn=collate_fn)
# 3. Define Model Architecture
class Model(BaseModel):
def __init__(self) -> None:
super().__init__()
self.bert = build_transformer_model(config_path, checkpoint_path, with_pool=True)
self.dropout = nn.Dropout(0.1)
self.dense = nn.Linear(768, 2)
def forward(self, token_ids, segment_ids):
# build_transformer_model returns [hidden_states, pooled_output] if with_pool=True
# Input must be wrapped in a list/tuple if there is only one argument
hidden_states, pooled_output = self.bert([token_ids, segment_ids])
output = self.dropout(pooled_output)
output = self.dense(output)
return output
model = Model().to(device)
# 4. Compile Model
model.compile(
loss=nn.CrossEntropyLoss(),
optimizer=optim.Adam(model.parameters(), lr=2e-5),
scheduler=None,
clip_gram_norm=1.0,
grad_accumulation_steps=2,
metrics=['accuracy']
)
# 5. Define Evaluation and Callbacks
class Evaluator(Callback):
def __init__(self):
self.best_val_acc = 0.
def on_epoch_end(self, global_step, epoch, logs=None):
val_acc = evaluate(valid_dataloader)
if val_acc > self.best_val_acc:
self.best_val_acc = val_acc
model.save_weights('best_model.pt')
print(f'val_acc: {val_acc:.5f}, best_val_acc: {self.best_val_acc:.5f}\n')
# 6. Train
if __name__ __name__ == '__main__':
model.fit(train_dataloader, epochs=20, steps_per_epoch=100,
callbacks=[Evaluator(), AdversarialTraining('fgm'), Logger('./test/test.log'), Tensorboard('./test/')])