From 65409455b16d1647c8f2b77600f060286a195442 Mon Sep 17 00:00:00 2001 From: webbrain-one <295484252+webbrain-one@users.noreply.github.com> Date: Mon, 10 Aug 2026 10:21:56 +0300 Subject: [PATCH] fix: add missing trainer module to resolve import error Adds trainer.py to fix the ModuleNotFoundError reported in #5. --- trainer.py | 55 ++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 55 insertions(+) create mode 100644 trainer.py diff --git a/trainer.py b/trainer.py new file mode 100644 index 0000000..d0c6055 --- /dev/null +++ b/trainer.py @@ -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()