High-resolution benthic surveys and gigapixel reef mosaics contain thousands of spatial patches. Passing all patch embeddings into an autoregressive language model overwhelms its attention context window and drives computational latency beyond practical limits.
The Perceiver Resampler (introduced in Flamingo) addresses this scaling challenge. It uses a fixed set of learnable latent queries and alternating cross-attention and self-attention blocks to resample dense visual feature grids into a compact, fixed-size latent sequence (such as 64 tokens).
This lesson implements a functional PyTorch Perceiver Resampler module from scratch. You will inspect cross-attention resampling, verify latent budget invariance on dense inputs, ingest high-resolution coral reef imagery from the LILA BC Community Fish Detection compilation (Coralscapes survey), and benchmark latent compression tradeoffs.
Learning Objectives¶
Implement a functional PyTorch Perceiver Resampler module with latent queries and multi-head cross-attention.
Resample dense visual feature grids (1024 patches) into a compact budget of 64 latent tokens.
Inspect the mathematical properties of latent cross-attention versus spatial self-attention.
Ingest high-resolution coral reef survey frames from the LILA BC Community Fish Detection dataset.
Execute zero-shot benthic habitat assessment using compact open-weight VLMs.
Benchmark latency and computational cost across varying latent budget allocations.
Part 1 - Setup¶
1.1 Install Dependencies¶
We install transformers, accelerate, pillow, opencv-python, pandas, matplotlib, and qwen-vl-utils for vision processing.
# Install required dependencies
!pip install -q transformers>=4.40.0 accelerate pillow opencv-python-headless huggingface_hub pandas matplotlib qwen-vl-utils
1.2 Check GPU and Hardware Fallback¶
A GPU provides high-throughput tensor operations. The module automatically selects CUDA when available and falls back to CPU execution.
import torch
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32
if torch.cuda.is_available():
gpu_name = torch.cuda.get_device_name(0)
vram_gb = torch.cuda.get_device_properties(0).total_memory / 1e9
print(f"[OK] GPU: {gpu_name} | VRAM: {vram_gb:.1f} GB | Precision: {dtype}")
else:
print("[INFO] No GPU detected. Running on CPU with float32 precision.")
1.3 Hugging Face Authentication Handling¶
Qwen2-VL-2B-Instruct is public and ungated. Authentication is optional.
import os
hf_token = os.environ.get("HF_TOKEN", None)
if hf_token:
try:
from huggingface_hub import login
login(token=hf_token)
print("[OK] Hugging Face token applied.")
except Exception as e:
print(f"[INFO] Login skipped: {e}")
else:
print("[INFO] No HF_TOKEN supplied. Using unauthenticated public access.")
Part 2 - Architecture and Module Implementation¶
2.1 The Perceiver Resampler Concept¶
Let represent patch embeddings from a high-resolution image, where may be 1024, 2048, or more.
The Perceiver Resampler defines learned latent query vectors:
where (typically ). In each Perceiver block:
Cross-Attention: Latents query the concatenated visual features and latents:
Self-Attention: Latents interact with one another:
Feed-Forward Network:
Regardless of input patch count , the resampler produces exactly tokens projected to language embedding dimension .
2.2 PyTorch Module Implementation of Perceiver Resampler¶
We implement the complete Perceiver Resampler architecture using PyTorch modules.
import torch
import torch.nn as nn
class PerceiverBlock(nn.Module):
'''Perceiver block combining visual cross-attention, latent self-attention, and FFN.'''
def __init__(self, latent_dim: int, visual_dim: int, num_heads: int = 8):
super().__init__()
self.norm_latents = nn.LayerNorm(latent_dim)
self.norm_context = nn.LayerNorm(visual_dim)
self.cross_attn = nn.MultiheadAttention(
embed_dim=latent_dim,
kdim=visual_dim,
vdim=visual_dim,
num_heads=num_heads,
batch_first=True,
)
self.norm_sa = nn.LayerNorm(latent_dim)
self.self_attn = nn.MultiheadAttention(
embed_dim=latent_dim,
num_heads=num_heads,
batch_first=True,
)
self.norm_ffn = nn.LayerNorm(latent_dim)
self.ffn = nn.Sequential(
nn.Linear(latent_dim, latent_dim * 4),
nn.GELU(),
nn.Linear(latent_dim * 4, latent_dim),
)
def forward(self, latents: torch.Tensor, visual_features: torch.Tensor):
# Cross-attention: latents attend to visual inputs
l_norm = self.norm_latents(latents)
v_norm = self.norm_context(visual_features)
ca_out, _ = self.cross_attn(l_norm, v_norm, v_norm)
latents = latents + ca_out
# Self-attention among latents
sa_norm = self.norm_sa(latents)
sa_out, _ = self.self_attn(sa_norm, sa_norm, sa_norm)
latents = latents + sa_out
# Feed-forward network
ffn_norm = self.norm_ffn(latents)
latents = latents + self.ffn(ffn_norm)
return latents
class PerceiverResampler(nn.Module):
'''Resamples variable-length visual tokens into a fixed budget of latent tokens.'''
def __init__(
self,
num_latents: int = 64,
latent_dim: int = 1024,
visual_dim: int = 1280,
llm_dim: int = 1536,
depth: int = 3,
num_heads: int = 8,
):
super().__init__()
self.num_latents = num_latents
self.latents = nn.Parameter(torch.randn(1, num_latents, latent_dim) * 0.02)
self.blocks = nn.ModuleList([
PerceiverBlock(latent_dim=latent_dim, visual_dim=visual_dim, num_heads=num_heads)
for _ in range(depth)
])
self.proj = nn.Linear(latent_dim, llm_dim)
self.proj_to_llm = self.proj
def forward(self, visual_features: torch.Tensor):
bs = visual_features.size(0)
latents = self.latents.expand(bs, -1, -1)
for block in self.blocks:
latents = block(latents, visual_features)
output_tokens = self.proj(latents)
return output_tokens
resampler = PerceiverResampler(
num_latents=64,
latent_dim=1024,
visual_dim=1280,
llm_dim=1536,
depth=3,
num_heads=8,
).to(device=device, dtype=dtype)
print("[OK] PerceiverResampler module initialized.")
n_params = sum(p.numel() for p in resampler.parameters())
print(f"Total Perceiver Resampler Parameters: {n_params:,}")
2.3 Verify Latent Token Compression on Dense Visual Inputs¶
We simulate dense benthic imagery producing 1024 patch tokens and verify that the resampler compresses the stream into exactly 64 tokens.
mock_dense_patches = torch.randn(1, 1024, 1280, dtype=dtype, device=device)
with torch.no_grad():
compressed_tokens = resampler(mock_dense_patches)
print(f"Input dense visual tokens: {list(mock_dense_patches.shape)}")
print(f"Compressed latent tokens: {list(compressed_tokens.shape)}")
compression_factor = mock_dense_patches.size(1) / compressed_tokens.size(1)
print(f"Token Compression Factor: {compression_factor:.1f}x reduction")
assert compressed_tokens.shape == (1, 64, 1536), "Unexpected output shape from resampler."
print("[OK] Fixed latent budget verified.")
2.4 Load Multimodal Model for Downstream Marine Analysis¶
We load Qwen2-VL-2B-Instruct to run downstream zero-shot reasoning on the loaded survey imagery.
from transformers import Qwen2VLForConditionalGeneration, AutoProcessor
MODEL_ID = "Qwen/Qwen2-VL-2B-Instruct"
print(f"Loading {MODEL_ID}...")
processor = AutoProcessor.from_pretrained(MODEL_ID)
model = Qwen2VLForConditionalGeneration.from_pretrained(
MODEL_ID,
torch_dtype=dtype,
device_map="auto" if torch.cuda.is_available() else None,
low_cpu_mem_usage=True,
).eval()
if not torch.cuda.is_available():
model = model.to(device)
print("[OK] Downstream VLM ready.")
Part 3 - LILA BC Marine Dataset Integration¶
3.1 Coralscapes Benthic Reef Survey Dataset Attribution¶
We integrate high-resolution coral reef benthic survey imagery from the Coralscapes dataset, part of the Community Fish Detection compilation on LILA BC.
Dataset Name: Coralscapes Benthic Reef Survey (Community Fish Detection Compilation)
Source Catalog: https://
lila .science /datasets /community -fish -detection -dataset/ Attribution: Filippo Varini et al. (Community Fish Detector project), Coralscapes diver-borne reef survey compilation
License: Apache 2.0 / Community Data License Agreement - Permissive (CDLA-Permissive-1.0)
Azure Blob Storage: https://
lilawildlife .blob .core .windows .net /lila -wildlife /community -fish -detection -dataset /JPEGImages/
The survey captures high-resolution (2048x1024) benthic habitat imagery across global coral reef systems using diver-borne imaging systems, documenting benthic substrate complexity, coral structural morphology, and associated marine fauna.
3.2 Ingest Dense Coastal Habitat Images¶
We download dense coral reef survey sample images directly from Azure storage. We set a default slice of 3 sample frames with an adjustable parameter to expand data volume on demand. Network requests are protected with offline fallback handling.
import urllib.request
import io
from PIL import Image, ImageDraw
BASE_AZURE_URL = "https://lilawildlife.blob.core.windows.net/lila-wildlife/community-fish-detection-dataset/JPEGImages/"
# Authentic high-resolution Coralscapes benthic reef survey frames (10 entries for volume expansion)
COASTAL_SAMPLE_URLS = [
("coralscapes_site10_000001_012400_leftImg8bit.png", "https://lilawildlife.blob.core.windows.net/lila-wildlife/community-fish-detection-dataset/JPEGImages/coralscapes_site10_000001_012400_leftImg8bit.png"),
("coralscapes_site10_000002_008400_leftImg8bit.png", "https://lilawildlife.blob.core.windows.net/lila-wildlife/community-fish-detection-dataset/JPEGImages/coralscapes_site10_000002_008400_leftImg8bit.png"),
("coralscapes_site10_000003_017912_leftImg8bit.png", "https://lilawildlife.blob.core.windows.net/lila-wildlife/community-fish-detection-dataset/JPEGImages/coralscapes_site10_000003_017912_leftImg8bit.png"),
("coralscapes_site10_000004_017518_leftImg8bit.png", "https://lilawildlife.blob.core.windows.net/lila-wildlife/community-fish-detection-dataset/JPEGImages/coralscapes_site10_000004_017518_leftImg8bit.png"),
("coralscapes_site10_000005_017343_leftImg8bit.png", "https://lilawildlife.blob.core.windows.net/lila-wildlife/community-fish-detection-dataset/JPEGImages/coralscapes_site10_000005_017343_leftImg8bit.png"),
("coralscapes_site10_000006_005752_leftImg8bit.png", "https://lilawildlife.blob.core.windows.net/lila-wildlife/community-fish-detection-dataset/JPEGImages/coralscapes_site10_000006_005752_leftImg8bit.png"),
("coralscapes_site10_000007_006849_leftImg8bit.png", "https://lilawildlife.blob.core.windows.net/lila-wildlife/community-fish-detection-dataset/JPEGImages/coralscapes_site10_000007_006849_leftImg8bit.png"),
("coralscapes_site10_000008_005094_leftImg8bit.png", "https://lilawildlife.blob.core.windows.net/lila-wildlife/community-fish-detection-dataset/JPEGImages/coralscapes_site10_000008_005094_leftImg8bit.png"),
("coralscapes_site10_000009_009890_leftImg8bit.png", "https://lilawildlife.blob.core.windows.net/lila-wildlife/community-fish-detection-dataset/JPEGImages/coralscapes_site10_000009_009890_leftImg8bit.png"),
("coralscapes_site10_000010_014085_leftImg8bit.png", "https://lilawildlife.blob.core.windows.net/lila-wildlife/community-fish-detection-dataset/JPEGImages/coralscapes_site10_000010_014085_leftImg8bit.png"),
]
# Adjustable data slice parameter (default: 3)
num_samples = 3 # Increase up to len(COASTAL_SAMPLE_URLS) to download additional frames
dense_samples = []
for fname, url in COASTAL_SAMPLE_URLS[:num_samples]:
print(f"Fetching {fname}...")
try:
req = urllib.request.Request(url, headers={"User-Agent": "Mozilla/5.0"})
with urllib.request.urlopen(req, timeout=15) as resp:
data = resp.read()
pil_img = Image.open(io.BytesIO(data)).convert("RGB")
dense_samples.append({
"filename": fname,
"image": pil_img,
"width": pil_img.width,
"height": pil_img.height,
"bytes": len(data),
"is_fallback": False,
})
print(f" [OK] Loaded {fname} ({pil_img.width}x{pil_img.height}, {len(data):,} bytes)")
except Exception as err:
print(f" [WARN] Remote fetch failed for {fname} ({err}). Generating synthetic reef fallback.")
fallback_img = Image.new("RGB", (640, 480), color=(10, 65, 80))
draw = ImageDraw.Draw(fallback_img)
draw.rectangle([30, 30, 610, 450], outline=(40, 160, 190), width=3)
draw.text((50, 50), f"Synthetic Reef Fallback: {fname}", fill=(200, 240, 255))
dense_samples.append({
"filename": fname,
"image": fallback_img,
"width": 640,
"height": 480,
"bytes": 0,
"is_fallback": True,
})
print(f"\nTotal coastal survey frames loaded: {len(dense_samples)}")
3.3 Visualize Survey Images¶
We inspect the high-resolution benthic habitat images to assess substrate rugosity, turf algae, and coral structural cover.
import matplotlib.pyplot as plt
n_samples = max(1, len(dense_samples))
fig, axes = plt.subplots(1, n_samples, figsize=(5 * n_samples, 4))
if n_samples == 1:
axes = [axes]
for ax, sample in zip(axes, dense_samples):
ax.imshow(sample["image"])
title = sample["filename"] if not sample.get("is_fallback") else f"{sample['filename']} (Fallback)"
ax.set_title(title, fontsize=8)
ax.axis("off")
plt.suptitle("Coralscapes Benthic Habitat Survey Frames", fontsize=11)
plt.tight_layout()
plt.show()
Part 4 - Inference and Benthic Habitat Analysis Pipeline¶
4.1 Define Benthic Classification Engine¶
We define an inference engine configured for dense habitat characterization, assessing coral structural complexity and benthic coverage categories.
from qwen_vl_utils import process_vision_info
@torch.no_grad()
def run_benthic_inference(
image: Image.Image,
prompt_text: str,
max_new_tokens: int = 140,
) -> str:
messages = [
{
"role": "user",
"content": [
{"type": "image", "image": image},
{"type": "text", "text": prompt_text},
],
}
]
text = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
image_inputs, video_inputs = process_vision_info(messages)
inputs = processor(
text=[text],
images=image_inputs,
videos=video_inputs,
padding=True,
return_tensors="pt",
).to(device)
generated_ids = model.generate(**inputs, max_new_tokens=max_new_tokens)
trimmed_ids = [out[len(inp):] for inp, out in zip(inputs.input_ids, generated_ids)]
response = processor.batch_decode(
trimmed_ids,
skip_special_tokens=True,
clean_up_tokenization_spaces=False,
)[0]
return response.strip()
print("[OK] Benthic inference engine defined.")
4.2 Substrate Rugosity and Coral Health Analysis¶
We run zero-shot benthic evaluation on the first survey image, prompting the model for substrate classification, coral growth forms, and structural rugosity.
if not dense_samples:
print("No coastal survey frames available for inference (num_samples=0). Increase num_samples to at least 1.")
else:
target_sample = dense_samples[0]
prompt = (
"Analyze this coral reef benthic survey image. Identify visible coral growth forms "
"(branching, massive, encrusting), estimate dominant substrate coverage "
"(hard coral, macroalgae, sand, rubble), and evaluate overall structural rugosity."
)
print(f"Analyzing {target_sample['filename']}...")
response = run_benthic_inference(target_sample["image"], prompt, max_new_tokens=160)
print("\nBenthic Ecological Assessment:")
print("-" * 60)
print(response)
print("-" * 60)Part 5 - Batch Evaluation and Latency Benchmarking¶
5.1 Batch Processing Across Survey Frames¶
We evaluate all loaded survey images, recording latency and extracting ecological observations into a structured pandas DataFrame.
import time
import pandas as pd
benthic_records = []
print(f"Processing {len(dense_samples)} benthic habitat frames...")
for idx, sample in enumerate(dense_samples):
t0 = time.time()
resp = run_benthic_inference(
sample["image"],
"Identify dominant substrate types and coral health conditions visible in this frame.",
max_new_tokens=100,
)
elapsed = time.time() - t0
benthic_records.append({
"frame_id": idx + 1,
"filename": sample["filename"],
"width": sample["width"],
"height": sample["height"],
"latency_sec": round(elapsed, 3),
"is_fallback": sample["is_fallback"],
"assessment": resp,
})
print(f" Frame {idx+1}/{len(dense_samples)} processed in {elapsed:.2f}s")
df_benthic = pd.DataFrame(benthic_records)
csv_file = "benthic_survey_log.csv"
df_benthic.to_csv(csv_file, index=False)
print(f"\n[OK] Survey assessments exported to {csv_file}.")
if not df_benthic.empty:
print(df_benthic[["frame_id", "filename", "latency_sec", "assessment"]])
else:
print("No benthic survey records available to summarize (num_samples=0).")
5.2 Latent Budget Tradeoff Modeling¶
We benchmark computational cost as a function of the latent query count . The resampler cross-attention complexity scales as , while downstream language model self-attention complexity scales as .
# Compare theoretical attention operations for different latent budgets
input_patches = 1024
prompt_tokens = 32
latent_budgets = [16, 32, 64, 128, 256]
comparison_records = []
for M in latent_budgets:
# Cross-attention operations in resampler: M * input_patches
ca_ops = M * input_patches
# LLM self-attention operations per layer: (M + prompt_tokens)^2
llm_ops = (M + prompt_tokens)**2
# Compression ratio
comp_ratio = input_patches / M
comparison_records.append({
"latent_budget_M": M,
"compression_ratio": f"{comp_ratio:.1f}x",
"resampler_ca_ops": ca_ops,
"llm_self_attn_ops": llm_ops,
"relative_llm_cost": round(llm_ops / ((64 + prompt_tokens)**2), 2),
})
df_tradeoffs = pd.DataFrame(comparison_records)
print("Latent Budget Scaling Dynamics (Input = 1024 Patches):")
print(df_tradeoffs.to_string(index=False))
Part 6 - Discussion and Marine Science Applications¶
6.1 Perceiver Resampler vs Q-Former Comparison¶
While both architectures compress variable visual inputs into fixed token counts:
Depth and Cross-Attention: A standard Q-Former uses cross-attention followed by self-attention in every block, querying visual tokens at each layer. The Perceiver Resampler interleaves visual cross-attention with multiple dense latent self-attention layers, allowing deep feature abstraction before language model injection.
Flamingo vs BLIP-2: Flamingo uses Perceiver Resamplers gated with cross-attention layers interleaved into frozen language models. BLIP-2 concatenates Q-Former tokens directly with text prompt tokens in the input embedding layer.
Resolution scaling: Perceiver Resamplers decouple vision encoder output resolution from LLM input sequence length, making them ideal for processing multi-gigabyte benthic orthomosaics.
6.2 Gigapixel Survey Processing¶
Strategies for handling vast underwater survey datasets:
Hierarchical tiling: Split multi-gigapixel orthomosaics into overlapping 1024x1024 tiles, resample each tile to 64 tokens, and concatenate latent representations to form a comprehensive survey context.
Habitat classification standards: Train the resampler on Catlin Seaview Survey and Coralscapes datasets to standardize Coral Reef Coverage Category (CRCC) classifications across international marine protected areas.
Edge survey execution: Resampling to 64 tokens enables complex multimodal models to evaluate benthic biodiversity aboard Autonomous Underwater Vehicles (AUVs) with strict compute and battery constraints.
References¶
Flamingo: Alayrac et al., A Visual Language Model for Few-Shot Learning (NeurIPS 2022). https://
arxiv .org /abs /2206 .14045 Perceiver IO: Jaegle et al., A General Architecture for Structured Inputs & Outputs (ICML 2022). https://
arxiv .org /abs /2107 .14795 Community Fish Detection Dataset (Coralscapes Compilation): https://
lila .science /datasets /community -fish -detection -dataset/