VJyzCELERY/ObjectClassificationPlayground
0
1import torch2import torch.nn as nn3import cv24import numpy as np5from dataclasses import dataclass6from skimage.feature import hog,local_binary_pattern7import itertools8import torch.nn.functional as F9import matplotlib.pyplot as plt10import os11import io12from PIL import Image13 14@dataclass15class Config:16 img_size=(256,256)17 in_channels=318 fc_num_layers=319 conv_hidden_dim=220 conv_kernel_size=321 dropout=0.222 classical_downsample=123 # HOG24 hog_orientations = 925 hog_pixels_per_cell = (16, 16)26 hog_cells_per_block = (2, 2)27 hog_block_norm = 'L2-Hys'28 29 # Canny30 canny_sigma = 1.031 canny_low = 10032 canny_high = 20033 34 # Gaussian35 gaussian_ksize = (3, 3)36 gaussian_sigmaX = 1.037 gaussian_sigmaY = 1.038 39 # Harris corners40 harris_block_size = 241 harris_ksize = 342 harris_k = 0.0443 44 45 # LBP46 lbp_P = 8 47 lbp_R = 1 48 49 # Gabor filters50 gabor_ksize = 2151 gabor_sigma = 552 gabor_theta = 053 gabor_lambda = 1054 gabor_gamma = 0.555 56 # Sobel57 sobel_ksize=358 59 60class CNNFeatureExtractor(nn.Module):61 def __init__(self,config : Config):62 super().__init__()63 layers = []64 self.in_channels = config.in_channels65 in_channel = config.in_channels66 self.img_size = config.img_size67 out_channel = 3268 for i in range(config.conv_hidden_dim):69 layers.append(nn.Conv2d(in_channels=in_channel,out_channels=out_channel,kernel_size=config.conv_kernel_size,stride=1,padding=config.conv_kernel_size // 2))70 layers.append(nn.BatchNorm2d(out_channel))71 layers.append(nn.ReLU())72 layers.append(nn.MaxPool2d((2,2)))73 in_channel=out_channel74 out_channel*=275 self.layers = nn.Sequential(*layers)76 def get_device(self):77 return next(self.parameters()).device78 def forward(self,x,**kwargs):79 if isinstance(x, list):80 if isinstance(x[0], np.ndarray):81 x = np.stack(x, axis=0) 82 if isinstance(x,np.ndarray):83 if len(x.shape) == 2:84 x = x[:, :, None] 85 x = np.expand_dims(x, 0)86 x = x.transpose(2, 0, 1) 87 elif len(x.shape) == 3:88 x = x.transpose(2, 0, 1)89 x = np.expand_dims(x, 0)90 elif x.ndim == 4:91 x = x.transpose(0, 3, 1, 2) # Change to (B,C,H,W)92 x = torch.from_numpy(x).float()93 elif isinstance(x, torch.Tensor):94 if x.ndim == 3:95 x = x.unsqueeze(0)96 x=x.to(self.get_device())97 return self.layers(x) # Always expects (B,C,H,W)98 def output(self):99 self.eval()100 101 with torch.no_grad():102 x = torch.zeros(103 (1, self.in_channels, self.img_size[1], self.img_size[0]),104 device=self.get_device()105 )106 107 out = self(x)108 109 return out110 def visualize(111 self,112 input_image,113 max_channels=8,114 couple=False,115 show=True,116 **kwargs117 ):118 self.eval()119 device = self.get_device()120 121 if isinstance(input_image, np.ndarray):122 x = torch.from_numpy(input_image).permute(2, 0, 1).float().unsqueeze(0).to(device)123 elif isinstance(input_image, torch.Tensor):124 x = input_image.unsqueeze(0).to(device) if input_image.ndim == 3 else input_image.to(device)125 else:126 raise TypeError("input_image must be np.ndarray or torch.Tensor")127 128 conv_layers = [129 (name, module)130 for name, module in self.named_modules()131 if isinstance(module, nn.ReLU)132 ]133 134 all_layer_images = []135 136 for name, layer in conv_layers:137 activations = []138 139 def hook_fn(module, input, output):140 activations.append(output.detach().cpu())141 142 handle = layer.register_forward_hook(hook_fn)143 _ = self(x)144 handle.remove()145 146 act = activations[0][0] # (C, H, W)147 C, H, W = act.shape148 149 # --------------------------------------------------150 # COUPLED RGB VISUALIZATION151 # --------------------------------------------------152 if couple:153 max_rgb = max_channels // 3154 num_rgb = min(C // 3, max_rgb)155 rem = min(C - num_rgb * 3, max_channels - num_rgb * 3)156 157 total_tiles = num_rgb + rem158 cols = min(4, total_tiles)159 rows = int(np.ceil(total_tiles / cols))160 161 fig, axes = plt.subplots(162 rows, cols,163 figsize=(3 * cols, 3 * rows)164 )165 166 axes = np.atleast_2d(axes)167 168 tile_idx = 0169 170 # ---------------------------171 # RGB COUPLED CHANNELS172 # ---------------------------173 for i in range(num_rgb):174 r = tile_idx // cols175 c = tile_idx % cols176 177 rgb = act[i*3:(i+1)*3].clone()178 179 for ch in range(3):180 v = rgb[ch]181 rgb[ch] = (v - v.min()) / (v.max() - v.min() + 1e-8)182 183 rgb = rgb.permute(1, 2, 0).numpy()184 185 axes[r, c].imshow(rgb)186 axes[r, c].axis("off")187 axes[r, c].set_title(f"RGB {i*3}-{i*3+2}", fontsize=9)188 189 tile_idx += 1190 191 start = num_rgb * 3192 for j in range(rem):193 r = tile_idx // cols194 c = tile_idx % cols195 196 ch = act[start + j]197 ch = (ch - ch.min()) / (ch.max() - ch.min() + 1e-8)198 199 axes[r, c].imshow(ch, cmap="gray")200 axes[r, c].axis("off")201 axes[r, c].set_title(f"Ch {start + j}", fontsize=9)202 203 tile_idx += 1204 205 for idx in range(tile_idx, rows * cols):206 r = idx // cols207 c = idx % cols208 axes[r, c].axis("off")209 210 fig.suptitle(f"Layer: {name} (Coupled RGB + Grayscale)", fontsize=14)211 plt.tight_layout()212 213 # --------------------------------------------------214 # STANDARD GRAYSCALE VISUALIZATION215 # --------------------------------------------------216 else:217 num_channels = min(C, max_channels)218 cols = min(8, num_channels)219 rows = int(np.ceil(num_channels / cols))220 221 fig, axes = plt.subplots(222 rows, cols,223 figsize=(3 * cols, 3 * rows)224 )225 226 axes = np.atleast_2d(axes)227 228 for idx in range(num_channels):229 r = idx // cols230 c = idx % cols231 axes[r, c].imshow(act[idx], cmap="gray")232 axes[r, c].axis("off")233 234 for idx in range(num_channels, rows * cols):235 r = idx // cols236 c = idx % cols237 axes[r, c].axis("off")238 239 fig.suptitle(f"Layer: {name}", fontsize=14)240 plt.tight_layout()241 242 if show:243 plt.show()244 245 buf = io.BytesIO()246 fig.savefig(buf, format="png", dpi=150, bbox_inches="tight")247 buf.seek(0)248 img = Image.open(buf).convert("RGB")249 all_layer_images.append(np.array(img))250 plt.close(fig)251 252 return all_layer_images253 254class ClassicalFeatureExtractor(nn.Module):255 def __init__(self, config : Config):256 super().__init__()257 self.img_size = config.img_size # (H, W)258 self.hog_orientations = config.hog_orientations259 self.num_downsample = config.classical_downsample260 self.config = config261 self.device = 'cpu'262 self.convolution=None263 264 def get_device(self):265 return next(self.parameters()).device if len(list(self.parameters())) > 0 else self.device266 267 def render_subplots(self,items, max_cols=8, figsize_per_cell=3):268 n = len(items)269 cols = min(max_cols, n)270 rows = int(np.ceil(n / cols))271 fig, axes = plt.subplots(272 rows, cols,273 figsize=(cols * figsize_per_cell, rows * figsize_per_cell)274 )275 axes = np.atleast_2d(axes)276 for idx, (img, title, cmap) in enumerate(items):277 r = idx // cols278 c = idx % cols279 ax = axes[r, c]280 ax.imshow(img, cmap=cmap)281 ax.set_title(title, fontsize=9)282 ax.axis("off")283 for idx in range(n, rows * cols):284 r = idx // cols285 c = idx % cols286 axes[r, c].axis("off")287 288 plt.tight_layout()289 return fig290 291 def extract_features(self, img,visualize=False,**kwargs):292 cfg = self.config293 # Convert to grayscale294 gray = cv2.cvtColor((img*255).astype(np.uint8), cv2.COLOR_RGB2GRAY)295 for _ in range(self.num_downsample):296 gray = cv2.pyrDown(gray)297 gray = cv2.GaussianBlur(gray, cfg.gaussian_ksize, sigmaX=cfg.gaussian_sigmaX, sigmaY=cfg.gaussian_sigmaY)298 valid_H, valid_W = gray.shape[:2]299 300 301 feature_list = []302 vis_items=[]303 # DEPRECATED304 # H, W = gray.shape305 # cell_h, cell_w = cfg.hog_pixels_per_cell306 # block_h, block_w = cfg.hog_cells_per_block307 308 # min_h = cell_h * block_h309 # min_w = cell_w * block_w310 # use_hog = False311 # # 1. HOG312 # if use_hog:313 # hog_descriptors, hog_image = hog(314 # gray,315 # orientations=cfg.hog_orientations,316 # pixels_per_cell=cfg.hog_pixels_per_cell,317 # cells_per_block=cfg.hog_cells_per_block,318 # block_norm=cfg.hog_block_norm,319 # visualize=True,320 # feature_vector=False321 # )322 323 # hog_cells = hog_descriptors.mean(axis=(2, 3))324 325 # cell_h, cell_w = cfg.hog_pixels_per_cell326 # hog_pixel = np.repeat(327 # np.repeat(hog_cells, cell_h, axis=0),328 # cell_w, axis=1329 # )330 # hog_pixel = hog_pixel[:gray.shape[0], :gray.shape[1]]331 # hog_energy = np.sum(hog_pixel, axis=2)332 # dominant_bin = np.argmax(hog_pixel, axis=2)333 # dominant_strength = np.max(hog_pixel, axis=2)334 # dominant_weighted = dominant_bin * dominant_strength335 # valid_H, valid_W = hog_pixel.shape[:2]336 # if visualize:337 # vis_items.append((hog_energy, "HOG Energy",'gray'))338 # vis_items.append((dominant_bin, "HOG Dominant Bin",'hsv'))339 # vis_items.append((dominant_weighted, "HOG Weighted Dominant Bin",'gray'))340 # vis_items.append((hog_image[:valid_H, :valid_W], f"HoG",'gray'))341 # for b in range(hog_pixel.shape[2]):342 # feature_list.append(hog_pixel[:, :, b])343 344 345 # 2. Canny edges346 edges = cv2.Canny(gray, cfg.canny_low, cfg.canny_high) / 255.0347 feature_list.append(edges[:valid_H, :valid_W])348 if visualize:349 vis_items.append((edges[:valid_H, :valid_W], "Canny Edge", "gray"))350 # 3. Harris corners351 harris = cv2.cornerHarris(gray, blockSize=cfg.harris_block_size, ksize=cfg.harris_ksize, k=cfg.harris_k)352 harris = cv2.dilate(harris, None)353 harris = np.clip(harris, 0, 1)354 feature_list.append(harris[:valid_H, :valid_W])355 if visualize:356 vis_items.append((harris[:valid_H, :valid_W], "Harris Corner", "gray"))357 358 # 4. LBP359 lbp = local_binary_pattern(gray, P=cfg.lbp_P, R=cfg.lbp_R, method='uniform')360 lbp = lbp / lbp.max() if lbp.max() != 0 else lbp361 # feature_list.append(lbp.ravel())362 feature_list.append(lbp[:valid_H, :valid_W])363 if visualize:364 # figs.append(plot_feature(lbp[:valid_H, :valid_W], "LBP"))365 vis_items.append((lbp[:valid_H, :valid_W], "LBP", "gray"))366 # 5. Gabor filter367 for theta in [0, np.pi/4, np.pi/2]:368 kernel = cv2.getGaborKernel(369 (cfg.gabor_ksize, cfg.gabor_ksize),370 cfg.gabor_sigma, theta,371 cfg.gabor_lambda, cfg.gabor_gamma372 )373 g = cv2.filter2D(gray, cv2.CV_32F, kernel)374 g = np.abs(g)375 g /= g.max() + 1e-8376 feature_list.append(g[:valid_H, :valid_W])377 if visualize:378 vis_items.append((g[:valid_H, :valid_W], f"Gabor θ={theta:.2f}", "gray"))379 # 6. Sobel380 sobelx = cv2.Sobel(gray, cv2.CV_32F, 1, 0, ksize=cfg.sobel_ksize)381 sobely = cv2.Sobel(gray, cv2.CV_32F, 0, 1, ksize=cfg.sobel_ksize)382 383 sobelx = np.abs(sobelx)384 sobely = np.abs(sobely)385 386 sobelx /= sobelx.max() + 1e-8387 sobely /= sobely.max() + 1e-8388 389 feature_list.append(sobelx[:valid_H, :valid_W])390 feature_list.append(sobely[:valid_H, :valid_W])391 if visualize:392 vis_items.append((sobelx[:valid_H, :valid_W], "Sobel X",'gray'))393 vis_items.append((sobely[:valid_H, :valid_W], "Sobel Y",'gray'))394 # 7. Laplacian395 lap = cv2.Laplacian(gray, cv2.CV_32F)396 lap = np.abs(lap)397 lap /= lap.max() + 1e-8398 399 feature_list.append(lap[:valid_H, :valid_W])400 401 if visualize:402 vis_items.append((lap[:valid_H, :valid_W], "Laplacian",'gray'))403 404 # 8. Gradient Magnitude405 gx = cv2.Sobel(gray, cv2.CV_32F, 1, 0, ksize=cfg.sobel_ksize)406 gy = cv2.Sobel(gray, cv2.CV_32F, 0, 1, ksize=cfg.sobel_ksize)407 408 grad_mag = np.sqrt(gx**2 + gy**2)409 grad_mag /= grad_mag.max() + 1e-8410 411 feature_list.append(grad_mag[:valid_H, :valid_W])412 413 if visualize:414 vis_items.append((grad_mag[:valid_H, :valid_W], "Gradient Magnitude",'gray'))415 416 # Stack all features along channel axis417 features = np.stack(feature_list, axis=0)418 if visualize:419 return features.astype(np.float32),[self.render_subplots(vis_items, max_cols=8)]420 return features.astype(np.float32)421 422 423 def forward(self, x, **kwargs):424 if isinstance(x, list):425 x = np.stack(x, axis=0)426 427 if isinstance(x, torch.Tensor):428 x = x.cpu().numpy()429 430 if isinstance(x, np.ndarray):431 if x.ndim == 3: 432 x = x[None]433 elif x.ndim != 4:434 raise ValueError(435 f"Expected input of shape HWC or BHWC, got {x.shape}"436 )437 feats = []438 for img in x:439 if img.shape[2] != 3:440 img = np.repeat(img[:, :, None], 3, axis=2)441 feats.append(self.extract_features(img))442 443 feats = np.stack(feats, axis=0)444 feats = torch.from_numpy(feats).float().to(self.get_device())445 return feats446 447 def visualize(self, img, show_original=True,show=True):448 if img.ndim != 3 or img.shape[2] != 3:449 img = np.repeat(img[:, :, None], 3, axis=2)450 451 outputs = [] 452 453 def fig_to_pil(fig):454 buf = io.BytesIO()455 fig.savefig(buf, format="png", dpi=150, bbox_inches="tight")456 buf.seek(0)457 458 pil_img = Image.open(buf).copy()459 460 buf.close()461 plt.close(fig)462 463 return pil_img464 465 if show_original:466 fig = plt.figure(figsize=(4, 4))467 plt.imshow(img)468 plt.title("Original")469 plt.axis("off")470 if show:471 plt.show() 472 outputs.append(fig_to_pil(fig)) 473 feature_stack,figs = self.extract_features(img,visualize=True)474 if show:475 plt.show() 476 for fig in figs:477 outputs.append(fig_to_pil(fig)) 478 479 return outputs480 481 482 def output(self):483 dummy = np.zeros(484 (self.img_size[1], self.img_size[0], 3),485 dtype=np.float32486 )487 488 feats = self.forward(dummy)489 490 return feats491 492 493 494class FullyConnectedHead(nn.Module):495 def __init__(self,in_features,classes,config:Config):496 super().__init__()497 num_classes = len(classes)498 self.classes = classes499 layers = []500 hidden_dim =1024501 for _ in range(config.fc_num_layers):502 layers.append(nn.Linear(in_features, hidden_dim))503 layers.append(nn.BatchNorm1d(hidden_dim))504 layers.append(nn.ReLU())505 layers.append(nn.Dropout(config.dropout))506 507 in_features = hidden_dim508 hidden_dim = max(hidden_dim // 2, num_classes * 2)509 layers.append(nn.Linear(in_features,num_classes))510 self.layers = nn.Sequential(*layers)511 def get_device(self):512 return next(self.parameters()).device513 def forward(self,x : torch.Tensor,**kwargs):514 x=x.to(self.get_device())515 return self.layers(x)516 517class Classifier(nn.Module):518 def __init__(self,backbone,classes,config : Config):519 super().__init__()520 self.config=config521 self.classes=classes522 self.backbone = backbone523 self.flatten = nn.Flatten()524 feat = backbone.output()525 flat = self.flatten(feat)526 in_features = flat.shape[1]527 self.head = FullyConnectedHead(in_features,classes,config)528 def get_device(self):529 return next(self.parameters()).device530 531 @torch.no_grad()532 def predict(self, x):533 self.eval()534 target_size = self.config.img_size535 x = cv2.resize(x, target_size)536 logits = self.forward(x) 537 probs = torch.softmax(logits,dim=1)538 pred_idx = torch.argmax(probs, dim=1).item()539 540 return self.classes[pred_idx]541 542 def forward(self,x,**kwargs):543 feat = self.backbone(x,**kwargs)544 feat = self.flatten(feat,**kwargs)545 return self.head(feat,**kwargs)546 def visualize_feature(self,img,return_img=True,**kwargs):547 target_size = self.config.img_size548 img = cv2.resize(img, target_size)549 if return_img:550 return self.backbone.visualize(img,**kwargs)551 else:552 self.backbone.visualize(img,**kwargs)553 def save(self, path: str):554 os.makedirs(os.path.dirname(path), exist_ok=True)555 torch.save({556 'model_state_dict': self.state_dict(),557 'classes': self.classes,558 'config': self.config559 }, path)560 print(f"Model saved to {path}")561 562@staticmethod563def load(path: str, backbone_class, device='cpu'):564 checkpoint = torch.load(path, map_location=device,weights_only=False)565 config = checkpoint['config']566 classes = checkpoint['classes']567 backbone = backbone_class(config).to(device)568 model = Classifier(backbone, classes, config).to(device)569 model.load_state_dict(checkpoint['model_state_dict'])570 model.eval()571 print(f"Model loaded from {path}")572 return model