Compile hot, fixed-shape modules with mode="reduce-overhead". Mirrors the targets in gct_profile.py:compile_model. Unlike the profile script, `model.point_head` is **kept** — the demo needs world_points for visualization.
(model)
| 168 | # ============================================================================= |
| 169 | |
| 170 | def compile_model(model): |
| 171 | """Compile hot, fixed-shape modules with mode="reduce-overhead". |
| 172 | |
| 173 | Mirrors the targets in gct_profile.py:compile_model. Unlike the profile script, |
| 174 | `model.point_head` is **kept** — the demo needs world_points for visualization. |
| 175 | """ |
| 176 | agg = model.aggregator |
| 177 | for i, b in enumerate(agg.frame_blocks): |
| 178 | agg.frame_blocks[i] = torch.compile(b, mode="reduce-overhead") |
| 179 | for i, b in enumerate(agg.patch_embed.blocks): |
| 180 | agg.patch_embed.blocks[i] = torch.compile(b, mode="reduce-overhead") |
| 181 | for b in agg.global_blocks: |
| 182 | if hasattr(b, 'attn_pre'): |
| 183 | b.attn_pre = torch.compile(b.attn_pre, mode="reduce-overhead") |
| 184 | if hasattr(b, 'ffn_residual'): |
| 185 | b.ffn_residual = torch.compile(b.ffn_residual, mode="reduce-overhead") |
| 186 | b.attn.proj = torch.compile(b.attn.proj, mode="reduce-overhead") |
| 187 | |
| 188 | |
| 189 | def _warm_streaming(model, images, scale_frames, warm_stream_n, dtype, |