rnagabh/gemma4-vision-encoder
6179
Gemma 4 Vision Encoder (27-layer ViT with 2D RoPE)
Standalone extraction of the vision encoder from Google's Gemma 4 31B multimodal model. This is a 569.6M parameter Vision Transformer with learned 2D positional embeddings, RoPE, QK-norms, and gated MLP — a significant upgrade from the SigLIP encoder used in Gemma 3.
License: Apache 2.0 (inherited from Gemma 4 — no restrictions)
Architecture
What's New vs Gemma 3 (SigLIP)
Not Shared with E2B/E4B
Unlike the audio encoder (which is identical across E2B and E4B), the vision encoders differ:
Usage
import torch
from transformers import Gemma4VisionModel, Gemma4ImageProcessor
from PIL import Image
# Load vision encoder directly from this repo
vision_model = Gemma4VisionModel.from_pretrained(
"rnagabh/gemma4-vision-encoder",
torch_dtype=torch.bfloat16,
)
vision_model.to("cuda")
vision_model.eval()
# Load image processor (saved in this repo)
image_processor = Gemma4ImageProcessor.from_pretrained("rnagabh/gemma4-vision-encoder")
# Process an image
img = Image.open("your_image.jpg")
processed = image_processor(images=[img], return_tensors="pt")
pixel_values = processed["pixel_values"].to(dtype=torch.bfloat16, device="cuda")
position_ids = processed["image_position_ids"].to(device="cuda")
tokens_per_image = processed["num_soft_tokens_per_image"] # for splitting batch output
with torch.no_grad():
output = vision_model(pixel_values=pixel_values, pixel_position_ids=position_ids)
embeddings = output.last_hidden_state # (num_tokens, 1152)
# Mean-pool for a single image vector
image_embedding = embeddings.float().mean(dim=0) # (1152,)Important: Always use the Gemma4ImageProcessor included in this repo for preprocessing. It handles resizing, patchification, position ID generation, and pixel normalization. Manual patchification without this processor will produce significantly degraded results.Benchmark Results (frozen 1152-dim embeddings, linear probe)
CIFAR-10 Classification
Strong performance across all classes: airplane (0.98 F1), ship (0.98 F1), truck (0.97 F1), automobile (0.97 F1). Weakest class is cat (0.86 F1) — a fine-grained category that is inherently harder.
Files in This Repo
Limitations
- End-to-end trained for LLM decoding: The encoder was trained to produce features for Gemma 4's text decoder. The 1152-dim output is the pure vision representation; the
embed_visionprojection maps to the 31B's text hidden space (5376-dim). - Requires image processor: Use the
Gemma4ImageProcessorincluded in this repo for preprocessing. The model expects pre-patchified(B, num_patches, 768)tensors with explicit 2D position IDs — the processor handles this automatically. - Variable aspect ratio support: The 2D position embeddings enable non-square images. The processor generates correct position IDs for any aspect ratio.
- Output shape note: The pooler strips padding and collapses the batch dimension, returning
(num_valid_tokens, 1152). For batched inference, usenum_soft_tokens_per_imagefrom the processor to split the output back into per-image embeddings.
Extraction Details
- Extracted from
google/gemma-4-31B-itby downloading only the shard containing vision tower weights (model-00001-of-00002.safetensors) - No full model load required — targeted tensor extraction
- Weights loaded with
strict=True— perfect match - Forward pass verified: 864×864 image → (324, 1152) output
- All architecture specs verified against the live model config
