VJyzCELERY/ObjectClassificationPlayground
0
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