InternVL3-9B-CVRR / source_spatial.py
dmis-lab's picture
Add files using upload-large-folder tool
a381a62 verified
Raw
History Blame Contribute Delete
14.5 kB
"""Question-conditioned recurrence over the complete native visual field.
The state is the dynamic-resolution sequence emitted by Qwen2.5-VL's vision
merger, before those image embeddings enter the language model. Tokens keep
their native two-dimensional grid throughout the recurrence. One shared cell
combines directional local messages with visual-to-question cross-attention::
V_0 = VisionMerger(image)
V_{k+1} = V_k + Cell(V_k, question, grid), k = 0, ..., T - 1
Only ``V_T`` is inserted into the ordinary Qwen multimodal prefix. The cell's
output projection is exactly zero at initialization, so every horizon,
including the default T=8, is initially bitwise identical to the base visual
embedding path. The full language model is then run from scratch; no cache
created while constructing the visual field is available to answer decoding.
"""
from __future__ import annotations
from dataclasses import dataclass
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers.models.qwen2.modeling_qwen2 import Qwen2RMSNorm
@dataclass
class SpatialVisualTelemetry:
"""Detached recurrence diagnostics; no visual trajectory is retained."""
update_rms: torch.Tensor # [T]
update_relative: torch.Tensor # [T]
step_drift: torch.Tensor # [T], 1 - cos(V_k, V_{k+1})
final_relative: torch.Tensor # scalar, RMS(V_T - V_0) / RMS(V_0)
class SpatialVisualRecurrentCell(nn.Module):
"""Shared O(N) spatial cell for packed, variable-resolution image tokens.
Global N-by-N visual attention is deliberately absent: a sample may carry
8192 merged visual tokens. Four directional grid messages preserve local
geometry in linear time, while low-rank visual-to-question attention makes
every patch update task dependent. The recurrent *state* remains at the
native language-model width; ``inner_width`` limits only the update
operator's compute and is not a state bottleneck.
"""
def __init__(
self,
width: int,
inner_width: int,
num_heads: int,
*,
rms_norm_eps: float,
residual_scale: float,
) -> None:
super().__init__()
if width < 1 or inner_width < 1:
raise ValueError("width and inner_width must be positive")
if num_heads < 1 or inner_width % num_heads:
raise ValueError(
"inner_width must be divisible by the positive num_heads"
)
if not 0.0 <= residual_scale <= 1.0:
raise ValueError("residual_scale must lie in [0, 1]")
self.width = int(width)
self.inner_width = int(inner_width)
self.num_heads = int(num_heads)
self.head_dim = self.inner_width // self.num_heads
self.residual_scale = float(residual_scale)
self.state_norm = Qwen2RMSNorm(self.width, eps=rms_norm_eps)
self.question_norm = Qwen2RMSNorm(self.width, eps=rms_norm_eps)
# One projection contains distinct north/south/west/east maps. This
# keeps orientation identifiable without allocating a dense 3x3 dxd
# convolution over the 3584-wide native visual state.
self.center_proj = nn.Linear(self.width, self.inner_width, bias=False)
self.neighbor_proj = nn.Linear(
self.width, 4 * self.inner_width, bias=False
)
self.query_proj = nn.Linear(self.width, self.inner_width, bias=False)
self.key_proj = nn.Linear(self.width, self.inner_width, bias=False)
self.value_proj = nn.Linear(self.width, self.inner_width, bias=False)
self.output_proj = nn.Linear(self.inner_width, self.width, bias=False)
def zero_output(self) -> None:
"""Restore the exact base-model identity after generic HF init."""
with torch.no_grad():
self.output_proj.weight.zero_()
@staticmethod
def _token_layout(
merged_grid_thw: torch.LongTensor,
image_to_batch: torch.LongTensor,
batch_size: int,
expected_tokens: int,
) -> tuple[torch.LongTensor, torch.LongTensor, torch.LongTensor]:
"""Return token-to-batch ids, within-sample positions and counts."""
if merged_grid_thw.ndim != 2 or merged_grid_thw.shape[-1] != 3:
raise ValueError("merged_grid_thw must have shape [num_images, 3]")
if image_to_batch.ndim != 1 or image_to_batch.shape[0] != len(
merged_grid_thw
):
raise ValueError("image_to_batch must contain one id per image")
per_image = merged_grid_thw.prod(dim=-1).long()
if int(per_image.sum()) != int(expected_tokens):
raise ValueError(
"merged grids and visual field disagree: "
f"grid tokens={int(per_image.sum())}, state tokens={expected_tokens}"
)
token_to_batch = torch.repeat_interleave(image_to_batch, per_image)
if token_to_batch.numel() and (
int(token_to_batch.min()) < 0
or int(token_to_batch.max()) >= batch_size
):
raise ValueError("image_to_batch contains an out-of-range batch id")
if token_to_batch.numel() > 1 and bool(
(token_to_batch[1:] < token_to_batch[:-1]).any()
):
raise ValueError(
"images must follow the same batch-major order as Qwen inputs"
)
counts = torch.bincount(token_to_batch, minlength=batch_size)
starts = counts.cumsum(0) - counts
within = torch.arange(
expected_tokens, device=token_to_batch.device
) - torch.repeat_interleave(starts, counts)
return token_to_batch, within, counts
@staticmethod
def _neighbor_edges(
merged_grid_thw: torch.LongTensor,
expected_tokens: int,
) -> tuple[torch.LongTensor, torch.LongTensor, torch.LongTensor]:
"""Build packed directed 4-neighbour edges once for all recurrent steps.
Direction ids are 0/1/2/3 = north/south/west/east *as seen by the
destination token*. Temporal slices are separate 2-D fields.
"""
device = merged_grid_thw.device
destinations: list[torch.Tensor] = []
sources: list[torch.Tensor] = []
directions: list[torch.Tensor] = []
offset = 0
for t_value, h_value, w_value in merged_grid_thw.detach().cpu().tolist():
t, h, w = int(t_value), int(h_value), int(w_value)
count = t * h * w
if min(t, h, w) < 1:
raise ValueError(f"invalid merged visual grid {(t, h, w)}")
index = torch.arange(
offset, offset + count, device=device, dtype=torch.long
).view(t, h, w)
def append(dst: torch.Tensor, src: torch.Tensor, direction: int) -> None:
if dst.numel() == 0:
return
destinations.append(dst.reshape(-1))
sources.append(src.reshape(-1))
directions.append(
torch.full(
(dst.numel(),), direction, device=device, dtype=torch.long
)
)
append(index[:, 1:, :], index[:, :-1, :], 0) # north
append(index[:, :-1, :], index[:, 1:, :], 1) # south
append(index[:, :, 1:], index[:, :, :-1], 2) # west
append(index[:, :, :-1], index[:, :, 1:], 3) # east
offset += count
if offset != expected_tokens:
raise ValueError(
f"edge grids contain {offset} tokens, expected {expected_tokens}"
)
if not destinations:
empty = torch.empty(0, device=device, dtype=torch.long)
return empty, empty, empty
return (
torch.cat(destinations),
torch.cat(sources),
torch.cat(directions),
)
def _question_attention(
self,
normalized_state: torch.Tensor,
question_keys: torch.Tensor,
question_values: torch.Tensor,
question_valid: torch.BoolTensor,
token_to_batch: torch.LongTensor,
within_sample: torch.LongTensor,
token_counts: torch.LongTensor,
) -> torch.Tensor:
"""Visual-query/text-memory SDPA for a packed visual sequence."""
batch_size = question_keys.shape[0]
max_tokens = int(token_counts.max()) if token_counts.numel() else 0
if max_tokens == 0:
return normalized_state.new_empty(0, self.inner_width)
query_flat = self.query_proj(normalized_state)
query = query_flat.new_zeros(
batch_size, max_tokens, self.inner_width
)
query[token_to_batch, within_sample] = query_flat
bsz, question_length, _ = question_keys.shape
query = query.view(
bsz, max_tokens, self.num_heads, self.head_dim
).transpose(1, 2)
key = question_keys.view(
bsz, question_length, self.num_heads, self.head_dim
).transpose(1, 2)
value = question_values.view(
bsz, question_length, self.num_heads, self.head_dim
).transpose(1, 2)
# Boolean SDPA masks use True for entries that are allowed to attend.
allowed = question_valid[:, None, None, :]
attended = F.scaled_dot_product_attention(
query,
key,
value,
attn_mask=allowed,
dropout_p=0.0,
is_causal=False,
)
attended = attended.transpose(1, 2).reshape(
bsz, max_tokens, self.inner_width
)
return attended[token_to_batch, within_sample]
def forward(
self,
visual_state: torch.Tensor,
merged_grid_thw: torch.LongTensor,
image_to_batch: torch.LongTensor,
question_embeddings: torch.Tensor,
question_attention_mask: torch.Tensor | None,
*,
steps: int,
) -> tuple[torch.Tensor, SpatialVisualTelemetry]:
if visual_state.ndim != 2 or visual_state.shape[-1] != self.width:
raise ValueError(
f"visual_state must be [N,{self.width}], got "
f"{tuple(visual_state.shape)}"
)
if question_embeddings.ndim != 3 or question_embeddings.shape[-1] != self.width:
raise ValueError(
f"question_embeddings must be [B,Q,{self.width}]"
)
if steps < 0:
raise ValueError("steps must be non-negative")
batch_size = question_embeddings.shape[0]
token_to_batch, within_sample, token_counts = self._token_layout(
merged_grid_thw,
image_to_batch,
batch_size,
visual_state.shape[0],
)
edge_dst, edge_src, edge_direction = self._neighbor_edges(
merged_grid_thw, visual_state.shape[0]
)
if question_attention_mask is None:
question_valid = torch.ones(
question_embeddings.shape[:2],
dtype=torch.bool,
device=question_embeddings.device,
)
else:
if question_attention_mask.shape != question_embeddings.shape[:2]:
raise ValueError("question attention mask shape mismatch")
question_valid = question_attention_mask > 0
if not bool(question_valid.any(dim=-1).all()):
raise ValueError("every image-bearing sample needs a question token")
normalized_question = self.question_norm(question_embeddings)
question_keys = self.key_proj(normalized_question)
question_values = self.value_proj(normalized_question)
initial = visual_state
state = visual_state
update_rms: list[torch.Tensor] = []
update_relative: list[torch.Tensor] = []
step_drift: list[torch.Tensor] = []
for _ in range(steps):
normalized = self.state_norm(state)
center = self.center_proj(normalized)
directional = self.neighbor_proj(normalized).view(
state.shape[0], 4, self.inner_width
)
spatial = center.new_zeros(center.shape)
degree = center.new_zeros(center.shape[0], 1)
if edge_dst.numel():
messages = directional[edge_src, edge_direction]
spatial.index_add_(0, edge_dst, messages)
degree.index_add_(
0,
edge_dst,
torch.ones(
edge_dst.shape[0], 1, dtype=center.dtype, device=center.device
),
)
spatial = spatial / degree.clamp_min(1.0)
question = self._question_attention(
normalized,
question_keys,
question_values,
question_valid,
token_to_batch,
within_sample,
token_counts,
)
mixed = F.silu((center + spatial + question) / (3.0**0.5))
update = self.output_proj(mixed) * self.residual_scale
next_state = state + update
with torch.no_grad():
state_rms = state.float().square().mean().sqrt().clamp_min(1e-8)
update_norm = update.float().square().mean().sqrt()
cosine = F.cosine_similarity(
state.float(), next_state.float(), dim=-1
).mean()
update_rms.append(update_norm.detach())
update_relative.append((update_norm / state_rms).detach())
step_drift.append((1.0 - cosine).detach())
state = next_state
with torch.no_grad():
initial_rms = initial.float().square().mean().sqrt().clamp_min(1e-8)
final_relative = (
(state.float() - initial.float()).square().mean().sqrt()
/ initial_rms
).detach()
empty = initial.new_empty(0, dtype=torch.float32)
telemetry = SpatialVisualTelemetry(
update_rms=(torch.stack(update_rms) if update_rms else empty),
update_relative=(
torch.stack(update_relative) if update_relative else empty
),
step_drift=(torch.stack(step_drift) if step_drift else empty),
final_relative=final_relative,
)
return state, telemetry
__all__ = ["SpatialVisualRecurrentCell", "SpatialVisualTelemetry"]