File size: 37,918 Bytes
c335050
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
# Retrieval-Uncertainty Loss & Memory-Retrieval Improvements

This document gives (1) a concrete, codebase-accurate implementation of the
**retrieval-uncertainty loss** you described, and (2) a set of additional
methods to improve the model's ability to retrieve preceding (context) frames.

> **Loss you asked for**
> `L_unc = e^(-um) · sg(MSE(VAE(Target View), predicted_x0)) + um`
> where `um` is the model's predicted (log-)uncertainty about the current
> prediction and `sg(·)` is stop-gradient.
>
> This is the **Kendall–Gal heteroscedastic-uncertainty** objective. With the
> MSE stop-gradiented, the uncertainty head learns a *per-token retrieval
> confidence* (low `um` where the model reconstructs the target well, high `um`
> where it fails / "forgets"). That confidence map is then reused to **reweight
> the main denoising loss**, focusing capacity on the tokens the memory pathway
> currently fails to retrieve.

Read `CLAUDE.md` and `attention.md` first for the two-chunk paradigm and the
public-repo constraints.

---

## 0. Where the loss lives in this codebase

The training loss is **flow-matching MSE**, computed in
`diffsynth/pipelines/wan_video_new.py → WanVideoPipeline.training_loss(...)`.
The relevant facts (verified against the code):

- The scheduler is `FlowMatchScheduler` (`diffsynth/schedulers/flow_match.py`):
  - `add_noise`: `x_t = (1 - σ)·x0 + σ·noise`
  - `training_target`: `v = noise − x0`**the model predicts velocity**, not x0
- In `training_mode == "context"` (the default for memory baselines), the
  target tokens are `input_latents` (already VAE-encoded — this **is**
  `VAE(Target View)` in latent space), and the loss is:

  ```python
  noisy_target_latents = self.scheduler.add_noise(target_latents, target_noise, timestep)   # x_t
  training_target      = self.scheduler.training_target(target_latents, target_noise, timestep)  # v = noise - x0
  ...
  noise_pred           = self.model_fn(**inputs, timestep=timestep)   # v_pred over [target | context] tokens
  target_noise_pred    = noise_pred[:, :, :target_latents.shape[2], :, :]   # suffix layout: target first
  loss = F.mse_loss(target_noise_pred.float(), training_target.float())
  loss = loss * self.scheduler.training_weight(timestep)
  ```

**Recovering `predicted_x0`** (needed for your loss) is exact under flow
matching — no extra forward pass:

```
x0_pred = x_t − σ · v_pred = noisy_target_latents − σ · target_noise_pred
```

So `MSE(VAE(Target View), predicted_x0) = MSE(target_latents, x0_pred)`, computed
entirely in latent space — no VAE decode is needed during training.

> `σ` for the sampled `timestep` is `self.scheduler.sigmas[timestep_id]` where
> `timestep_id = argmin(|scheduler.timesteps − timestep|)`. There are two
> equivalent target-token layouts (`context_position == "suffix"` vs `prefix`);
> the slice that selects target tokens is already computed as
> `target_noise_pred` — reuse it.

---

## 1. The module (already added)

`diffsynth/models/memory/uncertainty.py` (exported from
`diffsynth/models/memory/__init__.py`). Key pieces:

- `UncertaintyHead(in_channels)` — a 1×1×1 Conv3d MLP mapping the per-token x0
  prediction `(B, C, T, H, W)` → log-uncertainty `um` `(B, 1, T, H, W)`. The
  final conv is **zero-initialised**, so `um ≡ 0` at init ⇒ `exp(−um) ≡ 1`
  (a no-op weighting) ⇒ the existing training dynamics are unchanged on step 0.
- `recover_x0_flow_match(noisy_target, v_pred, sigma)` — `x_t − σ·v`.
- `per_token_mse_map(x0_pred, x0_target)` — channel-mean MSE → `(B,1,T,H,W)`.
- `retrieval_uncertainty_loss(um, mse_map, detach_mse=True)` — returns
  `mean(exp(−um)·sg(MSE) + um)`.

Smoke-tested: `um` is `(B,1,T,H,W)`, `==0` at init, loss is finite, gradients
flow into the head.

---

## 2. Wiring it into training

Five small edits. All are additive and gated by a new flag so default behaviour
is untouched.

### 2a. `diffsynth/pipelines/wan_video_new.py` — `WanVideoPipeline.__init__`

Create the head lazily (the latent channel count `C` is known from the DiT
`in_dim`, typically 16 for Wan 2.1). Add near where other memory attributes are
set:

```python
self.use_retrieval_uncertainty = False
self.retrieval_uncertainty_weight = 1.0
self.uncertainty_head = None   # nn.Module, built on first use
```

### 2b. `WanVideoPipeline.training_loss` — compute the extra loss

In **both** the `"context"` and `"predict"` branches, right after the existing
`loss = ... * training_weight(timestep)`, insert:

