# VisualEncoderHead.py import os import json import torch import torch.nn as nn import utils.configs_loader as ProjectConfigs from PIL import Image from transformers import AutoModel, AutoProcessor # ------------------------------------------------------------ # Config # ------------------------------------------------------------ SIGLIP2_LARGE_CKPT = "google/siglip2-large-patch16-384" # ------------------------------------------------------------ # Visual Encoder: SigLIP2 Large # ------------------------------------------------------------ class VisualEncoder(nn.Module): """ SigLIP2 Large vision encoder wrapper. Output: image_features: [B, N, vision_dim] For google/siglip2-large-patch16-384: image size: 384x384 patch size: 16 patch tokens: roughly 24 * 24 = 576 """ def __init__( self, model_name: str = SIGLIP2_LARGE_CKPT, device: str | torch.device | None = None, dtype: torch.dtype | None = None, freeze: bool = True, ): super().__init__() self.device = torch.device(device or ProjectConfigs.get_device()) if dtype is None: if self.device.type == "cuda": dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16 else: dtype = torch.float32 self.dtype = dtype self.model_name = model_name self.processor = AutoProcessor.from_pretrained(model_name) self.encoder = AutoModel.from_pretrained( model_name, torch_dtype=self.dtype, attn_implementation="sdpa", ).to(self.device) self.encoder.eval() if freeze: self.freeze() # Usually SigLIP/SigLIP2 has this available. self.vision_dim = self.encoder.config.vision_config.hidden_size def freeze(self): for p in self.encoder.parameters(): p.requires_grad = False self.encoder.eval() def unfreeze(self): for p in self.encoder.parameters(): p.requires_grad = True self.encoder.train() @torch.no_grad() def encode_images(self, images): """ Args: images: - PIL.Image - list[PIL.Image] - torch.Tensor already shaped [B, C, H, W] Returns: image_features: [B, N, vision_dim] """ if isinstance(images, torch.Tensor): pixel_values = images.to(self.device, dtype=self.dtype) else: if isinstance(images, Image.Image): images = [images] inputs = self.processor( images=images, return_tensors="pt", ) pixel_values = inputs["pixel_values"].to(self.device, dtype=self.dtype) outputs = self.encoder.vision_model( pixel_values=pixel_values, output_hidden_states=False, return_dict=True, ) # [B, N, D] image_features = outputs.last_hidden_state return image_features def forward(self, images): return self.encode_images(images) # ------------------------------------------------------------ # Projector Head # ------------------------------------------------------------ class ProjectorHead(nn.Module): """ Projects SigLIP2 visual tokens into Qwen hidden size. Example: SigLIP2 Large vision_dim -> Qwen2.5-1.5B hidden_dim For Qwen2.5-1.5B, projection_dim is usually 1536. """ def __init__( self, embed_dim: int, projection_dim: int = 1536, hidden_mult: int = 4, dropout: float = 0.0, device: str | torch.device | None = None, dtype: torch.dtype | None = None, ): super().__init__() self.device = torch.device(device or ProjectConfigs.get_device()) if dtype is None: if self.device.type == "cuda": dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16 else: dtype = torch.float32 self.dtype = dtype hidden_dim = embed_dim * hidden_mult self.net = nn.Sequential( nn.LayerNorm(embed_dim), nn.Linear(embed_dim, hidden_dim), nn.GELU(), nn.Dropout(dropout), nn.Linear(hidden_dim, projection_dim), ).to(self.device, dtype=self.dtype) self.projection_dim = projection_dim def forward(self, x: torch.Tensor) -> torch.Tensor: """ Args: x: [B, N, embed_dim] Returns: projected: [B, N, projection_dim] """ x = x.to(self.device, dtype=self.dtype) return self.net(x) # ------------------------------------------------------------ # Compressor Layer # ------------------------------------------------------------ class QueryResampler(nn.Module): def __init__( self, dim, num_queries=64, num_heads=8, device=None, dtype=None, ): super().__init__() self.device = torch.device(device or ProjectConfigs.get_device()) if dtype is None: if self.device.type == "cuda": dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16 else: dtype = torch.float32 self.dtype = dtype self.queries = nn.Parameter(torch.randn(1, num_queries, dim) * 0.02) self.attn = nn.MultiheadAttention(dim, num_heads, batch_first=True) self.norm_q = nn.LayerNorm(dim) self.norm_x = nn.LayerNorm(dim) self.ff = nn.Sequential( nn.LayerNorm(dim), nn.Linear(dim, dim * 4), nn.GELU(), nn.Linear(dim * 4, dim), ) self.to(self.device, dtype=self.dtype) def forward(self, x): x = x.to(self.device, dtype=self.dtype) B = x.size(0) q = self.queries.expand(B, -1, -1) y, _ = self.attn( query=self.norm_q(q), key=self.norm_x(x), value=self.norm_x(x), need_weights=False, ) y = y + q y = y + self.ff(y) return y # ------------------------------------------------------------ # Visual Layer # ------------------------------------------------------------ class VisualLayer(nn.Module): def __init__( self, freeze_encoder=False, ): super().__init__() self.encoder = VisualEncoder(freeze=freeze_encoder) self.projector = ProjectorHead( embed_dim=self.encoder.vision_dim, dropout=0.1 ) self.compressor = QueryResampler( dim=self.projector.projection_dim ) def forward(self, x): image_features = self.encoder(x) projected_x = self.projector(image_features) compressed_x = self.compressor(projected_x) return compressed_x def save(self, path="outputs/visual_projection_layer"): os.makedirs(path, exist_ok=True) ckpt = { "projector": self.projector.state_dict(), "compressor": self.compressor.state_dict(), "meta": { "vision_model": self.encoder.model_name, "vision_dim": self.encoder.vision_dim, "projection_dim": self.projector.projection_dim, "dtype": str(self.projector.dtype), }, } torch.save(ckpt, os.path.join(path, "visual_layer.pt")) with open(os.path.join(path, "meta.json"), "w", encoding="utf-8") as f: json.dump(ckpt["meta"], f, indent=2) print(f"[VisualLayer] saved to {path}") def load( self, path="outputs/visual_projection_layer", map_location=None, strict=True, ): ckpt_path = os.path.join(path, "visual_layer.pt") if map_location is None: map_location = self.encoder.device ckpt = torch.load(ckpt_path, map_location=map_location) self.projector.load_state_dict(ckpt["projector"], strict=strict) self.compressor.load_state_dict(ckpt["compressor"], strict=strict) self.projector.to(self.encoder.device, dtype=self.projector.dtype) self.compressor.to(self.encoder.device, dtype=self.projector.dtype) print(f"[VisualLayer] loaded from {ckpt_path}") return ckpt.get("meta", {})