Team Ai
Apppublic

VJyzCELERY/ObjectClassificationPlayground

sourceHugging Facemitupdated 10mo agoView on Hugging Face
0likes
model.py572 linesDownload Raw Back to src
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