abhi02072005/JEPA_backend
0
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 