Team Ai
Apppublic

VJyzCELERY/ObjectClassificationPlayground

sourceHugging Facemitupdated 10mo agoView on Hugging Face
0likes
trainer.py359 linesDownload Raw Back to src
1from src.model import Classifier2from src.dataloader import ImageDataset,collate_fn3from torch.utils.data import DataLoader4import torch.optim as optim5import torch.nn.functional as F6from tqdm import tqdm7import matplotlib.pyplot as plt8import torch9import random10import numpy as np11import torch.nn as nn12import time13from sklearn.metrics import (14    confusion_matrix,15    classification_report,16    roc_curve,17    auc18)19from sklearn.preprocessing import label_binarize20def seed_worker(worker_id):21    worker_seed = torch.initial_seed() % 2**3222    np.random.seed(worker_seed)23    random.seed(worker_seed)24 25def model_evaluation(model, val_set, device,batch_size=32,num_workers=0, class_names=None):26 27    model.eval()28    all_preds = []29    all_probs = []30    all_labels = []31    val_loader = DataLoader(32        val_set,33        batch_size=batch_size,34        shuffle=False,35        num_workers=num_workers36    )37    with torch.no_grad():38        for images, labels in val_loader:39            if images.ndim == 4 and images.shape[-1] in (1, 3):40                images = images.permute(0, 3, 1, 2)41            images = images.to(device)42            labels = labels.to(device)43            logits = model(images)44            probs = torch.softmax(logits, dim=1)45            preds = torch.argmax(probs, dim=1)46 47            all_preds.append(preds.cpu().numpy())48            all_probs.append(probs.cpu().numpy())49            all_labels.append(labels.cpu().numpy())50 51    y_true = np.concatenate(all_labels)52    y_pred = np.concatenate(all_preds)53    y_prob = np.concatenate(all_probs)54 55    num_classes = y_prob.shape[1]56 57    if class_names is None:58        class_names = [f"Class {i}" for i in range(num_classes)]59 60    cm = confusion_matrix(y_true, y_pred)61 62    cm_fig, ax = plt.subplots(figsize=(6, 6))63    im = ax.imshow(cm)64 65    ax.set_title("Confusion Matrix")66    ax.set_xlabel("Predicted")67    ax.set_ylabel("True")68    ax.set_xticks(range(num_classes))69    ax.set_yticks(range(num_classes))70    ax.set_xticklabels(class_names, rotation=75)71    ax.set_yticklabels(class_names)72 73    for i in range(num_classes):74        for j in range(num_classes):75            ax.text(j, i, cm[i, j], ha="center", va="center")76 77    plt.tight_layout()78 79    report = classification_report(80        y_true, y_pred,81        target_names=class_names,82        output_dict=True83    )84 85    cr_fig, ax = plt.subplots(figsize=(12, 8))86    ax.axis("off")87 88    table_data = []89    headers = ["Class", "Precision", "Recall", "F1", "Support"]90 91    for cls in class_names:92        row = report[cls]93        table_data.append([94            cls,95            f"{row['precision']:.3f}",96            f"{row['recall']:.3f}",97            f"{row['f1-score']:.3f}",98            int(row['support'])99        ])100 101    accuracy = report["accuracy"]102    macro_avg = report["macro avg"]103    weighted_avg = report["weighted avg"]104 105    table_data.append([106        "Accuracy",107        f"{accuracy:.3f}",108        "",109        "",110        ""111    ])112 113    table_data.append([114        "Macro Avg",115        f"{macro_avg['precision']:.3f}",116        f"{macro_avg['recall']:.3f}",117        f"{macro_avg['f1-score']:.3f}",118        f"{int(macro_avg['support'])}" if 'support' in macro_avg else ""119    ])120 121    table_data.append([122        "Weighted Avg",123        f"{weighted_avg['precision']:.3f}",124        f"{weighted_avg['recall']:.3f}",125        f"{weighted_avg['f1-score']:.3f}",126        f"{int(weighted_avg['support'])}" if 'support' in weighted_avg else ""127    ])128 129    table = ax.table(130        cellText=table_data,131        colLabels=headers,132        loc="center"133    )134 135    table.scale(1, 2)136    ax.set_title("Classification Report")137 138    y_true_bin = label_binarize(y_true, classes=list(range(num_classes)))139 140    roc_fig, ax = plt.subplots(figsize=(6, 6))141 142    for i in range(num_classes):143        fpr, tpr, _ = roc_curve(y_true_bin[:, i], y_prob[:, i])144        roc_auc = auc(fpr, tpr)145        ax.plot(fpr, tpr, label=f"{class_names[i]} (AUC={roc_auc:.3f})")146 147    ax.plot([0, 1], [0, 1], linestyle="--")148    ax.set_xlabel("False Positive Rate")149    ax.set_ylabel("True Positive Rate")150    ax.set_title("ROC-AUC Curve")151    ax.legend()152    ax.grid(True)153 154    return cm_fig, cr_fig, roc_fig155 156class ModelTrainer:157    def __init__(self,model : Classifier,train_set : ImageDataset,val_set : ImageDataset = None, batch_size=32,lr = 1e-3,device='cpu',return_fig=False, seed=None):158        g = torch.Generator()159        if seed is not None:160            g.manual_seed(seed)161        162        self.train_loader = DataLoader(163            train_set, 164            batch_size, 165            shuffle=True, 166            collate_fn=collate_fn,167            worker_init_fn=seed_worker,168            generator=g 169        )170        171        self.device = device172        173        if val_set is not None:174            self.val_loader = DataLoader(175                val_set, 176                batch_size, 177                shuffle=False, 178                collate_fn=collate_fn,179                worker_init_fn=seed_worker180            )181        else:182            self.val_loader = None183        self.class_names = model.classes184        self.model = model185        self.lr = lr186        self.optim = optim.Adam(model.parameters(), lr=lr, weight_decay=1e-4)187        self.optim.zero_grad()188        self.criterion = nn.CrossEntropyLoss()189        self.return_fig=return_fig190        self.best_model_state = None191        self.best_val_acc = 0.0192        self.interrupt=False193    194    def visualize_batch(self, imgs, preds, labels, class_names=None, max_samples=4):195        196        first_image = imgs197        if isinstance(imgs, list):198            imgs = np.stack(imgs, axis=0)199            imgs = torch.from_numpy(imgs).permute(0, 3, 1, 2).float()200 201        imgs_np = imgs.cpu().numpy()202        preds = preds.cpu().numpy()203        labels = labels.cpu().numpy()204 205        batch_size = imgs_np.shape[0]206        indices = random.sample(range(batch_size), min(max_samples, batch_size))207        first_image = first_image[indices[0]]208        fig_pred = plt.figure(figsize=(6 * len(indices), 5))209        grid = fig_pred.add_gridspec(1, len(indices))210 211        for col, idx in enumerate(indices):212            ax = fig_pred.add_subplot(grid[0, col])213            ax.imshow(imgs_np[idx].transpose(1, 2, 0))214 215            if class_names:216                title = f"P: {class_names[preds[idx]]} | T: {class_names[labels[idx]]}"217            else:218                title = f"P: {preds[idx]} | T: {labels[idx]}"219 220            ax.set_title(title)221            ax.axis("off")222 223        fig_pred.tight_layout()224        raw_features = self.model.visualize_feature(first_image,show=False)225        feature_figs = []226 227        for f in raw_features:228            229            if isinstance(f, plt.Figure):230                feature_figs.append(f)231                continue232 233            if hasattr(f, "mode"):234                f = np.array(f)235            h, w = f.shape[:2]236 237            dpi = 100238            fig_w = max(4, w / dpi)239            fig_h = max(4, h / dpi)240            fig = plt.figure(figsize=(fig_w, fig_h), dpi=dpi)241            ax = fig.add_subplot(111)242            ax.imshow(f)243            ax.axis("off")244            feature_figs.append(fig)245            246 247        all_figs = [fig_pred] + feature_figs248        if not self.return_fig:249            plt.show()250        plt.close(fig_pred)251        if self.return_fig:252            return all_figs253        else:254            return None255 256 257    def train_one_epoch(self,epoch):258        self.model.train()259        total_loss = 0260        train_pbar = tqdm(self.train_loader, desc="Training",leave=False)261        correct = 0262        total = 0263        for imgs, labels in train_pbar:264            if self.interrupt:265                break266            labels = labels.to(self.device)267 268            # Forward269            outputs = self.model(imgs)270            loss = self.criterion(outputs, labels)271 272            # Backward273            self.optim.zero_grad()274            loss.backward()275            self.optim.step()276            preds = outputs.argmax(dim=1)277            correct += (preds == labels).sum().item()278            total += labels.size(0)279            total_loss += loss.item()280            train_pbar.set_postfix(acc=correct/total,loss=loss.item())281 282        avg_loss = total_loss / len(self.train_loader)283        avg_acc = correct / total284        return avg_loss,avg_acc285    def train(self, epochs=10, visualize_every=5):286        train_losses=[]287        train_accuracies=[]288        val_losses=[]289        val_accuracies=[]290        for epoch in range(1, epochs + 1):291            train_loss,train_acc = self.train_one_epoch(epoch)292            if self.interrupt:293                return294            train_losses.append(train_loss)295            train_accuracies.append(train_acc)296            if self.val_loader is not None:297                val_loss,val_acc,fig=self.validate(epoch, visualize=(epoch % visualize_every == 0 or epoch == 1))298                if self.interrupt:299                    return300                val_losses.append(val_loss)301                val_accuracies.append(val_acc)302                print(f"Epoch {epoch} Train Loss: {train_loss:.4f} | Train Acc : {train_acc:.4f} | Val Loss : {val_loss:.4f} | Val Acc : {val_acc:.4f}")303                if val_acc > self.best_val_acc:304                    print(f"New best model found at epoch {epoch} (Val Acc: {val_acc:.4f})")305                    self.best_val_acc = val_acc306                    self.best_model_state = {k: v.clone() for k, v in self.model.state_dict().items()}307 308                yield train_loss,train_acc,val_loss,val_acc,fig309            else:310                print(f"Epoch {epoch} Train Loss: {train_loss:.4f} | Train Acc : {train_acc:.4f}")311                yield train_loss,train_acc,None,None,None312        if self.best_model_state is not None:313            self.model.load_state_dict(self.best_model_state)314            print(f"Best model (Val Acc: {self.best_val_acc:.4f}) loaded into trainer.model")315        yield train_losses,train_accuracies,val_losses,val_accuracies,None316 317    def validate(self,epoch, visualize=False):318        if self.val_loader is None:319            return320 321        self.model.eval()322        total_loss = 0323        correct = 0324        total = 0325 326        val_imgs_display = None327        val_preds_display = None328        val_labels_display = None329 330        val_pbar = tqdm(self.val_loader, desc="Validation",leave=False)331        fig = None332        with torch.no_grad():333            for imgs, labels in val_pbar:334                if self.interrupt:335                    break336                labels = labels.to(self.device)337 338                outputs = self.model(imgs)339                loss = self.criterion(outputs, labels)340                total_loss += loss.item()341 342                preds = outputs.argmax(dim=1)343                correct += (preds == labels).sum().item()344                total += labels.size(0)345 346                if visualize and val_imgs_display is None:347                    val_imgs_display = imgs348                    val_preds_display = preds349                    val_labels_display = labels350 351                val_pbar.set_postfix(loss=loss.item(), acc=correct / total)352 353        avg_loss = total_loss / len(self.val_loader)354        acc = correct / total355 356        if visualize and val_imgs_display is not None:357            fig = self.visualize_batch(val_imgs_display, val_preds_display, val_labels_display, self.class_names)358 359        return avg_loss,acc,fig