```python
if getattr(self, "use_retrieval_uncertainty", False):
    from diffsynth.models.memory.uncertainty import (
        UncertaintyHead, recover_x0_flow_match, per_token_mse_map,
        retrieval_uncertainty_loss,
    )
    # σ for this timestep (flow-match scheduler).
    sched = self.scheduler
    timestep_id = torch.argmin(
        (sched.timesteps - timestep.to(sched.timesteps.device)).abs())
    sigma = sched.sigmas[timestep_id].to(target_noise_pred.dtype)

    # x_t and v_pred for the TARGET tokens only (reuse the existing slice).
    x_t_target = noisy_target_latents                       # (B,C,T,H,W)
    x0_pred    = recover_x0_flow_match(x_t_target, target_noise_pred, sigma)
    x0_target  = target_latents                             # = VAE(Target View)

    # Lazily build the head with the right channel count, on the right device.
    if self.uncertainty_head is None:
        self.uncertainty_head = UncertaintyHead(x0_pred.shape[1]).to(
            device=x0_pred.device, dtype=torch.float32)
    um = self.uncertainty_head(x0_pred)                     # (B,1,T,H,W)

    mse_map = per_token_mse_map(x0_pred, x0_target)
    loss_unc = retrieval_uncertainty_loss(um, mse_map, detach_mse=True)

    # (Optional, Method A) reweight the MAIN denoising loss by confidence so the
    # denoiser/memory pathway focuses on hard-to-retrieve tokens. Detach um here
    # so this term trains the denoiser, not the head.
    # per_token_main = (target_noise_pred.float() - training_target.float()).pow(2).mean(1, keepdim=True)
    # w = torch.exp(-um.detach()).clamp(0.1, 10.0)
    # loss = (w * per_token_main).mean() * self.scheduler.training_weight(timestep)

    loss = loss + self.retrieval_uncertainty_weight * loss_unc
```

> `target_latents`, `noisy_target_latents`, `target_noise_pred`, and `timestep`
> all already exist as locals in that scope — no signature changes needed.

### 2c. `src/model_training/train.py` — argparse flags

Add to the numeric-default tuples (near `--timestep_shift`):

```python
("--retrieval_uncertainty_weight", dict(type=float, default=1.0)),
```

and add `"--use_retrieval_uncertainty"` to the list of store-true flags
(alongside `"--use_block_wise_ssm"`, ~line 1510).

### 2d. `src/model_training/train.py` — push flags onto the pipe

In the trainer `__init__` (where `self.pipe.use_spatial_memory = ...` is set,
~line 984), add:

```python
self.pipe.use_retrieval_uncertainty = bool(use_retrieval_uncertainty)
self.pipe.retrieval_uncertainty_weight = float(retrieval_uncertainty_weight)
```

and thread the two values in from `_arg(...)` at the trainer construction call
(~line 1671), mirroring `timestep_shift`.

### 2e. Make the head trainable **and saved**

The optimizer collects `model.trainable_modules()` = all params with
`requires_grad=True`, and `--save_full_model` exports the whole DiT state
(otherwise only `requires_grad` params are exported via
`export_trainable_state_dict`). The head lives on `self.pipe`, **not** inside
`dit`, so:

- Its params are created with `requires_grad=True` by default ✔ (so AdamW will
  pick them up **provided the optimizer is built after the head exists**). The
  head is built lazily on the first `training_loss` call, which is *after*
  `torch.optim.AdamW(model.trainable_modules(), ...)` (~line 1794). **Fix:**
  build the head eagerly so it is registered before the optimizer is created —
  add this right after the block-replacement section in `train.py`:

  ```python
  if _arg('use_retrieval_uncertainty', False):
      from diffsynth.models.memory.uncertainty import UncertaintyHead
      _c = int(getattr(model.pipe.dit, "in_dim", 16))
      model.pipe.uncertainty_head = UncertaintyHead(_c).to(
          device=next(model.pipe.dit.parameters()).device, dtype=torch.float32)
  ```

- For checkpointing, the head is an attribute of `self.pipe`, which is a
  submodule of the training `model`, so `accelerator.get_state_dict(model)`
  includes `pipe.uncertainty_head.*` keys. With `--save_full_model` they are
  saved; the keys are ignored at inference (the head is not needed to generate).
  If you do **not** use `--save_full_model`, ensure the head params have
  `requires_grad=True` (they do) so `export_trainable_state_dict` keeps them.

### 2f. Launcher

Copy an existing memory launcher (e.g.
`train/memory_baselines_basic/run_spatial_memory_baseline.sh`) and add:

```bash
  --use_retrieval_uncertainty --retrieval_uncertainty_weight 1.0 \
```

Keep every other hyperparameter identical to the baseline row you are comparing
against — this is a controlled ablation; only the loss should change.

### 2g. Sanity check before a full run

