Haldhar/Malware-classification
1
1import torch2import torch.nn as nn3import torchvision4import wandb5from sklearn.metrics import accuracy_score, f1_score, precision_score, recall_score6 7 8class Train:9 """10 Training engine for the LeNet malware classifier with W&B logging.11 """12 13 def __init__(14 self,15 model: nn.Module,16 epoch: int,17 train_dataloader: torch.utils.data.DataLoader,18 val_dataloader: torch.utils.data.DataLoader,19 test_dataloader: torch.utils.data.DataLoader,20 loss_fn: nn.Module,21 optimizer: torch.optim.Optimizer,22 scheduler: torch.optim.lr_scheduler.LRScheduler,23 device: torch.device,24 config: dict,25 experiment_name: str = "baseline",26 ):27 self.model = model.to(device)28 self.epoch = epoch29 self.loss_fn = loss_fn30 self.optimizer = optimizer31 self.scheduler = scheduler32 self.train_dataloader = train_dataloader33 self.val_dataloader = val_dataloader34 self.test_dataloader = test_dataloader35 self.device = device36 37 self.run = wandb.init(38 project="malware-lenet",39 name=experiment_name,40 config=config,41 reinit=True,42 )43 44 wandb.watch(self.model, self.loss_fn, log="all", log_freq=10)45 46 images, labels = next(iter(self.train_dataloader))47 grid = torchvision.utils.make_grid(images[:16])48 grid_np = grid.permute(1, 2, 0).numpy()49 wandb.log({"sample_images": wandb.Image(grid_np, caption="Sample training batch")})50 51 print(f"W&B run initialised - {self.run.url}")52 53 def train_step(self, epoch_idx: int):54 self.model.train()55 pred_label, true_label = [], []56 train_loss = 0.057 58 for batch, (x, y) in enumerate(self.train_dataloader):59 x, y = x.to(self.device), y.to(self.device)60 y_pred = self.model(x)61 loss = self.loss_fn(y_pred, y)62 train_loss += loss.item()63 64 self.optimizer.zero_grad()65 loss.backward()66 self.optimizer.step()67 68 pred = torch.argmax(y_pred, dim=1)69 pred_label.extend(pred.cpu().numpy())70 true_label.extend(y.cpu().numpy())71 72 if batch % 10 == 0:73 global_step = epoch_idx * len(self.train_dataloader) + batch74 wandb.log({"batch/train_loss": loss.item(), "batch/step": global_step})75 76 self.scheduler.step()77 78 train_acc = accuracy_score(true_label, pred_label)79 train_prec = precision_score(true_label, pred_label, average="macro", zero_division=0)80 train_rec = recall_score(true_label, pred_label, average="macro", zero_division=0)81 train_f1 = f1_score(true_label, pred_label, average="macro", zero_division=0)82 83 return train_loss / len(self.train_dataloader), train_acc, train_prec, train_rec, train_f184 85 def val_step(self):86 self.model.eval()87 pred_label, true_label = [], []88 val_loss = 0.089 90 with torch.inference_mode():91 for x, y in self.val_dataloader:92 x, y = x.to(self.device), y.to(self.device)93 y_pred = self.model(x)94 val_loss += self.loss_fn(y_pred, y).item()95 pred = torch.argmax(y_pred, dim=1)96 pred_label.extend(pred.cpu().numpy())97 true_label.extend(y.cpu().numpy())98 99 val_acc = accuracy_score(true_label, pred_label)100 val_prec = precision_score(true_label, pred_label, average="macro", zero_division=0)101 val_rec = recall_score(true_label, pred_label, average="macro", zero_division=0)102 val_f1 = f1_score(true_label, pred_label, average="macro", zero_division=0)103 104 return val_loss / len(self.val_dataloader), val_acc, val_prec, val_rec, val_f1105 106 def test_step(self):107 self.model.eval()108 pred_label, true_label = [], []109 test_loss = 0.0110 111 with torch.inference_mode():112 for x, y in self.test_dataloader:113 x, y = x.to(self.device), y.to(self.device)114 y_pred = self.model(x)115 test_loss += self.loss_fn(y_pred, y).item()116 pred = torch.argmax(y_pred, dim=1)117 pred_label.extend(pred.cpu().numpy())118 true_label.extend(y.cpu().numpy())119 120 test_acc = accuracy_score(true_label, pred_label)121 test_prec = precision_score(true_label, pred_label, average="macro", zero_division=0)122 test_rec = recall_score(true_label, pred_label, average="macro", zero_division=0)123 test_f1 = f1_score(true_label, pred_label, average="macro", zero_division=0)124 125 return test_loss / len(self.test_dataloader), test_acc, test_prec, test_rec, test_f1126 127 def engine(self):128 best_val_accuracy = 0.0129 130 for i in range(self.epoch):131 t_loss, t_acc, t_prec, t_rec, t_f1 = self.train_step(i)132 v_loss, v_acc, v_prec, v_rec, v_f1 = self.val_step()133 current_lr = self.optimizer.param_groups[0]["lr"]134 135 wandb.log(136 {137 "epoch": i + 1,138 "train/loss": t_loss,139 "train/accuracy": t_acc,140 "train/precision": t_prec,141 "train/recall": t_rec,142 "train/f1": t_f1,143 "val/loss": v_loss,144 "val/accuracy": v_acc,145 "val/precision": v_prec,146 "val/recall": v_rec,147 "val/f1": v_f1,148 "train/lr": current_lr,149 }150 )151 152 print(153 f"Epoch [{i+1:>3}/{self.epoch}] "154 f"| Train Loss: {t_loss:.4f} Acc: {t_acc:.4f} F1: {t_f1:.4f} "155 f"| Val Loss: {v_loss:.4f} Acc: {v_acc:.4f} F1: {v_f1:.4f} "156 f"| LR: {current_lr:.6f}"157 )158 159 if v_acc > best_val_accuracy:160 best_val_accuracy = v_acc161 torch.save(self.model.state_dict(), "best_model.pth")162 wandb.save("best_model.pth")163 print(f" New best model saved (val_acc={v_acc:.4f})")164 165 print("\n" + "=" * 70)166 print("FINAL TEST EVALUATION (loaded from best_model.pth)")167 print("=" * 70)168 self.model.load_state_dict(torch.load("best_model.pth", map_location=self.device))169 test_loss, test_acc, test_prec, test_rec, test_f1 = self.test_step()170 171 print(f"Test Loss : {test_loss:.4f}")172 print(f"Test Accuracy : {test_acc:.4f}")173 print(f"Test Precision : {test_prec:.4f}")174 print(f"Test Recall : {test_rec:.4f}")175 print(f"Test F1 : {test_f1:.4f}")176 177 wandb.log(178 {179 "test/loss": test_loss,180 "test/accuracy": test_acc,181 "test/precision": test_prec,182 "test/recall": test_rec,183 "test/f1": test_f1,184 }185 )186 187 wandb.summary["test/accuracy"] = test_acc188 wandb.summary["test/f1"] = test_f1189 wandb.summary["test/precision"] = test_prec190 wandb.summary["test/recall"] = test_rec191 wandb.summary["best_val_acc"] = best_val_accuracy192 193 self.run.finish()194 print(f"\nW&B run finished - {self.run.url}")195 