Hritik045678's picture
Upload folder using huggingface_hub
296a506 verified
Raw
History Blame Contribute Delete
10.9 kB
"""
Frox AI Morph 1.1 — Vision Encoder
Encodes images into token embeddings the LLM can read.
Architecture: ViT + MLP projector → text embedding space
Improvements over Morph 1.0:
- Dynamic high-res tiling: images up to 1344px are split into up to
4×336px tiles + 1 global thumbnail (multi-tile high-resolution
encoding, similar in spirit to other tiled vision-language setups)
- 2D RoPE option for position encoding (better extrapolation than
learned absolute positions, kept learned as default for stability)
- Attention pooling head available for embedding-only use cases
- Interpolatable position embeddings (supports images the model
wasn't originally trained on the exact resolution of)
"""
from __future__ import annotations
import math
from typing import List, Optional, Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from model.attention.gqa import MorphRMSNorm
class VisionPatchEmbedding(nn.Module):
"""
Splits an image into patches and projects to hidden dim.
336px image / 14px patches = 24×24 = 576 patches (tokens).
NEW 1.1: interpolate_pos_encoding() allows using a different
input resolution than the model was initialized with.
"""
def __init__(
self,
image_size: int = 336,
patch_size: int = 14,
num_channels: int = 3,
hidden_size: int = 1024,
):
super().__init__()
self.image_size = image_size
self.patch_size = patch_size
self.num_patches = (image_size // patch_size) ** 2
self.num_patches_side = image_size // patch_size
self.projection = nn.Conv2d(
num_channels, hidden_size,
kernel_size=patch_size, stride=patch_size, bias=False,
)
self.cls_token = nn.Parameter(torch.zeros(1, 1, hidden_size))
self.position_embedding = nn.Embedding(self.num_patches + 1, hidden_size)
nn.init.trunc_normal_(self.cls_token, std=0.02)
nn.init.trunc_normal_(self.position_embedding.weight, std=0.02)
def interpolate_pos_encoding(self, x: torch.Tensor, h: int, w: int) -> torch.Tensor:
"""
Bicubic-interpolate learned position embeddings to a new grid size.
Lets the same weights serve 336px, 448px, 672px tiles, etc.
"""
num_patches = x.shape[1] - 1
num_positions = self.position_embedding.weight.shape[0] - 1
if num_patches == num_positions and h == w:
return self.position_embedding(torch.arange(x.shape[1], device=x.device))
class_pos = self.position_embedding.weight[:1]
patch_pos = self.position_embedding.weight[1:]
dim = x.shape[-1]
orig_side = int(math.sqrt(num_positions))
new_h, new_w = h // self.patch_size, w // self.patch_size
patch_pos = patch_pos.reshape(1, orig_side, orig_side, dim).permute(0, 3, 1, 2)
patch_pos = F.interpolate(
patch_pos, size=(new_h, new_w), mode="bicubic", align_corners=False
)
patch_pos = patch_pos.permute(0, 2, 3, 1).reshape(1, new_h * new_w, dim)
return torch.cat([class_pos.unsqueeze(0), patch_pos], dim=1).squeeze(0)
def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:
"""pixel_values: [B, 3, H, W] → [B, num_patches+1, hidden_size]"""
B, _, H, W = pixel_values.shape
x = self.projection(pixel_values) # [B, hidden, H/p, W/p]
x = rearrange(x, "b c h w -> b (h w) c") # [B, num_patches, hidden]
cls = self.cls_token.expand(B, -1, -1)
x = torch.cat([cls, x], dim=1)
if H == self.image_size and W == self.image_size:
positions = torch.arange(x.shape[1], device=x.device)
x = x + self.position_embedding(positions)
else:
x = x + self.interpolate_pos_encoding(x, H, W)
return x
class VisionAttention(nn.Module):
"""Standard bidirectional multi-head attention for the ViT encoder."""
def __init__(self, hidden_size: int, num_heads: int):
super().__init__()
self.num_heads = num_heads
self.head_dim = hidden_size // num_heads
self.scale = self.head_dim ** -0.5
self.qkv = nn.Linear(hidden_size, 3 * hidden_size, bias=False)
self.proj = nn.Linear(hidden_size, hidden_size, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
B, S, C = x.shape
qkv = self.qkv(x).reshape(B, S, 3, self.num_heads, self.head_dim)
qkv = qkv.permute(2, 0, 3, 1, 4)
q, k, v = qkv.unbind(0)
x = F.scaled_dot_product_attention(q, k, v, is_causal=False)
x = x.transpose(1, 2).contiguous().reshape(B, S, C)
return self.proj(x)
class VisionLayer(nn.Module):
"""Single ViT encoder layer (pre-norm, GELU MLP)."""
def __init__(self, hidden_size: int, num_heads: int, intermediate_size: int):
super().__init__()
self.norm1 = nn.LayerNorm(hidden_size, eps=1e-6)
self.attn = VisionAttention(hidden_size, num_heads)
self.norm2 = nn.LayerNorm(hidden_size, eps=1e-6)
self.mlp = nn.Sequential(
nn.Linear(hidden_size, intermediate_size),
nn.GELU(),
nn.Linear(intermediate_size, hidden_size),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x + self.attn(self.norm1(x))
x = x + self.mlp(self.norm2(x))
return x
class MorphVisionEncoder(nn.Module):
"""
Vision Transformer encoder.
Input: RGB image [B, 3, 336, 336]
Output: visual tokens [B, 576, 1024]
"""
def __init__(
self,
image_size: int = 336,
patch_size: int = 14,
num_channels: int = 3,
hidden_size: int = 1024,
num_layers: int = 24,
num_heads: int = 16,
intermediate_size: int = 4096,
):
super().__init__()
self.patch_embed = VisionPatchEmbedding(
image_size, patch_size, num_channels, hidden_size
)
self.layers = nn.ModuleList([
VisionLayer(hidden_size, num_heads, intermediate_size)
for _ in range(num_layers)
])
self.norm = nn.LayerNorm(hidden_size, eps=1e-6)
def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:
x = self.patch_embed(pixel_values)
for layer in self.layers:
x = layer(x)
x = self.norm(x)
return x[:, 1:, :] # drop CLS, return patch tokens only
class MorphVisionProjector(nn.Module):
"""Projects vision features into LLM embedding space (see fusion module for V2)."""
def __init__(self, vision_hidden_size: int = 1024, text_hidden_size: int = 3072):
super().__init__()
self.projector = nn.Sequential(
nn.Linear(vision_hidden_size, text_hidden_size, bias=False),
nn.GELU(),
nn.Linear(text_hidden_size, text_hidden_size, bias=False),
)
self.norm = MorphRMSNorm(text_hidden_size)
def forward(self, vision_features: torch.Tensor) -> torch.Tensor:
return self.norm(self.projector(vision_features))
class MorphVisionModule(nn.Module):
"""
Complete vision pipeline: image → ViT encoder → projector → LLM-space tokens.
NEW 1.1: dynamic high-res mode. Large images are tiled into up to
4 sub-images (2×2 grid of 336px crops) plus one global 336px
thumbnail — a multi-tile high-resolution encoding strategy.
Each tile produces 576 tokens; a learnable separator token marks
tile boundaries (see MorphMultimodalModel.image_sep_embed).
"""
def __init__(
self,
image_size: int = 336,
patch_size: int = 14,
vision_hidden: int = 1024,
vision_layers: int = 24,
vision_heads: int = 16,
vision_intermediate: int = 4096,
text_hidden: int = 3072,
max_image_size: int = 1344,
use_dynamic_hires: bool = True,
):
super().__init__()
self.encoder = MorphVisionEncoder(
image_size=image_size, patch_size=patch_size,
hidden_size=vision_hidden, num_layers=vision_layers,
num_heads=vision_heads, intermediate_size=vision_intermediate,
)
self.projector = MorphVisionProjector(vision_hidden, text_hidden)
self.image_size = image_size
self.max_image_size = max_image_size
self.use_dynamic_hires = use_dynamic_hires
self.num_image_tokens = (image_size // patch_size) ** 2 # 576 per tile
def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:
"""
pixel_values: [B, 3, H, W] — normalized RGB (already tiled if hi-res)
returns: [B, 576, text_hidden]
"""
features = self.encoder(pixel_values)
return self.projector(features)
def preprocess_image(
self,
image,
device: torch.device,
hires: Optional[bool] = None,
) -> torch.Tensor:
"""
Preprocess PIL image(s) to tensor(s).
If hires (or self.use_dynamic_hires) and the image is larger than
`image_size`, tiles it into up to 4 crops + 1 global thumbnail.
Returns: [num_tiles, 3, image_size, image_size]
"""
import torchvision.transforms as T
use_hires = self.use_dynamic_hires if hires is None else hires
transform = T.Compose([
T.Resize((self.image_size, self.image_size)),
T.ToTensor(),
T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
if not isinstance(image, list):
image = [image]
all_tiles = []
for img in image:
w, h = img.size
if use_hires and max(w, h) > self.image_size * 1.3:
all_tiles.extend(self._tile_image(img, transform))
else:
all_tiles.append(transform(img))
return torch.stack(all_tiles).to(device)
def _tile_image(self, image, transform) -> List[torch.Tensor]:
"""
Split a large image into a 2×2 grid of crops (each resized to
image_size) plus one global thumbnail of the full image.
Matches the max_image_size budget (4 tiles × 336 = 1344px).
"""
w, h = image.size
tile_grid = min(2, self.max_image_size // self.image_size) # up to 2×2
tiles = []
# Global thumbnail (full image, captures overall context)
tiles.append(transform(image))
# Grid crops (local detail)
tile_w, tile_h = w / tile_grid, h / tile_grid
for r in range(tile_grid):
for c in range(tile_grid):
box = (c * tile_w, r * tile_h, (c + 1) * tile_w, (r + 1) * tile_h)
crop = image.crop(box)
tiles.append(transform(crop))
return tiles