```bash
PYTHONPATH=. python3 tests/test_two_chunk_anchor_readout.py
# plus a tiny head test (identity-at-init + grad flow), e.g.:
PYTHONPATH=. python3 - <<'PY'
import importlib.util, torch
s=importlib.util.spec_from_file_location('u','diffsynth/models/memory/uncertainty.py')
u=importlib.util.module_from_spec(s); s.loader.exec_module(u)
h=u.UncertaintyHead(16); xt=torch.randn(1,16,21,44,80); v=torch.randn_like(xt)
x0=u.recover_x0_flow_match(xt,v,torch.tensor(.7)); um=h(x0)
assert float(um.abs().max())==0.0          # identity at init
L=u.retrieval_uncertainty_loss(um,u.per_token_mse_map(x0,torch.randn_like(x0)))
L.backward(); assert h.net[0].weight.grad is not None
print("ok", float(L))
PY
```

---

## 3. Why this helps retrieval (and how to read the signal)

`um` becomes a learned, per-token map of **where the model fails to reconstruct
the target from memory**. Two ways to exploit it:

- **Diagnostic** — log `um` heatmaps to W&B alongside the two-chunk
  left/right-rotation monitor (`--sampling_atomic_left_right`). High-`um`
  regions on the revisit tail localise *what* the memory is dropping (object
  identity vs background vs camera geometry).
- **Loss reweighting (Method A above)**`exp(−um.detach())` upweights the
  main denoising loss on tokens the model is *confident and wrong* about,
  pushing the memory pathway to fix systematic retrieval failures rather than
  averaging error uniformly.

---

## 4. Additional methods to improve preceding-frame retrieval

Ordered roughly by expected impact / effort. All are compatible with the
two-chunk setup and the existing memory families.

### Method A — Confidence-reweighted denoising loss
Already sketched in §2b. Uses `exp(−um.detach())` to focus the **denoiser** on
hard-to-retrieve tokens. Cheap, synergises directly with the uncertainty head.

### Method B — Explicit retrieval-consistency (anchor) loss
Add a term that directly penalises drift between the **first/anchor frame** and
the **revisit tail** in latent space, since revisit consistency is exactly what
the paper measures. After recovering `x0_pred` for the target tokens:

```
L_anchor = MSE( x0_pred[revisit_tail_tokens], context_latents[anchor_token] )
```
restricted to samples where the trajectory returns near the start pose (the
codebase already constructs loop-closure probes; reuse
`env/loop_utils.py` / the `replay` context source to identify revisit tokens).
This trains the memory pathway to *reproduce* stored content, not merely to
denoise plausibly. Gate behind `--use_anchor_consistency_loss`.

### Method C — Contrastive memory read-out (InfoNCE)
Make the memory read-out **discriminative**: the target token's retrieved
memory feature should match its *own* context frame more than other frames'.
Take per-frame pooled features from the context tokens (before they enter the
DiT blocks) and the corresponding target query features, and add an InfoNCE
loss pulling matched (target-frame ↔ source-frame) pairs together and pushing
mismatched pairs apart. This sharpens *which* preceding frame is retrieved —
particularly useful for the Spatial and Context-K families. Implement as a
small head reading the block hidden state (same hook point as block-wise SSM in
`DiTBlock_w_Action`, see `attention.md`).

### Method D — Harder/longer context sampling (curriculum)
Retrieval is only as good as the supervision distribution. Levers already in
the data path:
- Increase `--context_memory_frames` (K) and/or widen the temporal gap between
  context and target so the model must retrieve *distant* history, not adjacent
  frames (`--context_source replay`, `--prev_chunk_frames`).
- Curriculum: start with short gaps, anneal to longer gaps over training.
- Mix revisit-style samples (leave-and-return) more heavily — the two-chunk
  `--sampling_atomic_left_right` probe shows what to oversample.
Pure data/schedule change; no model edits.

### Method E — Memory dropout / robustness regularisation
Randomly drop or noise a subset of context tokens during training
(`--context_drop_prob`, `--context_noise_std` already exist). Forcing the model
to retrieve from partial memory improves robustness and prevents trivial
copy-through, which tends to help long-horizon revisit. Tune these existing
flags rather than adding code.

### Method F — Cross-attention readout supervision for Spatial memory
For the spatial family (`spatial_cross_attn_readout`), add an auxiliary loss
that encourages the read-out attention map to concentrate on the spatially
corresponding stored region (when camera RT gives a known correspondence).
This is a targeted version of Method C for the spatial grid memory.

### Recommended first experiment
1. Implement §1–§2 (uncertainty head + loss), train one row vs its baseline.
2. Turn on **Method A** (confidence reweighting) — likely the largest gain per
   line of code.
3. Add **Method B** (anchor consistency) if revisit MSE is still the bottleneck.
Evaluate all with the existing tiers:
```bash
export CKPT=outputs/<your_row>/epoch-0.safetensors
bash eval/v2/run_basic_replay_gt.sh
bash eval/v2/run_static_consistency_loop_and_revisit.sh
PHASE=stage1 OOD_DIR=assets/opendomain_revisit bash eval/v2/revisit_suite/run_one_click_revisit_eval.sh
```
Compare revisit-tail MSE / PSNR / LPIPS against the unmodified baseline row.

---

## 5. Pitfalls

