CoolFace
Apppublic

abhi02072005/JEPA_backend

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
encoder.py174 linesDownload Raw Back to models
1"""2encoder.py — Frozen (or partially fine-tuned) ViT-B/16 encoder.3 4Upgrade 1: partial_unfreeze() exposes last 4 transformer blocks5           for domain adaptation at a low learning rate.6"""7import sys8from pathlib import Path9 10import numpy as np11import torch12import torch.nn as nn13import timm14from torchvision import transforms15from typing import Tuple, List16 17sys.path.insert(0, str(Path(__file__).parent.parent))18import config as cfg19 20 21_IMAGENET_MEAN = (0.485, 0.456, 0.406)22_IMAGENET_STD  = (0.229, 0.224, 0.225)23 24 25class ViTEncoder(nn.Module):26    """27    ViT-B/16 visual encoder.28 29    Default: first 8 blocks frozen, last 4 blocks trainable (Upgrade 1).30    Set ENCODER_FINETUNE=False in config to fully freeze (original behaviour).31 32    Returns:33      cls_token  : [B, 768]     — global frame representation34      patch_tokens: [B, 196, 768] — local spatial tokens35    """36 37    def __init__(38        self,39        model_name: str = cfg.ENCODER_MODEL,40        device: str     = cfg.DEVICE,41        finetune: bool  = cfg.ENCODER_FINETUNE,42        freeze_blocks: int = cfg.ENCODER_FREEZE_BLOCKS,43    ):44        super().__init__()45        self.device = device46 47        self.model = timm.create_model(48            model_name, pretrained=True, num_classes=049        ).to(device)50 51        # Start fully frozen52        for param in self.model.parameters():53            param.requires_grad = False54 55        # Upgrade 1 — selectively unfreeze last N blocks56        if finetune:57            self.partial_unfreeze(freeze_blocks)58 59        # AUTO-LOAD fine-tuned encoder weights if they were saved during training60        ft_path = cfg.CHECKPOINTS_DIR / "encoder_finetune.pt"61        if ft_path.exists():62            self._load_finetuned(ft_path, device)63 64        self.transform = transforms.Compose([65            transforms.ToTensor(),66            transforms.Normalize(mean=_IMAGENET_MEAN, std=_IMAGENET_STD),67        ])68 69        self.eval()   # eval by default; trainer sets .train() on unfrozen subset70 71    # ─────────────────────────────────────────────────────────────────72    # Upgrade 1 — Partial Domain Adaptation73    # ─────────────────────────────────────────────────────────────────74    def partial_unfreeze(self, n_freeze: int = 8) -> None:75        """76        Freeze first n_freeze transformer blocks; unfreeze the rest.77        Also unfreezes the final LayerNorm so the output is adaptable.78        """79        total = len(self.model.blocks)80        for i, block in enumerate(self.model.blocks):81            for param in block.parameters():82                param.requires_grad = (i >= n_freeze)83 84        # Final norm85        for param in self.model.norm.parameters():86            param.requires_grad = True87 88        n_unfrozen = total - n_freeze89        trainable  = sum(p.numel() for p in self.model.parameters() if p.requires_grad)90        total_p    = sum(p.numel() for p in self.model.parameters())91        print(f"[Encoder] {n_freeze}/{total} blocks frozen | "92              f"{n_unfrozen} blocks trainable "93              f"({trainable:,} / {total_p:,} params, "94              f"{100*trainable/total_p:.1f}% unfrozen)")95 96    def trainable_parameters(self):97        """Yield only the unfrozen encoder parameters (for optimizer)."""98        return (p for p in self.model.parameters() if p.requires_grad)99 100    def _load_finetuned(self, path: Path, device: str):101        """Load partially fine-tuned encoder weights saved during training."""102        try:103            saved = torch.load(str(path), map_location=device)104            loaded = 0105            model_dict = dict(self.model.named_parameters())106            for name, param_data in saved.items():107                # Strip 'model.' prefix if present108                clean = name.replace('model.', '') if name.startswith('model.') else name109                if clean in model_dict:110                    model_dict[clean].data.copy_(param_data.data)111                    loaded += 1112            print(f"[Encoder] Loaded {loaded} fine-tuned parameter tensors from {path.name}")113        except Exception as e:114            print(f"[Encoder] Warning: could not load fine-tuned weights: {e}")115 116    # ─────────────────────────────────────────────────────────────────117    # Forward118    # ─────────────────────────────────────────────────────────────────119    def forward(self, pixel_values: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:120        """121        Args:122            pixel_values: [B, 3, 224, 224] normalised tensor123 124        Returns:125            cls_tokens:    [B, 768]126            patch_tokens:  [B, 196, 768]127        """128        features = self.model.forward_features(pixel_values)129        # timm ViT: features[:, 0] = CLS, features[:, 1:] = patches130        cls_tokens   = features[:, 0]         # [B, 768]131        patch_tokens = features[:, 1:]        # [B, 196, 768]132        return cls_tokens, patch_tokens133 134    # ─────────────────────────────────────────────────────────────────135    # NumPy helpers136    # ─────────────────────────────────────────────────────────────────137    @torch.no_grad()138    def encode_frame_np(139        self, frame_rgb: np.ndarray140    ) -> Tuple[np.ndarray, np.ndarray]:141        """142        Encode a single frame.143 144        Args:145            frame_rgb: uint8 [224, 224, 3]146 147        Returns:148            cls_emb   : float32 [768]149            patch_emb : float32 [196, 768]150        """151        t = self.transform(frame_rgb).unsqueeze(0).to(self.device)152        cls, patches = self.forward(t)153        return cls[0].cpu().numpy(), patches[0].cpu().numpy()154 155    @torch.no_grad()156    def encode_batch_numpy(157        self,158        frames: np.ndarray,        # [N, 224, 224, 3] uint8159        batch_size: int = 16,160    ) -> Tuple[np.ndarray, np.ndarray]:161        """162        Encode a numpy batch. Returns:163            cls   : [N, 768]164            patches: [N, 196, 768]165        """166        all_cls, all_patches = [], []167        for i in range(0, len(frames), batch_size):168            batch_np = frames[i : i + batch_size]169            tensors  = torch.stack([self.transform(f) for f in batch_np]).to(self.device)170            cls, pat = self.forward(tensors)171            all_cls.append(cls.cpu().numpy())172            all_patches.append(pat.cpu().numpy())173        return np.concatenate(all_cls, 0), np.concatenate(all_patches, 0)174