Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 55 additions & 0 deletions trainer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
import torch
import torch.nn as nn
import torch.optim as optim
from tqdm import tqdm
import wandb


class Trainer:
def __init__(self, model, loss_fn, lr=1e-3, T_max=1000, weight_decay=0, wandb_log=False):
self.model = model
self.loss_fn = loss_fn
self.optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)
self.scheduler = optim.lr_scheduler.CosineAnnealingLR(self.optimizer, T_max=T_max)
self.wandb_log = wandb_log

def plot_architecture(self, dataloader, depth=6):
try:
from torchview import draw_graph
sample_input = next(iter(dataloader))[0]
draw_graph(self.model, input_size=(1, *sample_input.shape[1:]), device=next(self.model.parameters()).device).save_graph(fname='model_graph_new.png', max_depth=depth)
except Exception as e:
print(f"Skipping plot_architecture: {e}")

def train_loop(self, train_loader, test_loader, n_epochs, val_metric=None):
device = next(self.model.parameters()).device
self.model.train()
for epoch in range(n_epochs):
running_loss = 0.0
for inputs, labels in tqdm(train_loader, desc=f"Epoch {epoch}"):
inputs, labels = inputs.to(device), labels.to(device)
self.optimizer.zero_grad()
outputs = self.model(inputs)
loss = self.loss_fn(outputs, labels)
loss.backward()
self.optimizer.step()
running_loss += loss.item()
self.scheduler.step()

avg_loss = running_loss / len(train_loader)
if self.wandb_log:
wandb.log({"epoch_loss": avg_loss})

if test_loader is not None and val_metric is not None:
self.model.eval()
val_metric.reset()
with torch.no_grad():
for inputs, labels in test_loader:
inputs, labels = inputs.to(device), labels.to(device)
outputs = self.model(inputs)
val_metric.update(outputs, labels)
val_acc = val_metric.compute().item()
print(f"Epoch {epoch}, Loss: {avg_loss:.4f}, Val Acc: {val_acc:.4f}")
if self.wandb_log:
wandb.log({"val_acc": val_acc})
self.model.train()