- **Predicting x0 vs velocity.** The model outputs **velocity** `v = noise − x0`.
  Do **not** feed `target_noise_pred` directly as `x0` — always recover via
  `x0 = x_t − σ·v` (`recover_x0_flow_match`). Getting this wrong silently
  inverts the uncertainty signal.
- **Non-zero `um` at init.** Keep the head's final layer zero-initialised; a
  non-zero init multiplies the main loss by an arbitrary factor on step 0 and
  destabilises early training.
- **Optimizer misses the head.** Build the head **before**
  `AdamW(model.trainable_modules())` is constructed (see §2e), or its params
  won't be optimised.
- **Stop-gradient.** With `detach_mse=True` the uncertainty term trains only the
  head. If you want it to also shape the denoiser, use Method A's
  `exp(−um.detach())` reweighting of the main loss — don't simply drop the
  stop-gradient on the MSE (that lets the model lower the loss by inflating
  `um`, i.e. "predict badly on purpose").
- **dtype/autocast.** Compute the head and the loss in fp32 (the module already
  casts), and the `um` clamp keeps `exp(−um)` finite under bf16 autocast.
- **Public-repo constraints** (`CLAUDE.md`): no machine-local paths, minimal
  diffs, don't commit `outputs/` or weights.

---

## 6. Where the retrieval target comes from (what "correct retrieval" means)

A confidence map is only meaningful relative to a **target** that defines correct
retrieval. There is no single target — there is a hierarchy of increasingly
strict definitions, and which one you pick decides what your confidence map
actually measures. This codebase already computes the geometric ones.

### Background: what is RT?

**RT = Rotation + Translation = the camera extrinsics** (the rigid-body pose of
the camera). In this repo an RT is a **12-dim row-major vector**
`[t_x, t_y, t_z, R_11, R_12, R_13, R_21, R_22, R_23, R_31, R_32, R_33]`
— a 3×1 translation `t` followed by a flattened 3×3 rotation `R`
(`src/model_training/rt_utils.py` docstring). The `MLP_CamPose(pose_dim=12)`
inside `DiTBlock_w_Action` consumes exactly this 12-vector per latent frame.

The **relative RT** between a context frame *i* and the reference (target) frame
maps points from one camera frame into the other — this is what lets you
*reproject* context content into the target view. It is computed by
`rt_utils.convert_rt_to_relative(rt_list_all, ref_rt)`:

```
R_rel = R_ref⁻¹ · R_i ,   t_rel = R_ref⁻¹ · t_i + (−R_ref⁻¹ · t_ref)
```

(`R_ref⁻¹ = R_refᵀ` since rotations are orthonormal). Camera poses come from the
per-frame JSONs via `pose_to_rt(pose)` (paper default: XY translation + Z-axis
yaw only). Enabled in training by `--use_rt_relative` (env `USE_RT_RELATIVE`).

### Level 0 — Reconstruction target (what the base loss already uses)

Weakest definition: *correct retrieval = the token was denoised well.* Target is
`x0_target = target_latents = VAE(Target View)`; supervision is the per-token MSE
already computed in `training_loss`. **Limitation:** it conflates two failure
modes — (a) the memory pathway *failed to retrieve* the right context, vs.
(b) the content is *genuinely novel / newly revealed* and no memory could help.
For a *retrieval* confidence map you want to isolate (a), so Level 0 alone is the
wrong target.

### Level 1 — Geometric co-visibility (already in the repo)

The retrieval target here is **not learned** — it is a precomputed geometric
label answering *"which past frames actually share field-of-view with the
current frame?"*:

- **`overlap_labels/{video_name}/{frame_idx}.json`** →
  `{"overlapping_frames": ["2796", "2797", ...]}`. For each frame, the list of
  historical frames that co-observe the same scene. **This is the ground-truth
  retrieval target at frame granularity.** Loaded by
  `fov_retrieval.load_overlap_frames()` and consumed by
  `fov_training_integration.retrieve_fov_context_frames()` to *select* the
  context frames during training.
- **`fov_retrieval.compute_fov_overlap_3d(pose1, pose2, fov=52.67°)`** computes a
  continuous overlap score in [0,1] from the 6-DoF poses (mutual visibility +
  forward-direction similarity). This is the function that *generates* the
  labels.

Use it as a **target for the confidence head**: a token whose co-visible content
the model reproduced ⇒ confidence high; a token that *had* co-visible support
but was reproduced wrong ⇒ retrieval failure (what you want to flag). The
co-visibility set also gives a **mask**: only ask "did you retrieve correctly?"
where retrieval was geometrically possible.

### Level 2 — Reprojection correspondence (the per-token target you asked for)

Level 1 is frame-level + coarse-region. Level 2 tightens it to **per latent
token**: use the relative RT to warp the co-visible context latent into the
current view, then define correct retrieval token-by-token. This removes the
Level-0 ambiguity (novel regions are masked out of the retrieval loss).

#### 6.1 Token ↔ pixel geometry in this stack (must get this right)

To reproject *into latent-token space* you need the compression factors:

