gaaaaaaaaaaa's picture
Document Reasoning Training Push for VAST
a822d67 verified
Raw History Blame Contribute Delete
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()
@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", {})