LumiSign / train_nn.py
anthony01's picture
update: applying traininng loss and validation loss for best epoch
9589c44
Raw
History Blame Contribute Delete
11.2 kB
import os
import logging
import json
import torch
import torch.nn.functional as F
from torch.utils import data
from sklearn.metrics import accuracy_score
from tqdm import tqdm
from models import CNN, LSTM, Transformer
from configs import CnnConfig, LstmConfig, TransformerConfig
from utils import (
seed_everything,
AverageMeter,
EarlyStopping,
load_label_map,
get_experiment_name,
load_json,
)
from dataset import KeypointsDataset, FeaturesDatset
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def train(dataloader, model, optimizer, device):
model.train()
losses = AverageMeter()
accuracy = AverageMeter()
pbar = tqdm(dataloader, desc="Training")
for batch in pbar:
input_data = batch["data"].to(device)
label = batch["label"].to(device)
optimizer.zero_grad()
preds = model(input_data)
loss = F.cross_entropy(preds, label)
loss.backward()
optimizer.step()
losses.update(loss.item())
preds = preds.detach().cpu()
accuracy.update(
accuracy_score(
label.cpu().numpy(),
torch.argmax(torch.softmax(preds, dim=-1), dim=-1).numpy(),
)
)
pbar.set_postfix(loss=losses.avg, accuracy=accuracy.avg)
torch.cuda.empty_cache()
loss_avg = losses.avg
accuracy_avg = accuracy.avg
return loss_avg, accuracy_avg
@torch.no_grad()
def validate(dataloader, model, device):
model.eval()
losses = AverageMeter()
accuracy = AverageMeter()
pbar = tqdm(dataloader, desc="Eval")
for batch in pbar:
input_data = batch["data"].to(device)
label = batch["label"].to(device)
preds = model(input_data)
loss = F.cross_entropy(preds, label)
losses.update(loss.item())
preds = preds.detach().cpu()
accuracy.update(
accuracy_score(
label.cpu().numpy(),
torch.argmax(torch.softmax(preds, dim=-1), dim=-1).numpy(),
)
)
pbar.set_postfix(loss=losses.avg, accuracy=accuracy.avg)
torch.cuda.empty_cache()
loss_avg = losses.avg
accuracy_avg = accuracy.avg
return loss_avg, accuracy_avg
def change_max_pos_embd(args, new_mpe_size, n_classes):
config = TransformerConfig(
size=args.transformer_size, max_position_embeddings=new_mpe_size
)
if args.use_cnn:
config.input_size = CnnConfig.output_dim
model = Transformer(config=config, n_classes=n_classes)
model = model.to(device)
return model
def pretrained_name(args):
load_modelName = args.dataset
if args.use_cnn:
load_modelName += "_use_cnn"
else:
load_modelName += "_no_cnn"
if args.model == "lstm":
load_modelName += "_lstm.pth"
elif args.model == "transformer":
load_modelName += "_transformer"
if args.transformer_size == "large":
load_modelName += "_large.pth"
elif args.transformer_size == "small":
load_modelName += "_small.pth"
return load_modelName
def load_pretrained(args, n_classes, model, optimizer=None, scheduler=None):
load_modelName = pretrained_name(args)
pretrained_model_links = load_json("pretrained_links.json")
if not os.path.isfile(load_modelName):
link = pretrained_model_links[load_modelName]
torch.hub.download_url_to_file(link, load_modelName, progress=True)
if args.model == "transformer":
model = change_max_pos_embd(args, new_mpe_size=256, n_classes=n_classes)
ckpt = torch.load(load_modelName, weights_only=False)
model.load_state_dict(ckpt["model"])
if args.use_pretrained == "resume_training":
optimizer.load_state_dict(ckpt["optimizer"])
scheduler.load_state_dict(ckpt["scheduler"])
return model, optimizer, scheduler
def fit(args):
exp_name = get_experiment_name(args)
logging_path = os.path.join(args.save_path, exp_name) + ".log"
logging.basicConfig(filename=logging_path, level=logging.INFO, format="%(message)s")
seed_everything(args.seed)
label_map = load_label_map(args.dataset)
early_stop_metric = getattr(args, "early_stop_metric", "val_acc")
early_stop_patience = int(getattr(args, "early_stop_patience", 15))
if early_stop_metric not in ("val_loss", "val_acc", "loss_gap"):
raise ValueError(
f"Unsupported early_stop_metric: {early_stop_metric}. "
"Expected val_loss or val_acc, or loss_gap."
)
if early_stop_patience < 1:
raise ValueError("--early_stop_patience must be >= 1.")
metric_mode = "min" if early_stop_metric in ("val_loss", "loss_gap") else "max"
print(
f"Early stopping config: metric={early_stop_metric}, "
f"mode={metric_mode}, patience={early_stop_patience}"
)
logging.info(
f"Early stopping config: metric={early_stop_metric}, "
f"mode={metric_mode}, patience={early_stop_patience}"
)
if args.use_cnn:
train_dataset = FeaturesDatset(
features_dir=os.path.join(args.data_dir, f"{args.dataset}_train_features"),
label_map=label_map,
mode="train",
)
val_dataset = FeaturesDatset(
features_dir=os.path.join(args.data_dir, f"{args.dataset}_val_features"),
label_map=label_map,
mode="val",
)
else:
train_dataset = KeypointsDataset(
keypoints_dir=os.path.join(
args.data_dir, f"{args.dataset}_train_keypoints"
),
use_augs=args.use_augs,
label_map=label_map,
mode="train",
max_frame_len=args.max_frame_len,
)
val_dataset = KeypointsDataset(
keypoints_dir=os.path.join(args.data_dir, f"{args.dataset}_val_keypoints"),
use_augs=False,
label_map=label_map,
mode="val",
max_frame_len=args.max_frame_len,
)
if len(train_dataset) == 0:
raise ValueError(
"No training samples found. "
"Check that keypoints were generated and that --data_dir points to "
f"{args.data_dir}/{args.dataset}_train_keypoints with JSON files."
)
if len(val_dataset) == 0:
raise ValueError(
"No validation samples found. "
"Check that keypoints were generated and that --data_dir points to "
f"{args.data_dir}/{args.dataset}_val_keypoints with JSON files."
)
train_dataloader = data.DataLoader(
train_dataset,
batch_size=args.batch_size,
shuffle=True,
num_workers=4,
pin_memory=True,
)
val_dataloader = data.DataLoader(
val_dataset,
batch_size=args.batch_size,
shuffle=False,
num_workers=4,
pin_memory=True,
)
n_classes = len(label_map)
if args.model == "lstm":
config = LstmConfig()
if args.use_cnn:
config.input_size = CnnConfig.output_dim
model = LSTM(config=config, n_classes=n_classes)
else:
config = TransformerConfig(size=args.transformer_size)
if args.use_cnn:
config.input_size = CnnConfig.output_dim
model = Transformer(config=config, n_classes=n_classes)
model = model.to(device)
optimizer = torch.optim.AdamW(
model.parameters(), lr=args.learning_rate, weight_decay=0.01
)
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
optimizer, mode=metric_mode, factor=0.2
)
if args.use_pretrained == "resume_training":
model, optimizer, scheduler = load_pretrained(
args, n_classes, model, optimizer, scheduler
)
model_path = os.path.join(args.save_path, exp_name) + ".pth"
es = EarlyStopping(patience=early_stop_patience, mode=metric_mode)
for epoch in range(args.epochs):
print(f"Epoch: {epoch+1}/{args.epochs}")
train_loss, train_acc = train(train_dataloader, model, optimizer, device)
val_loss, val_acc = validate(val_dataloader, model, device)
loss_gap = abs(train_loss - val_loss)
if early_stop_metric == "val_loss":
tracked_metric_value = val_loss
elif early_stop_metric == "val_acc":
tracked_metric_value = val_acc
else: # loss_gap
tracked_metric_value = loss_gap
logging.info(
f"epoch:{epoch} train_loss:{train_loss:.6f} val_loss:{val_loss:.6f} "
f"train_acc:{train_acc:.6f} val_acc:{val_acc:.6f} loss_gap:{loss_gap:.6f}"
)
scheduler.step(tracked_metric_value)
es(
model_path=model_path,
epoch_score=tracked_metric_value,
model=model,
optimizer=optimizer,
scheduler=scheduler,
)
if es.early_stop:
print("Early stopping")
break
print("### Training Complete ###")
def evaluate(args):
label_map = load_label_map(args.dataset)
n_classes = len(label_map)
eval_split = getattr(args, "eval_split", "test")
split_folder = f"{args.dataset}_{eval_split}"
if args.use_cnn:
dataset = FeaturesDatset(
features_dir=os.path.join(args.data_dir, f"{split_folder}_features"),
label_map=label_map,
mode=eval_split,
)
else:
dataset = KeypointsDataset(
keypoints_dir=os.path.join(args.data_dir, f"{split_folder}_keypoints"),
use_augs=False,
label_map=label_map,
mode=eval_split,
max_frame_len=args.max_frame_len,
)
if len(dataset) == 0:
raise ValueError(
f"No {eval_split} samples found. "
"Check that keypoints were generated and that --data_dir points to "
f"{args.data_dir}/{split_folder}_keypoints with JSON files."
)
dataloader = data.DataLoader(
dataset,
batch_size=args.batch_size,
shuffle=False,
num_workers=4,
pin_memory=True,
)
if args.model == "lstm":
config = LstmConfig()
if args.use_cnn:
config.input_size = CnnConfig.output_dim
model = LSTM(config=config, n_classes=n_classes)
else:
config = TransformerConfig(size=args.transformer_size)
if args.use_cnn:
config.input_size = CnnConfig.output_dim
model = Transformer(config=config, n_classes=n_classes)
model = model.to(device)
if args.use_pretrained == "evaluate":
model, _, _ = load_pretrained(args, n_classes, model)
print("### Model loaded ###")
else:
exp_name = get_experiment_name(args)
model_path = os.path.join(args.save_path, exp_name) + ".pth"
ckpt = torch.load(model_path)
model.load_state_dict(ckpt["model"])
print("### Model loaded ###")
test_loss, test_acc = validate(dataloader, model, device)
print(f"Evaluation Results ({eval_split} split):")
print(f"Loss: {test_loss}, Accuracy: {test_acc}")