| Stage | Factor | Source |
| --- | --- | --- |
| VAE spatial downsample | **÷8** (three 2× `downsample2d/3d` blocks) | `wan_video_vae.py` Resample blocks |
| VAE temporal downsample | **÷4** (`temperal_downsample=[True,True,False]`, +1 for the first frame) | `wan_video_vae.py:284` |
| DiT patchify | **(1, 2, 2)** | `wan_video_dit.py:511` `patch_size=(1,2,2)` |

So one latent **token** covers a `8·2 = 16` px × `16` px region of the original
frame (spatially), and the latent grid for a `352×640` frame is
`H_lat = 352/8 = 44`, `W_lat = 640/8 = 80`, then patchified by 2 →
`22 × 40` token grid per latent frame. **Reproject at the latent-pixel grid
(44×80), then patch-pool to the token grid (22×40)** to match `um`'s resolution.

#### 6.2 Building the per-token reprojection-confidence map

Given a context frame *i* and the target frame, with relative pose `(R_rel,
t_rel)` from `convert_rt_to_relative`, the homltography/flow that maps target
latent-pixel `(u,v)` ↔ context latent-pixel depends on scene depth. Two regimes:

- **Depth available** (SpatialVID has more geometry than the static pool):
  full reprojection `p_ctx = K · (R_rel · (depth · K⁻¹ · p_tgt) + t_rel)`.
