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.
|