(
self,
x: Tensor,
pos=None,
num_patches=None,
num_special=None,
num_frames=None,
enable_3d_rope=False,
kv_cache=None,
global_idx=0,
num_frame_per_block=1,
num_frame_for_scale=-1,
num_register_tokens=4,
)
| 227 | return self.attn.prepare_qkv(self.norm1(x), pos=pos, enable_3d_rope=enable_3d_rope) |
| 228 | |
| 229 | def forward( |
| 230 | self, |
| 231 | x: Tensor, |
| 232 | pos=None, |
| 233 | num_patches=None, |
| 234 | num_special=None, |
| 235 | num_frames=None, |
| 236 | enable_3d_rope=False, |
| 237 | kv_cache=None, |
| 238 | global_idx=0, |
| 239 | num_frame_per_block=1, |
| 240 | num_frame_for_scale=-1, |
| 241 | num_register_tokens=4, |
| 242 | ) -> Tensor: |
| 243 | # Phase 2 (streaming): single-frame FlashInfer paged attention. |
| 244 | # Handle inline so attn_pre (norm1+prepare_qkv) can be compiled as one CUDA graph. |
| 245 | is_streaming = (kv_cache is not None and (num_frames is None or num_frames <= 1)) |
| 246 | if is_streaming: |
| 247 | manager = kv_cache |
| 248 | # Compiled: norm1 + qkv linear + reshape + q_norm + k_norm + RoPE + format |
| 249 | q_nhd, k_nhd, v_nhd = self.attn_pre(x, pos=pos, enable_3d_rope=enable_3d_rope) |
| 250 | |
| 251 | # Non-keyframe path: attend to cache+current but don't persist the |
| 252 | # current frame. FlashInfer paged attention can only read from the |
| 253 | # paged cache, so we temporarily append (with eviction deferred so |
| 254 | # it stays clean), attend, and then roll back the append. Mirrors |
| 255 | # the ``skip_append`` behavior of the SDPA dict path. |
| 256 | skip_append = getattr(manager, '_skip_append', False) |
| 257 | if skip_append: |
| 258 | prev_defer = manager._defer_eviction |
| 259 | manager._defer_eviction = True |
| 260 | manager.append_frame(global_idx, k_nhd, v_nhd) |
| 261 | attn_x = manager.compute_attention(global_idx, q_nhd) |
| 262 | manager.rollback_last_frame(global_idx) |
| 263 | manager._defer_eviction = prev_defer |
| 264 | else: |
| 265 | # Eager: write frame K/V to paged cache |
| 266 | manager.append_frame(global_idx, k_nhd, v_nhd) |
| 267 | # CPU-only: update eviction state (deque ops, no GPU kernel) |
| 268 | manager.evict_frames( |
| 269 | block_idx=global_idx, |
| 270 | scale_frames=self.attn.kv_cache_scale_frames, |
| 271 | sliding_window=self.attn.kv_cache_sliding_window, |
| 272 | cross_frame_special=self.attn.kv_cache_cross_frame_special, |
| 273 | include_scale_frames=self.attn.kv_cache_include_scale_frames, |
| 274 | camera_only=self.attn.kv_cache_camera_only, |
| 275 | num_register_tokens=num_register_tokens, |
| 276 | ) |
| 277 | # Eager: FlashInfer BatchPrefillWithPagedKVCacheWrapper |
| 278 | attn_x = manager.compute_attention(global_idx, q_nhd) |
| 279 | |
| 280 | # [tpf, H, D] -> [B, tpf, C] (B=1 in streaming, contiguous from FlashInfer output) |
| 281 | attn_x = attn_x.reshape(x.shape[0], q_nhd.shape[0], |
| 282 | self.attn.num_heads * self.attn.head_dim) |
| 283 | # Compiled: output projection |
| 284 | attn_x = self.attn.proj(attn_x) |
| 285 | x = x + self.ls1(attn_x) |
| 286 | else: |
nothing calls this directly
no test coverage detected