echo / code /retrieve.md
amonshano's picture
Add Echo-Memory codebase used for this run (CC BY 4.0, JD Echo Team) (part 4)
c335050 verified
|
Raw
History Blame Contribute Delete
37.9 kB

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 − x0the 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:

    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.pyWanVideoPipeline.__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:

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:

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):

("--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:

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:

    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:

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

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

# 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 directlyconf_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:

    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)

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