- **No depth / planar approximation** (static pool, paper's XY+yaw setting): a
  **homography** `Hℓ` suffices because motion is dominated by yaw + translation
  on a plane — exactly the regime `pose_to_rt(constrain_to_xy=True)` encodes.
  Build `Hℓ` from `(R_rel, t_rel)` and a reference plane normal/depth.

The confidence target is then: *warp the context latent into the target view and
measure how close the model's `x0_pred` is to that warped evidence, only where
co-visibility holds.*

```python
# diffsynth/models/memory/reproj_confidence.py  (sketch — add as a new module)
import torch
import torch.nn.functional as F


def latent_grid_hw(height_px: int, width_px: int):
    """Latent-pixel grid before DiT patchify: VAE divides spatial by 8."""
    return height_px // 8, width_px // 8


def warp_context_latent(ctx_latent, H_rel):
    """
    Warp a context latent frame into the target view via a 3x3 homography H_rel
    expressed in *latent-pixel* coordinates (44x80 for 352x640).

    ctx_latent: (B, C, Hl, Wl)   single context latent frame
    H_rel:      (B, 3, 3)        target-latent-pixel -> context-latent-pixel
    returns:    (B, C, Hl, Wl) warped context, (B,1,Hl,Wl) valid mask
    """
    B, C, Hl, Wl = ctx_latent.shape
    ys, xs = torch.meshgrid(
        torch.arange(Hl, device=ctx_latent.device, dtype=torch.float32),
        torch.arange(Wl, device=ctx_latent.device, dtype=torch.float32),
        indexing="ij",
    )
    ones = torch.ones_like(xs)
    grid = torch.stack([xs, ys, ones], dim=-1).reshape(1, Hl * Wl, 3).expand(B, -1, -1)

    src = torch.bmm(grid, H_rel.transpose(1, 2))            # (B, Hl*Wl, 3)
    src = src[..., :2] / src[..., 2:3].clamp(min=1e-6)      # homogeneous divide
    sx, sy = src[..., 0], src[..., 1]

    # Normalise to grid_sample's [-1, 1] coordinates.
    gx = (sx / (Wl - 1)) * 2 - 1
    gy = (sy / (Hl - 1)) * 2 - 1
    samp = torch.stack([gx, gy], dim=-1).reshape(B, Hl, Wl, 2)

    warped = F.grid_sample(ctx_latent, samp, mode="bilinear",
                           padding_mode="zeros", align_corners=True)
    valid = ((gx >= -1) & (gx <= 1) & (gy >= -1) & (gy <= 1)).float()
    return warped, valid.reshape(B, 1, Hl, Wl)              # in-FOV co-visibility mask


def reprojection_confidence_map(x0_pred, ctx_latents, H_rels,
                                patch=2, tau=1.0):
    """
    Per-token retrieval-confidence target from reprojection correspondence.

    x0_pred:     (B, C, T, Hl, Wl)   recovered x0 for TARGET tokens (per latent frame)
    ctx_latents: (B, C, K, Hl, Wl)   clean context latents (VAE-encoded history)
    H_rels:      (B, T, K, 3, 3)     target-frame t <- context-frame k homographies
                                     (latent-pixel coords), from convert_rt_to_relative
    Returns:
        conf_tok:  (B, 1, T, Hl//patch, Wl//patch)  in [0,1], token-resolution
        mask_tok:  (B, 1, T, Hl//patch, Wl//patch)  co-visibility (any context covers token)
    """
    B, C, T, Hl, Wl = x0_pred.shape
    K = ctx_latents.shape[2]
    best_err = x0_pred.new_full((B, 1, T, Hl, Wl), float("inf"))
    any_valid = x0_pred.new_zeros((B, 1, T, Hl, Wl))

    for t in range(T):
        for k in range(K):
            warped, valid = warp_context_latent(ctx_latents[:, :, k], H_rels[:, t, k])
            err = (x0_pred[:, :, t] - warped).pow(2).mean(dim=1, keepdim=True)   # (B,1,Hl,Wl)
            err = torch.where(valid > 0, err, best_err[:, :, t])
            best_err[:, :, t] = torch.minimum(best_err[:, :, t], err)            # best matching ctx frame
            any_valid[:, :, t] = torch.maximum(any_valid[:, :, t], valid)

    best_err = torch.where(torch.isfinite(best_err), best_err, torch.zeros_like(best_err))
    conf = torch.exp(-best_err / tau)              # low reprojection error -> high confidence
    conf = conf * any_valid                        # undefined where nothing is co-visible

    # Patch-pool latent-pixel grid (44x80) down to token grid (22x40) to match `um`.
    conf_tok = F.avg_pool3d(conf, kernel_size=(1, patch, patch))
    mask_tok = (F.avg_pool3d(any_valid, kernel_size=(1, patch, patch)) > 0).float()
    return conf_tok, mask_tok
```

**Where `H_rels` comes from.** In `training_loss` you already have (or can pass
through `inputs`) the per-frame RTs. For target latent frame *t* and context
frame *k*: `rel = convert_rt_to_relative([rt_k], ref_rt=rt_t)[0]`, parse into
`(R_rel, t_rel)`, and convert to a latent-pixel homography with the intrinsics
scaled by 1/8 (latent) — under the paper's XY+yaw planar setting a homography is
the correct first-order model. Precompute `H_rels` on CPU/numpy in the dataloader
(the RTs are already loaded for the action MLP) and pass them in as a tensor;
avoid per-step Python geometry in the hot loop.

#### 6.3 Two ways to use the reprojection map

This is the bridge to Category B / Category C from the discussion:

- **(B) As a near-ground-truth confidence map directly** — `conf_tok` *is* a
  retrieval-confidence map, no learning required. Use it to reweight the main
  loss (`w = conf_tok` upweights tokens that *should* be retrievable, focusing
  the memory pathway on co-visible content) or as an eval-time diagnostic over
  the revisit tail.

- **(C) As the supervision target for a predictive head** — train the
  `UncertaintyHead` (or a dedicated `ConfidenceHead`) to **predict `conf_tok`
  before generation**, supervised only on `mask_tok` tokens:

  ```python
  pred_conf = torch.sigmoid(-um)                      # head's confidence in [0,1]
  loss_conf = (mask_tok * (pred_conf - conf_tok.detach()).pow(2)).sum() \
              / mask_tok.sum().clamp(min=1.0)
  loss = loss + lambda_conf * loss_conf
  ```

  This gives a **calibrated, forward-time** confidence signal grounded in
  geometry, instead of the self-supervised heteroscedastic target — and it
  cleanly answers "do we know, during training, whether we retrieved the correct
  context?": *yes, because geometry tells us which tokens had retrievable
  support and reprojection tells us whether the model reproduced it.*

### How to know, during training, if retrieval is correct — summary

| Target level | Source in repo | "Correct" means | Strength / caveat |
| --- | --- | --- | --- |
| **0 Reconstruction** | `target_latents` (base loss) | low token MSE | weak — confounds forgetting vs. novelty |
| **1 Co-visibility** | `overlap_labels/*.json`, `compute_fov_overlap_3d` | model uses the geometrically co-visible frames | frame/region-level; FOV-frustum, **not depth/occlusion aware** |
| **2 Reprojection** | relative RT (`convert_rt_to_relative`) + warp | warped co-visible evidence matches `x0_pred`, masked to co-visible tokens | per-token, strongest; needs depth or planar/homography assumption |

**Honesty caveat (state this in any writeup):** the overlap labels and
`compute_fov_overlap_3d` are **camera-frustum** co-visibility from poses, *not*
depth-aware occlusion — two frames can be marked co-visible when an occluder
blocks the shared content. Level-2 reprojection inherits this: the homography
regime assumes near-planar / yaw-dominant motion (the paper's `constrain_to_xy`
setting). For truly metric per-token correspondence, use depth (better available
in the SpatialVID dynamic pool) and a full reprojection rather than a homography.

---

## 7. Depth-aware per-token confidence (the accurate version)

With depth you replace the §6.2 **homography** (planar, yaw-dominant
approximation) by a **full metric reprojection with occlusion reasoning**. This
removes the two failure modes of the homography path: (i) it handles arbitrary
3D scene geometry and 6-DoF motion, not just a reference plane, and (ii) it can
*detect occlusion* — telling apart "co-visible and the model retrieved it" from
"the frustum overlaps but an occluder hides the content" (the exact blind spot
of the FOV-frustum labels in §6, Level 1).

### 7.1 Pose convention in this repo (get the direction right)

`fov_retrieval.compute_fov_overlap_3d` treats `position` as the **camera centre
in world coordinates** `C` and the third column `R[:,2]` as the world-space
**forward** axis. So the stored 12-dim RT `[t | R]` is **camera-to-world**:

```
X_world = R · X_cam + C          # R = R_cam→world,  C = camera centre = t
X_cam   = Rᵀ · (X_world − C)     # world → camera (inverse)
```

(`R⁻¹ = Rᵀ` for a rotation). This is the opposite direction from a
"world-to-camera extrinsic" `[R|t]` convention — using the wrong one silently
flips the reprojection, so anchor on this.

### 7.2 The reprojection (target token → 3D → context frame)

For a target latent token at pixel `p_t=(u,v)` in latent-pixel coords with
metric depth `d`:

1. **Back-project to the target camera ray, scale by depth:**
   `X_cam_t = d · K⁻¹ · [u, v, 1]ᵀ`
2. **Target camera → world** (camera-to-world):
   `X_world = R_t · X_cam_t + C_t`
3. **World → context camera k:**
   `X_cam_k = R_kᵀ · (X_world − C_k)`
4. **Project into context frame k:**
   `p_k = K · X_cam_k / z_k`, where `z_k = X_cam_k.z`

`K` is the **latent-resolution** intrinsic: build it from the FOV
(`fov=52.67°`, the same constant `compute_fov_overlap_3d` uses) and divide focal
length + principal point by the VAE spatial factor 8 (so it acts on the 44×80
latent-pixel grid, matching §6.1). Then sample the context latent at `p_k` and,
crucially, **also sample the context depth at `p_k`** for the occlusion test.

### 7.3 Occlusion test (what depth buys you)

A target point is genuinely visible in context frame k only if its reprojected
depth `z_k` matches the context frame's own recorded depth at `p_k`. If the
context depth is *closer* than `z_k`, something else occludes the point — mark
it **not co-visible** even though the frustum overlaps:

```
visible_k = (z_k ≤ depth_ctx_k(p_k) · (1 + occ_thresh))
```

This is a forward z-buffer check (`occ_thresh` ~0.05–0.1 absorbs depth noise).
It is exactly the discriminator the §6 Level-1 labels lack.

### 7.4 Sketch (syntax-checked)

```python
# diffsynth/models/memory/reproj_confidence_depth.py  (sketch)
import torch
import torch.nn.functional as F


def intrinsics_latent(width_px, height_px, fov_deg=52.67, vae_down=8):
    """Latent-resolution pinhole intrinsics K (focal & principal point ÷ VAE factor)."""
    import math
    Wl, Hl = width_px // vae_down, height_px // vae_down
    f_px = (width_px / 2.0) / math.tan(math.radians(fov_deg) / 2.0)
    f_lat = f_px / vae_down
    K = torch.tensor([[f_lat, 0.0, Wl / 2.0],
                      [0.0, f_lat, Hl / 2.0],
                      [0.0, 0.0, 1.0]])
    return K, Hl, Wl


def reproject_target_to_context(depth_t, R_t, C_t, R_k, C_k, K, Kinv):
    """
    Map every target latent-pixel into context frame k via depth + camera-to-world RT.
    depth_t: (B,1,Hl,Wl) metric depth of TARGET latent frame
    R_t,R_k: (B,3,3) camera->world rotations ;  C_t,C_k: (B,3) camera centres
    returns: grid (B,Hl,Wl,2) for grid_sample, in_fov mask (B,1,Hl,Wl),
             z_k_map (B,1,Hl,Wl) reprojected depth in context camera
    """
    B, _, Hl, Wl = depth_t.shape
    dev = depth_t.device
    ys, xs = torch.meshgrid(torch.arange(Hl, device=dev, dtype=torch.float32),
                            torch.arange(Wl, device=dev, dtype=torch.float32),
                            indexing="ij")
    ones = torch.ones_like(xs)
    pix = torch.stack([xs, ys, ones], -1).reshape(1, Hl * Wl, 3).expand(B, -1, -1)
    ray = torch.bmm(pix, Kinv.transpose(1, 2))             # K^-1 [u,v,1]
    d = depth_t.reshape(B, Hl * Wl, 1)
    Xc_t = ray * d                                         # target camera coords
    Xw = torch.bmm(Xc_t, R_t.transpose(1, 2)) + C_t.reshape(B, 1, 3)   # cam->world
    Xc_k = torch.bmm(Xw - C_k.reshape(B, 1, 3), R_k)       # world->context cam (R_k^T via right-mul)
    z_k = Xc_k[..., 2:3].clamp(min=1e-6)
    proj = torch.bmm(Xc_k / z_k, K.transpose(1, 2))
    u, v = proj[..., 0], proj[..., 1]
    gx = (u / (Wl - 1)) * 2 - 1
    gy = (v / (Hl - 1)) * 2 - 1
    grid = torch.stack([gx, gy], -1).reshape(B, Hl, Wl, 2)
    in_fov = ((gx >= -1) & (gx <= 1) & (gy >= -1) & (gy <= 1)).float().reshape(B, 1, Hl, Wl)
    return grid, in_fov, z_k.reshape(B, 1, Hl, Wl)


def depth_aware_confidence(x0_pred, ctx_latents, ctx_depths, depth_t,
                           R_t, C_t, R_k_list, C_k_list, K,
                           patch=2, tau=1.0, occ_thresh=0.1):
    """
    Per-token retrieval-confidence map via metric reprojection + occlusion test.
    x0_pred:    (B,C,T,Hl,Wl)   recovered x0 for target tokens
    ctx_latents:(B,C,K,Hl,Wl)   clean context latents ;  ctx_depths:(B,1,K,Hl,Wl)
    depth_t:    (B,1,T,Hl,Wl)   target-frame metric depth
    R_t,C_t:    (B,T,3,3),(B,T,3)   target cam->world per latent frame
    R_k_list,C_k_list: lists of (B,3,3),(B,3) per context frame
    """
    B, C, T, Hl, Wl = x0_pred.shape
    Kb = K.unsqueeze(0).expand(B, -1, -1)
    Kinv = torch.inverse(K).unsqueeze(0).expand(B, -1, -1)
    best_err = x0_pred.new_full((B, 1, T, Hl, Wl), float("inf"))
    any_valid = x0_pred.new_zeros((B, 1, T, Hl, Wl))
    for t in range(T):
        for k in range(len(R_k_list)):
            grid, in_fov, z_proj = reproject_target_to_context(
                depth_t[:, :, t], R_t[:, t], C_t[:, t], R_k_list[k], C_k_list[k], Kb, Kinv)
            warped = F.grid_sample(ctx_latents[:, :, k], grid, mode="bilinear",
                                   padding_mode="zeros", align_corners=True)
            ctx_z = F.grid_sample(ctx_depths[:, :, k], grid, mode="bilinear",
                                  padding_mode="zeros", align_corners=True)
            visible = (z_proj <= ctx_z * (1.0 + occ_thresh)).float()   # z-buffer occlusion test
            valid = in_fov * visible
            err = (x0_pred[:, :, t] - warped).pow(2).mean(1, keepdim=True)
            err = torch.where(valid > 0, err, best_err[:, :, t])
            best_err[:, :, t] = torch.minimum(best_err[:, :, t], err)
            any_valid[:, :, t] = torch.maximum(any_valid[:, :, t], valid)
    best_err = torch.where(torch.isfinite(best_err), best_err, torch.zeros_like(best_err))
    conf = torch.exp(-best_err / tau) * any_valid
    conf_tok = F.avg_pool3d(conf, (1, patch, patch))
    mask_tok = (F.avg_pool3d(any_valid, (1, patch, patch)) > 0).float()
    return conf_tok, mask_tok
```

### 7.5 Getting depth into latent-token space

- **Source.** The SpatialVID dynamic pool carries richer geometry than the
  static pool; if per-frame metric depth is not already exported, run a monocular
  depth estimator offline and cache it (do **not** add it to the training hot
  loop). The static in-domain pool only has camera poses, so depth-aware
  confidence is primarily a **dynamic-pool** technique.
- **Resolution.** Downsample depth to the **latent-pixel grid** (÷8 → 44×80) by
  *area/min pooling* (min-pool preserves near surfaces for the occlusion test;
  avoid bilinear across depth discontinuities, which invents mid-air depths).
- **Scale.** Metric consistency matters — `z_k` (reprojected) and
  `depth_ctx_k` must be in the **same units**. If depth is up-to-scale
  (monocular), fit a per-video scale so it is consistent with the RT translation
  units, or make the occlusion test **relative** (compare normalised depth
  ranks) instead of absolute.
- **Plumbing.** Precompute and pass `depth_t`, `ctx_depths`, and the per-frame
  `(R, C)` through `inputs` (the RTs are already loaded for the action MLP — see
  §6). Keep the double loop over `T×K` out of the innermost step by vectorising
  over `k`, or restrict `k` to the top-N co-visible frames from the §6 Level-1
  overlap labels (cheaper and removes obviously-irrelevant frames first).

### 7.6 Accuracy ladder (how the targets compare)

| Variant | Geometry model | Occlusion | Needs | Accuracy |
| --- | --- | --- | --- | --- |
| §6 Level-1 co-visibility | camera frustum (poses only) | ✗ | poses | frame/region |
| §6.2 homography | planar / yaw-dominant | ✗ | poses + plane | per-token, approx |
| **§7 depth reprojection** | full 6-DoF metric | **✓ (z-buffer)** | poses + **depth** | **per-token, metric** |

The depth-aware map plugs into the **same two consumers** as §6.3: use `conf_tok`
directly to reweight the main loss, or as the supervision target for a
forward-time `ConfidenceHead` (masked on `mask_tok`). The only change is a
strictly more accurate, occlusion-aware target.

**Caveats specific to depth.** Reprojection confidence is now bounded by *depth
quality*: noisy/biased monocular depth produces false occlusions and warp
errors. Mitigate with a tolerant `occ_thresh`, min-pooled latent depth, and —
when in doubt — fall back to the §6.2 homography or §6 Level-1 mask for frames
whose depth is low-confidence. Dynamic/independently-moving objects also break
the static-scene assumption of any reprojection (the point moved between
frames); mask known-dynamic regions out of the retrieval loss where you can
detect them.