File size: 14,497 Bytes
a381a62
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
"""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"]