CoolFace
Apppublic

Haldhar/Malware-classification

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes
train.py195 linesDownload Raw Back to src
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