Download models/visual.py from gaaaaaaaaaaa/multimodal-reasoning: direct link, hf CLI and curl.
- Browser
- Download file 8.54 kB
-
https://huggingface.co/gaaaaaaaaaaa/multimodal-reasoning/resolve/main/models/visual.py
- Command line
-
hf download hf://gaaaaaaaaaaa/multimodal-reasoning/models/visual.py
-
curl -L -o visual.py https://huggingface.co/gaaaaaaaaaaa/multimodal-reasoning/resolve/main/models/visual.py
8.54 kB
| # 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() | |
| 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", {}) |