from collections.abc import Mapping, Sequence import csv from dataclasses import dataclass import os import random import time from pathlib import Path import numpy as np import torch import torch.distributed as dist import torch.nn.functional as F import torchvision.transforms.functional as TF from torchvision.transforms import InterpolationMode from PIL import Image, ImageDraw from einops import rearrange from tqdm import tqdm from omegaconf import DictConfig, open_dict from lightning.pytorch.utilities.types import STEP_OUTPUT from algorithms.common.metrics import ( LearnedPerceptualImagePatchSimilarity, ) from datasets.video.memory_selection import ( _build_shared_fov_candidate_pool, _dynamic_multiview_selector, _dynamic_policy, _event_triggered_anchor_candidates_from_deltas, _pose_delta_values, _select_dynamic_from_stream, _select_anchor, _select_dynamic_by_policy, _select_revisit, ) from utils.logging_utils import log_video, get_validation_metrics_for_videos from .df_base import DiffusionForcingBase from .models.vae import VAE_models from .models.diffusion import Diffusion from .models.pose_prediction import PosePredictionNet import glob # Utility Functions _DEMEMWM_SEGMENT_KEYS = ("target", "anchor", "dynamic", "revisit") _DEMEMWM_STREAM_KEYS = ("anchor", "dynamic", "revisit") _DEMEMWM_REFERENCE_ATTN_MARKERS = (".r_attn_anchor.", ".r_attn_dynamic.", ".r_attn_revisit.") _DEMEMWM_GEOMETRY_PROJ_MARKERS = (".query_pose_proj.", ".key_pose_proj.", ".timestamp_embedding.") _DEMEMWM_ADALN_MLP_MARKERS = (".r_adaLN_modulation.", ".r_mlp.") @dataclass(frozen=True) class RevisitPair: clip_id: str source_index: int target_index: int source_frame: int target_frame: int gap: int fov_overlap: float plucker_overlap: float position_distance: float yaw_delta_deg: float pitch_delta_deg: float positional: bool def _best_revisit_pair_by_target(pairs: Sequence[RevisitPair]) -> dict[int, RevisitPair]: best: dict[int, RevisitPair] = {} for pair in pairs: rank = (float(pair.fov_overlap), float(pair.plucker_overlap), int(pair.gap), -int(pair.source_index)) current = best.get(int(pair.target_index)) if current is None: best[int(pair.target_index)] = pair continue current_rank = ( float(current.fov_overlap), float(current.plucker_overlap), int(current.gap), -int(current.source_index), ) if rank > current_rank: best[int(pair.target_index)] = pair return best def _metric_psnr_from_mse(mse: torch.Tensor) -> torch.Tensor: return 10.0 * torch.log10(1.0 / mse.clamp_min(torch.finfo(mse.dtype).eps)) def _metric_central_crop(tensor: torch.Tensor, crop_fraction: float) -> torch.Tensor: crop_fraction = float(crop_fraction) if crop_fraction == 1.0: return tensor h = int(tensor.shape[-2]) w = int(tensor.shape[-1]) crop_h = max(1, int(round(h * crop_fraction))) crop_w = max(1, int(round(w * crop_fraction))) top = max(0, (h - crop_h) // 2) left = max(0, (w - crop_w) // 2) return tensor[..., top : top + crop_h, left : left + crop_w] def _metric_lpips_scalar(lpips_model, pred: torch.Tensor, target: torch.Tensor) -> float: if lpips_model is None: return float("nan") pred = torch.clamp(pred.float(), 0.0, 1.0) target = torch.clamp(target.float(), 0.0, 1.0).to(device=pred.device) if hasattr(lpips_model, "reset"): lpips_model.reset() with torch.no_grad(): try: value = lpips_model(pred, target) except TypeError: lpips_model.update(pred, target) value = lpips_model.compute() if hasattr(lpips_model, "reset"): lpips_model.reset() if torch.is_tensor(value): return float(value.detach().float().mean().cpu().item()) return float(value) def compute_frame_metric_rows( *, branch: str, clip_id: str, pred: torch.Tensor, gt: torch.Tensor | None, poses: torch.Tensor | None, pairs: Sequence[RevisitPair], context_frames: int, lpips_model=None, self_crop_fraction: float = 0.50, focal_length: float = 0.35, compute_plucker_similarity: bool = True, ) -> list[dict[str, object]]: if pred.ndim == 5: pred = pred[:, 0] if gt is not None and gt.ndim == 5: gt = gt[:, 0] pred = pred.detach() gt = None if gt is None else gt.detach().to(device=pred.device) best_pairs = _best_revisit_pair_by_target(pairs) gt_psnr_values = None if gt is not None: gt_mse = (torch.clamp(pred.float(), 0.0, 1.0) - torch.clamp(gt.float(), 0.0, 1.0)).square() gt_psnr_values = _metric_psnr_from_mse(gt_mse.flatten(start_dim=1).mean(dim=1)).detach().cpu().tolist() rows: list[dict[str, object]] = [] for out_index in range(int(pred.shape[0])): frame_index = int(context_frames) + out_index pair = best_pairs.get(frame_index) gt_lpips = float("nan") gt_psnr = float("nan") if gt is not None and gt_psnr_values is not None: gt_psnr = float(gt_psnr_values[out_index]) gt_lpips = _metric_lpips_scalar(lpips_model, pred[out_index : out_index + 1], gt[out_index : out_index + 1]) self_psnr = float("nan") self_lpips = float("nan") source_index = "" gap = "" fov_overlap = "" plucker_overlap = "" self_status = "not_revisit" if pair is not None: source_index = int(pair.source_index) gap = int(pair.gap) fov_overlap = float(pair.fov_overlap) plucker_overlap = float(pair.plucker_overlap) source_out_index = int(pair.source_index) - int(context_frames) if 0 <= source_out_index < int(pred.shape[0]): pred_source = _metric_central_crop(pred[source_out_index : source_out_index + 1], self_crop_fraction) pred_target = _metric_central_crop(pred[out_index : out_index + 1], self_crop_fraction) self_mse = (torch.clamp(pred_source.float(), 0.0, 1.0) - torch.clamp(pred_target.float(), 0.0, 1.0)).square().mean() self_psnr = float(_metric_psnr_from_mse(self_mse.reshape(1))[0].detach().cpu().item()) self_lpips = _metric_lpips_scalar(lpips_model, pred_target, pred_source) self_status = "central_crop" else: self_status = "source_before_prediction_horizon" rows.append( { "branch": str(branch), "clip_id": str(clip_id), "frame_index": int(frame_index), "output_index": int(out_index), "is_revisit": pair is not None, "source_index": source_index, "gap": gap, "fov_overlap": fov_overlap, "plucker_overlap": plucker_overlap, "gt_lpips": gt_lpips, "gt_psnr": gt_psnr, "self_lpips": self_lpips, "self_psnr": self_psnr, "self_plucker_similarity": float("nan"), "self_consistency_status": self_status, } ) return rows def _cfg_get(cfg, key: str, default=None): if cfg is None: return default if isinstance(cfg, Mapping): return cfg.get(key, default) return getattr(cfg, key, default) def _cuda_device_index(device): device = torch.device(device) if device.type != "cuda" or not torch.cuda.is_available(): return None return torch.cuda.current_device() if device.index is None else device.index def _cuda_vram_postfix(device): device_index = _cuda_device_index(device) if device_index is None: return None gib = 1024**3 allocated = torch.cuda.memory_allocated(device_index) / gib reserved = torch.cuda.memory_reserved(device_index) / gib peak = torch.cuda.max_memory_allocated(device_index) / gib return f"{allocated:.1f}/{reserved:.1f}G peak {peak:.1f}G" def _trainability_cfg(cfg): return _cfg_get(cfg, "trainability", {}) def _dememwm_full_dit_active(trainability, global_step: int) -> bool: if not bool(_cfg_get(trainability, "train_full_dit", False)): return False start_step = _cfg_get(trainability, "full_dit_start_step", 0) return start_step is None or int(global_step) >= int(start_step) def _is_dememwm_memory_trainable_parameter(name: str, trainability) -> bool: is_geometry = any(marker in name for marker in _DEMEMWM_GEOMETRY_PROJ_MARKERS) if bool(_cfg_get(trainability, "geometry_projections", True)) and is_geometry: return True if bool(_cfg_get(trainability, "reference_attention", True)) and not is_geometry: if any(marker in name for marker in _DEMEMWM_REFERENCE_ATTN_MARKERS): return True if bool(_cfg_get(trainability, "adaln_mlp", True)): return any(marker in name for marker in _DEMEMWM_ADALN_MLP_MARKERS) return False def _apply_dememwm_trainability(diffusion_model, vae, trainability, global_step: int) -> None: full_dit_active = _dememwm_full_dit_active(trainability, global_step) for name, param in diffusion_model.named_parameters(): param.requires_grad_(full_dit_active or _is_dememwm_memory_trainable_parameter(name, trainability)) if vae is not None: freeze_vae = bool(_cfg_get(trainability, "freeze_vae", True)) for param in vae.parameters(): param.requires_grad_(not freeze_vae) def _dememwm_optimizer_parameters(diffusion_model, vae, trainability): include_full_dit = bool(_cfg_get(trainability, "train_full_dit", False)) for name, param in diffusion_model.named_parameters(): if include_full_dit or _is_dememwm_memory_trainable_parameter(name, trainability): yield param if vae is not None and not bool(_cfg_get(trainability, "freeze_vae", True)): yield from vae.parameters() def _dememwm_group_target_lr(group_name: str, trainability, default_lr, global_step: int) -> float: lr_cfg = _cfg_get(trainability, "lr", {}) if group_name == "memory_modules": return float(_cfg_get(lr_cfg, "memory_modules", default_lr)) if group_name == "base_dit": if not _dememwm_full_dit_active(trainability, global_step): return 0.0 return float(_cfg_get(lr_cfg, "base_dit", default_lr)) return float(default_lr) def _dememwm_group_warmup_start_step(group_name: str, trainability) -> int: if group_name != "base_dit": return 0 start_step = _cfg_get(trainability, "full_dit_start_step", 0) return 0 if start_step is None else int(start_step) def _dememwm_optimizer_parameter_groups(diffusion_model, vae, trainability, default_lr, global_step: int): include_full_dit = bool(_cfg_get(trainability, "train_full_dit", False)) grouped = {"memory_modules": [], "base_dit": []} for name, param in diffusion_model.named_parameters(): if _is_dememwm_memory_trainable_parameter(name, trainability): grouped["memory_modules"].append(param) elif include_full_dit: grouped["base_dit"].append(param) param_groups = [] for group_name, params in grouped.items(): if params: target_lr = _dememwm_group_target_lr(group_name, trainability, default_lr, global_step) param_groups.append({ "params": params, "lr": target_lr, "target_lr": target_lr, "warmup_start_step": _dememwm_group_warmup_start_step(group_name, trainability), "name": group_name, }) if vae is not None and not bool(_cfg_get(trainability, "freeze_vae", True)): target_lr = float(_cfg_get(_cfg_get(trainability, "lr", {}), "vae", default_lr)) param_groups.append({"params": tuple(vae.parameters()), "lr": target_lr, "target_lr": target_lr, "name": "vae"}) return param_groups def _apply_dememwm_optimizer_group_lrs(optimizer, trainability, default_lr, global_step: int) -> None: for param_group in optimizer.param_groups: group_name = param_group.get("name") if group_name in {"memory_modules", "base_dit"}: param_group["target_lr"] = _dememwm_group_target_lr(group_name, trainability, default_lr, global_step) param_group["warmup_start_step"] = _dememwm_group_warmup_start_step(group_name, trainability) def _derive_memory_condition_length(cfg) -> int: memory_cfg = _cfg_get(cfg, "memory_selection") if memory_cfg is None: value = _cfg_get(cfg, "memory_condition_length") if value is None: raise ValueError("DeMemWM requires memory_selection or memory_condition_length") return int(value) return sum(int(_cfg_get(memory_cfg, f"max_{key}_frames", 0)) for key in _DEMEMWM_STREAM_KEYS) def _segment_value_to_int(value, key): if torch.is_tensor(value): flat = value.reshape(-1) if flat.numel() == 0: raise ValueError(f"memory_segments['{key}'] is empty") first = flat[0] if flat.numel() > 1 and not bool(torch.all(flat == first).item()): raise ValueError(f"memory_segments['{key}'] must be identical across the batch") return int(first.item()) if isinstance(value, (list, tuple)): values = [_segment_value_to_int(item, key) for item in value] if not values: raise ValueError(f"memory_segments['{key}'] is empty") if any(item != values[0] for item in values[1:]): raise ValueError(f"memory_segments['{key}'] must be identical across the batch") return values[0] return int(value) def _normalize_dememwm_image_hw(image_hw, batch_size): image_hw = image_hw if torch.is_tensor(image_hw) else torch.as_tensor(image_hw) image_hw = image_hw.to(dtype=torch.long) if image_hw.ndim == 1: image_hw = image_hw.unsqueeze(0) if image_hw.shape[0] == 1 and batch_size != 1: image_hw = image_hw.expand(batch_size, -1) return image_hw.contiguous() def _preprocess_dememwm_latent_batch(batch): memory_segments = { key: _segment_value_to_int(batch["memory_segments"][key], key) for key in _DEMEMWM_SEGMENT_KEYS } segment_lengths = dict(memory_segments) target_length = segment_lengths["target"] stream_lengths = {key: segment_lengths[key] for key in _DEMEMWM_STREAM_KEYS} memory_masks = {key: batch["memory_masks"][key] for key in _DEMEMWM_SEGMENT_KEYS} # Latent dataset batches are already VAE-encoded: B x T_all x C x H_lat x W_lat. latents = rearrange(batch["latents"], "b t c ... -> t b c ...").contiguous() actions = rearrange(batch["actions"], "b t d -> t b d").contiguous() # Dataset poses are the frame-memory geometry metadata in packed T_all x B x 5 order. poses = rearrange(batch["poses"], "b t d -> t b d").contiguous() frame_indices = rearrange(batch["frame_indices"], "b t -> t b").contiguous() image_hw = _normalize_dememwm_image_hw(batch["image_hw"], latents.shape[1]) segment_slices = {} start = 0 for key in _DEMEMWM_SEGMENT_KEYS: stop = start + segment_lengths[key] segment_slices[key] = slice(start, stop) start = stop target_slice = segment_slices["target"] stream_slices = {key: segment_slices[key] for key in _DEMEMWM_STREAM_KEYS} if start != latents.shape[0]: raise ValueError( f"memory_segments sum to {start} frames, but latent batch has {latents.shape[0]}" ) action_conditions = actions.clone() if target_length: action_conditions[target_slice.start:target_slice.start + 1] = 0 for stream_slice in stream_slices.values(): action_conditions[stream_slice] = 0 sequence_tensors = { "latents": latents, "actions": actions, "action_conditions": action_conditions, "poses": poses, "frame_indices": frame_indices, } # Packed sequence order stays [target][anchor][dynamic][revisit] in T x B layout. segments = { key: {name: tensor[segment_slices[key]] for name, tensor in sequence_tensors.items()} for key in _DEMEMWM_SEGMENT_KEYS } target_tensors = segments["target"] stream_tensors = {key: segments[key] for key in _DEMEMWM_STREAM_KEYS} # Keep original image H/W for later ray geometry; latent H/W is not a substitute. return { "latents": latents, "actions": actions, "action_conditions": action_conditions, "poses": poses, "frame_memory_pose": poses, "frame_indices": frame_indices, "memory_segments": memory_segments, "segment_lengths": segment_lengths, "target_length": target_length, "stream_lengths": stream_lengths, "segment_slices": segment_slices, "target_slice": target_slice, "stream_slices": stream_slices, "segments": segments, "target_tensors": target_tensors, "stream_tensors": stream_tensors, "memory_masks": memory_masks, "image_hw": image_hw, } def _gather_online_memory_tensor(source, indices, masks): output_shape = (indices.shape[1], indices.shape[0], *source.shape[2:]) if output_shape[0] == 0 or source.shape[0] == 0: return source.new_zeros(output_shape) gather_idx = indices.clamp(0, source.shape[0] - 1).T.to(source.device) batch_idx = torch.arange(indices.shape[0], device=source.device).expand_as(gather_idx) gathered = source[gather_idx, batch_idx] mask = masks.T.to(device=source.device, dtype=torch.bool) mask = mask.view(*mask.shape, *((1,) * (gathered.ndim - 2))) return gathered * mask.to(dtype=gathered.dtype) def _select_online_anchor_indices(n_context_frames: int, count: int, cfg, poses=None) -> np.ndarray: if count <= 0: return np.empty((0,), dtype=np.int64) context_count = max(0, int(n_context_frames)) if poses is not None: context_count = min(context_count, len(poses)) candidates = np.arange(0, context_count, dtype=np.int64) return _select_anchor(candidates, int(count), cfg, poses=poses) def _select_online_dynamic_indices( start_frame: int, count: int, cfg=None, latents=None, actions=None, poses=None, reference_frames=None, excluded=None, ) -> np.ndarray: return _select_dynamic_by_policy( int(start_frame), int(count), cfg, poses=poses, latents=latents, actions=actions, min_candidate_frame=0, reference_frames=reference_frames, excluded=excluded, ) def _new_online_event_cache(capacity=None): return { "stop": 0, "capacity": None if capacity is None else max(0, int(capacity)), "anchors": np.empty((0,), dtype=np.int64), "vectors": None, "d_vis": None, "d_pose": None, "d_act": None, } def _online_latent_frame_vectors(latents): if latents is None or len(latents) == 0: return None tensor = latents if torch.is_tensor(latents) else torch.as_tensor(np.asarray(latents), dtype=torch.float32) tensor = tensor.to(dtype=torch.float32) if tensor.ndim == 1: return tensor[:, None] if tensor.ndim == 2: return tensor return tensor.reshape(tensor.shape[0], tensor.shape[1], -1).mean(dim=-1) def _online_flat_frame_tensor(values, start: int, stop: int, device): if values is None or stop <= start: return None length = len(values) start = max(0, min(int(start), length)) stop = max(start, min(int(stop), length)) if stop <= start: return None sliced = values[start:stop] tensor = sliced if torch.is_tensor(sliced) else torch.as_tensor(np.asarray(sliced), dtype=torch.float32, device=device) tensor = tensor.to(device=device, dtype=torch.float32) if tensor.ndim == 1: return tensor[:, None] return tensor.reshape(tensor.shape[0], -1) def _online_l2_new_deltas(values, old_stop: int, stop: int, device): deltas = torch.zeros((stop - old_stop,), device=device, dtype=torch.float32) if values is None or stop <= old_stop: return deltas value_stop = max(old_stop, min(int(stop), len(values))) if value_stop <= old_stop: return deltas value_start = max(0, old_stop - 1) tensor = _online_flat_frame_tensor(values, value_start, value_stop, device) if tensor is None or tensor.shape[0] <= 1: return deltas consecutive = torch.linalg.vector_norm(tensor[1:] - tensor[:-1], dim=-1) frames = torch.arange(old_stop, value_stop, device=device, dtype=torch.long) valid = frames > 0 if bool(valid.any()): valid_frames = frames[valid] rows = valid_frames - (value_start + 1) deltas.index_copy_(0, valid_frames - old_stop, consecutive.index_select(0, rows)) return deltas def _append_online_event_values(cached, values, old_stop: int, stop: int, capacity): if cached is None: if capacity is None: return values output = values.new_empty((capacity, *values.shape[1:])) output[old_stop:stop] = values return output if capacity is not None: cached[old_stop:stop] = values.to(device=cached.device) return cached return torch.cat([cached, values.to(device=cached.device)], dim=0) def _extend_online_event_cache(cache, latents, actions, poses, stop: int, cfg=None) -> None: if latents is None: return old_stop = int(cache["stop"]) target_stop = max(0, min(int(stop), len(latents))) if target_stop <= old_stop: return capacity = cache["capacity"] if capacity is not None and target_stop > capacity: for key in ("vectors", "d_vis", "d_pose", "d_act"): if cache[key] is not None: cache[key] = cache[key][:old_stop].contiguous() cache["capacity"] = None capacity = None new_vectors = _online_latent_frame_vectors(latents[old_stop:target_stop]) if new_vectors is None: return cache["vectors"] = _append_online_event_values(cache["vectors"], new_vectors, old_stop, target_stop, capacity) vectors = cache["vectors"] device = vectors.device # The prefix grows monotonically online; only new consecutive deltas need the # latent/action/pose preprocessing, while robust-z event scoring stays exact. new_count = target_stop - old_stop frame_rows = torch.arange(old_stop, target_stop, device=device, dtype=torch.long) d_vis = torch.zeros((new_count,), device=device, dtype=torch.float32) valid_vis = frame_rows > 0 if bool(valid_vis.any()): curr_rows = frame_rows[valid_vis] curr = vectors.index_select(0, curr_rows) prev = vectors.index_select(0, curr_rows - 1) curr_norm = torch.linalg.vector_norm(curr, dim=-1) prev_norm = torch.linalg.vector_norm(prev, dim=-1) cosine = (curr * prev).sum(dim=-1) / (curr_norm * prev_norm).clamp_min(1e-6) valid_pair = (curr_norm > 1e-6) & (prev_norm > 1e-6) cosine = torch.where(valid_pair, cosine.clamp(-1.0, 1.0), torch.ones_like(cosine)) d_vis.index_copy_(0, curr_rows - old_stop, 1.0 - cosine) pose_frames = np.arange(old_stop, target_stop, dtype=np.int64) d_pose = _pose_delta_values(poses, pose_frames, max(0, old_stop - 1), target_stop, device) d_act = _online_l2_new_deltas(actions, old_stop, target_stop, device) cache["d_vis"] = _append_online_event_values(cache["d_vis"], d_vis, old_stop, target_stop, capacity) cache["d_pose"] = _append_online_event_values(cache["d_pose"], d_pose, old_stop, target_stop, capacity) cache["d_act"] = _append_online_event_values(cache["d_act"], d_act, old_stop, target_stop, capacity) cache["stop"] = target_stop frames = np.arange(0, target_stop, dtype=np.int64) event_anchors = _event_triggered_anchor_candidates_from_deltas( frames, cache["d_vis"][:target_stop], cache["d_pose"][:target_stop], cache["d_act"][:target_stop], cfg, ) stream = np.concatenate([np.asarray([0], dtype=np.int64), event_anchors.astype(np.int64, copy=False)]) stream = stream[(stream >= 0) & (stream < target_stop)] cache["anchors"] = np.unique(np.sort(stream.astype(np.int64, copy=False))) def _select_online_event_dynamic_from_cache( cache, start_frame: int, count: int, cfg, reference_frames=None, ) -> np.ndarray: if count <= 0: return np.empty((0,), dtype=np.int64) dynamic_cfg = _cfg_get(cfg, "dynamic", {}) max_event_anchors = _cfg_get(dynamic_cfg, "max_event_anchors") if max_event_anchors is not None: count = min(int(count), max(0, int(max_event_anchors))) if count <= 0: return np.empty((0,), dtype=np.int64) stop = min(max(0, int(start_frame)), int(cache["stop"])) anchors = np.asarray(cache.get("anchors", np.empty((0,), dtype=np.int64)), dtype=np.int64) eligible_stream = anchors[(anchors >= 0) & (anchors < stop)] return _select_dynamic_from_stream(eligible_stream, reference_frames, count) def _stabilized_sampling_levels(reference_levels, count: int): """Schedule level 0 enters sample_step's fixed stabilization path.""" levels = reference_levels[:, None] if reference_levels.ndim == 1 else reference_levels return levels.new_zeros((int(count), levels.shape[-1]), dtype=torch.long) def _memory_noise_levels_for_streams(cfg, diffusion_model, query_noise_levels, stream_lengths, mode: str): is_training = str(mode) == "training" noise_cfg = _cfg_get(cfg, "memory_noise") if noise_cfg is None: if not is_training: return { key: _stabilized_sampling_levels(query_noise_levels, int(stream_lengths[key])) for key in _DEMEMWM_STREAM_KEYS } value = int(getattr(diffusion_model, "stabilization_level", 0)) return { key: torch.full((int(stream_lengths[key]), query_noise_levels.shape[-1]), value, device=query_noise_levels.device, dtype=torch.long) for key in _DEMEMWM_STREAM_KEYS } noise_mode = str(_cfg_get(noise_cfg, "mode", "random_cleaner_fraction")) if noise_mode not in {"random_cleaner_fraction", "independent"}: raise ValueError("memory_noise.mode must be random_cleaner_fraction or independent") noisy_memory = bool(_cfg_get(noise_cfg, "enabled", True)) and ( is_training or bool(_cfg_get(noise_cfg, "validation_noisy_memory", False)) ) query = query_noise_levels[:, None] if query_noise_levels.ndim == 1 else query_noise_levels query_bounds = query.to(dtype=torch.long).clamp_min(0).min(dim=0).values if not noisy_memory: return { key: _stabilized_sampling_levels(query, max(0, int(stream_lengths[key]))) for key in _DEMEMWM_STREAM_KEYS } if noise_mode == "independent": if is_training: max_level = int(getattr(diffusion_model, "timesteps", 0) or 0) - 1 else: max_level = int( getattr(diffusion_model, "sampling_timesteps", 0) or getattr(diffusion_model, "timesteps", 0) or 1 ) bounds = query_bounds.new_full(query_bounds.shape, max(0, max_level)) else: bounds = query_bounds defaults = {"anchor": 0.25, "dynamic": 0.50, "revisit": 0.25} levels_by_stream = {} for key in _DEMEMWM_STREAM_KEYS: count = max(0, int(stream_lengths[key])) if count == 0: levels_by_stream[key] = bounds.new_zeros((0, bounds.numel())) continue if noise_mode == "independent": max_levels = bounds if noisy_memory else bounds.new_zeros(bounds.shape) else: fraction = float(_cfg_get(noise_cfg, f"{key}_max_fraction", defaults[key])) if noisy_memory else 0.0 max_levels = torch.floor(bounds.to(dtype=torch.float32) * fraction).to(dtype=torch.long).clamp_min(0) if is_training and noisy_memory and noise_mode == "independent": levels = torch.randint( 0, int(bounds.max().item()) + 1, (count, bounds.numel()), device=bounds.device, dtype=torch.long, ) levels = torch.minimum(levels, max_levels.reshape(1, -1).expand_as(levels)) elif is_training and noisy_memory: levels = torch.floor(torch.rand((count, bounds.numel()), device=bounds.device) * (max_levels[None].float() + 1)).to(dtype=torch.long) else: levels = max_levels.reshape(1, -1).expand(count, -1).contiguous() levels_by_stream[key] = levels return levels_by_stream def _apply_memory_route_masks(frame_memory_masks, query_noise_levels, cfg, diffusion_model, mode: str): route_cfg = _cfg_get(cfg, "noise_route") routes = {key: str(_cfg_get(route_cfg, key, "all")) for key in _DEMEMWM_STREAM_KEYS} if all(route == "all" for route in routes.values()): return frame_memory_masks invalid = [route for route in routes.values() if route not in {"all", "high", "low"}] if invalid: raise ValueError(f"noise_route values must be all, high, or low, got {invalid}") levels = query_noise_levels[:, None] if query_noise_levels.ndim == 1 else query_noise_levels levels = levels.to(device=query_noise_levels.device, dtype=torch.long) batch_size = int(levels.shape[-1]) if str(mode) != "training": timesteps = int(getattr(diffusion_model, "timesteps", 0) or 0) sampling_timesteps = int(getattr(diffusion_model, "sampling_timesteps", 0) or 0) if timesteps > 1 and sampling_timesteps > 0: real_steps = torch.linspace(-1, timesteps - 1, steps=sampling_timesteps + 1, device=levels.device).long() levels = real_steps[levels] timesteps = int(getattr(diffusion_model, "timesteps", 0) or 0) high = torch.zeros((batch_size,), device=levels.device, dtype=torch.bool) if timesteps > 1: high = levels.to(dtype=torch.float32).mean(dim=0) >= (float(timesteps) / 2.0) masks_by_route = {"all": torch.ones_like(high), "high": high, "low": ~high} routed = dict(frame_memory_masks) for key in _DEMEMWM_STREAM_KEYS: if routed.get(key) is not None: routed[key] = routed[key].to(device=levels.device, dtype=torch.bool) & masks_by_route[routes[key]][:, None] return routed def _pack_active_inference_memory_streams( target_latents, target_conditions, target_poses, target_frame_indices, target_mask, stream_latents_by_key, stream_poses_by_key, stream_frame_indices_by_key, routed_stream_masks, memory_noise_levels_by_key, ): active_streams = [] pruned_streams = [] active_latents = [target_latents] active_poses = [target_poses] active_frame_indices = [target_frame_indices] active_frame_memory_segments = {"target": int(target_latents.shape[0])} active_frame_memory_masks = {"target": target_mask} active_memory_noise_levels = [] active_memory_length = 0 for key in _DEMEMWM_STREAM_KEYS: routed_mask = routed_stream_masks[key] if bool(routed_mask.any().item()): active_streams.append(key) active_latents.append(stream_latents_by_key[key]) active_poses.append(stream_poses_by_key[key]) active_frame_indices.append(stream_frame_indices_by_key[key]) active_frame_memory_segments[key] = int(stream_latents_by_key[key].shape[0]) active_frame_memory_masks[key] = routed_mask active_memory_noise_levels.append(memory_noise_levels_by_key[key]) active_memory_length += int(stream_latents_by_key[key].shape[0]) else: pruned_streams.append(key) active_frame_memory_segments[key] = 0 active_frame_memory_masks[key] = routed_mask[:, :0] active_packed_latents = torch.cat(active_latents, dim=0) active_packed_conditions = torch.cat( [ target_conditions, target_conditions.new_zeros((active_memory_length, target_conditions.shape[1], target_conditions.shape[-1])), ], dim=0, ) active_frame_memory_pose = torch.cat(active_poses, dim=0) active_frame_indices = torch.cat(active_frame_indices, dim=0) if active_memory_noise_levels: active_memory_noise_levels = torch.cat(active_memory_noise_levels, dim=0) else: active_memory_noise_levels = target_frame_indices.new_zeros((0, target_frame_indices.shape[1])) return ( active_packed_latents, active_packed_conditions, active_frame_memory_pose, active_frame_indices, active_frame_memory_segments, active_frame_memory_masks, active_memory_noise_levels, active_streams, pruned_streams, ) def random_transform(tensor): """ Apply the same random translation, rotation, and scaling to all frames in the batch. Args: tensor (torch.Tensor): Input tensor of shape (F, B, 3, H, W). Returns: torch.Tensor: Transformed tensor of shape (F, B, 3, H, W). """ if tensor.ndim != 5: raise ValueError("Input tensor must have shape (F, B, 3, H, W)") F, B, C, H, W = tensor.shape # Generate random transformation parameters max_translate = 0.2 # Translate up to 20% of width/height max_rotate = 30 # Rotate up to 30 degrees max_scale = 0.2 # Scale change by up to +/- 20% translate_x = random.uniform(-max_translate, max_translate) * W translate_y = random.uniform(-max_translate, max_translate) * H rotate_angle = random.uniform(-max_rotate, max_rotate) scale_factor = 1 + random.uniform(-max_scale, max_scale) # Apply the same transformation to all frames and batches tensor = tensor.reshape(F*B, C, H, W) transformed_tensor = TF.affine( tensor, angle=rotate_angle, translate=(translate_x, translate_y), scale=scale_factor, shear=(0, 0), interpolation=InterpolationMode.BILINEAR, fill=0 ) transformed_tensor = transformed_tensor.reshape(F, B, C, H, W) return transformed_tensor def save_tensor_as_png(tensor, file_path): """ Save a 3*H*W tensor as a PNG image. Args: tensor (torch.Tensor): Input tensor of shape (3, H, W). file_path (str): Path to save the PNG file. """ if tensor.ndim != 3 or tensor.shape[0] != 3: raise ValueError("Input tensor must have shape (3, H, W)") # Convert tensor to PIL Image image = TF.to_pil_image(tensor) # Save image image.save(file_path) class DeMemWMMinecraft(DiffusionForcingBase): """ DeMemWM video generation for MineCraft with frame memory. """ def __init__(self, cfg: DictConfig): """ Initialize the DeMemWMMinecraft class with the given configuration. Args: cfg (DictConfig): Configuration object. """ if _cfg_get(cfg, "memory_condition_length") is None: with open_dict(cfg): cfg.memory_condition_length = _derive_memory_condition_length(cfg) self.n_tokens = cfg.n_frames // cfg.frame_stack # number of max tokens for the model self.n_frames = cfg.n_frames if hasattr(cfg, "n_tokens"): self.n_tokens = cfg.n_tokens // cfg.frame_stack self.memory_condition_length = cfg.memory_condition_length self.pose_cond_dim = getattr(cfg, "pose_cond_dim", 5) self.use_plucker = getattr(cfg, "use_plucker", True) self.relative_embedding = getattr(cfg, "relative_embedding", True) self.state_embed_only_on_qk = getattr(cfg, "state_embed_only_on_qk", True) self.use_memory_attention = getattr(cfg, "use_memory_attention", True) self.add_timestamp_embedding = getattr(cfg, "add_timestamp_embedding", False) self.memory_attention_key_only_geometry = getattr(cfg, "memory_attention_key_only_geometry", True) self.ref_mode = getattr(cfg, "ref_mode", 'sequential') self.log_curve = getattr(cfg, "log_curve", False) self.focal_length = getattr(cfg, "focal_length", 0.35) self.log_video = cfg.log_video self.save_local = getattr(cfg, "save_local", True) self.local_save_dir = getattr(cfg, "local_save_dir", None) self.lpips_batch_size = getattr(cfg, "lpips_batch_size", 16) self.next_frame_length = getattr(cfg, "next_frame_length", 1) self.require_pose_prediction = getattr(cfg, "require_pose_prediction", False) self.metric_report_segment = max(0, int(getattr(cfg, "metric_report_segment", 0) or 0)) self.log_per_frame_metrics = bool(getattr(cfg, "log_per_frame_metrics", False)) self.log_revisit_metrics = bool(getattr(cfg, "log_revisit_metrics", False)) self.log_memory_selection_sheet = bool(getattr(cfg, "log_memory_selection_sheet", False)) self.memory_sheet_generated_frames = tuple( int(x) for x in getattr(cfg, "memory_sheet_generated_frames", (1, 25, 50, 100, 200, 350, 500)) ) self.revisit_metrics_min_gap = int(getattr(cfg, "revisit_metrics_min_gap", 8)) self.revisit_metrics_gap_bands = tuple(int(x) for x in getattr(cfg, "revisit_metrics_gap_bands", (32, 128))) self.revisit_metrics_self_crop_fraction = float(getattr(cfg, "revisit_metrics_self_crop_fraction", 0.5)) memory_selection_cfg = _cfg_get(cfg, "memory_selection", {}) self.revisit_metrics_fov_overlap_threshold = float( getattr( cfg, "revisit_metrics_fov_overlap_threshold", _cfg_get(memory_selection_cfg, "fov_overlap_threshold", 0.6), ) ) super().__init__(cfg) def _build_model(self): self.diffusion_model = Diffusion( # DeMemWM injects memory through explicit frame_memory_segments; # keep the legacy reference_length path disabled by default. reference_length=0, x_shape=self.x_stacked_shape, action_cond_dim=self.action_cond_dim, pose_cond_dim=self.pose_cond_dim, is_causal=self.causal, cfg=self.cfg.diffusion, is_dit=True, use_plucker=self.use_plucker, relative_embedding=self.relative_embedding, state_embed_only_on_qk=self.state_embed_only_on_qk, use_memory_attention=self.use_memory_attention, add_timestamp_embedding=self.add_timestamp_embedding, ref_mode=self.ref_mode, focal_length=self.focal_length, memory_attention_key_only_geometry=self.memory_attention_key_only_geometry, ) self.validation_lpips_model = LearnedPerceptualImagePatchSimilarity(sync_on_compute=False) vae = VAE_models["vit-l-20-shallow-encoder"]() self.vae = vae.eval() if self.require_pose_prediction: self.pose_prediction_model = PosePredictionNet() self._apply_trainability() def _global_step_for_trainability(self) -> int: try: trainer = self.trainer except RuntimeError: return 0 return int(getattr(trainer, "global_step", 0) or 0) def _apply_trainability(self) -> None: _apply_dememwm_trainability( self.diffusion_model, getattr(self, "vae", None), _trainability_cfg(self.cfg), self._global_step_for_trainability(), ) def configure_optimizers(self): trainability = _trainability_cfg(self.cfg) self._apply_trainability() param_groups = _dememwm_optimizer_parameter_groups( self.diffusion_model, getattr(self, "vae", None), trainability, self.cfg.lr, self._global_step_for_trainability(), ) if not param_groups: raise ValueError("DeMemWM trainability selected no optimizer parameters") return torch.optim.AdamW( param_groups, weight_decay=self.cfg.weight_decay, betas=self.cfg.optimizer_beta ) def on_train_batch_start(self, batch, batch_idx, dataloader_idx=0) -> None: self._apply_trainability() trainability = _trainability_cfg(self.cfg) for optimizer in getattr(getattr(self, "trainer", None), "optimizers", []) or []: _apply_dememwm_optimizer_group_lrs(optimizer, trainability, self.cfg.lr, self._global_step_for_trainability()) def _generate_noise_levels(self, xs: torch.Tensor, masks = None) -> torch.Tensor: """ Generate noise levels for training. """ num_frames, batch_size, *_ = xs.shape match self.cfg.noise_level: case "random_all": # entirely random noise levels noise_levels = torch.randint(0, self.timesteps, (num_frames, batch_size), device=xs.device) case "same": noise_levels = torch.randint(0, self.timesteps, (num_frames, batch_size), device=xs.device) noise_levels[1:] = noise_levels[0] if masks is not None: # for frames that are not available, treat as full noise discard = torch.all(~rearrange(masks.bool(), "(t fs) b -> t b fs", fs=self.frame_stack), -1) noise_levels = torch.where(discard, torch.full_like(noise_levels, self.timesteps - 1), noise_levels) return noise_levels def training_step(self, batch, batch_idx) -> STEP_OUTPUT: """ Perform a single training step. This function processes the input batch, encodes the input frames, generates noise levels, and computes the loss using the diffusion model. Args: batch: Input batch of data containing frames, conditions, poses, etc. batch_idx: Index of the current batch. Returns: dict: A dictionary containing the training loss. """ if not isinstance(batch, Mapping): raise TypeError( "DeMemWM training requires the latent dict batch contract " "from video_minecraft_dememwm_latent; raw tuple batches are unsupported." ) preprocessed = self._preprocess_batch(batch) if isinstance(preprocessed, Mapping): xs = preprocessed["latents"] target_length = preprocessed["target_length"] conditions = preprocessed["action_conditions"].to(device=xs.device) frame_indices = preprocessed["frame_indices"].to(device=xs.device) frame_memory_pose = preprocessed["frame_memory_pose"].to(device=xs.device, dtype=xs.dtype) image_hw = preprocessed["image_hw"].to(device=xs.device) frame_memory_masks = { key: mask.to(device=xs.device) for key, mask in preprocessed["memory_masks"].items() } target_noise_levels = self._generate_noise_levels(xs[:target_length]) cfg = getattr(self, "cfg", None) frame_memory_masks = _apply_memory_route_masks( frame_memory_masks, target_noise_levels, cfg, self.diffusion_model, mode="training", ) memory_noise_levels = _memory_noise_levels_for_streams( cfg, self.diffusion_model, target_noise_levels, preprocessed["stream_lengths"], mode="training", ) noise_levels = torch.cat( [target_noise_levels, *[memory_noise_levels[key] for key in _DEMEMWM_STREAM_KEYS]], dim=0, ) _, loss = self.diffusion_model( xs, conditions, None, noise_levels=noise_levels, reference_length=0, frame_idx=frame_indices, frame_memory_segments=preprocessed["memory_segments"], frame_memory_masks=frame_memory_masks, frame_memory_pose=frame_memory_pose, image_hw=image_hw, ) loss = loss[:target_length] target_mask = frame_memory_masks.get("target") if target_mask is not None: target_mask = rearrange(target_mask.to(dtype=loss.dtype), "b t -> t b") target_mask = target_mask.view(*target_mask.shape, *((1,) * (loss.ndim - 2))) loss = (loss * target_mask).sum() / target_mask.expand_as(loss).sum().clamp_min(1.0) else: loss = self.reweight_loss(loss, None) if batch_idx % 20 == 0: self.log("training/loss", loss.detach(), prog_bar=True, sync_dist=True) return {"loss": loss} raise TypeError( "DeMemWM _preprocess_batch must return the latent dict contract; " "legacy raw training tuples are unsupported." ) def on_validation_epoch_start(self) -> None: self._reset_metric_accumulators() def on_test_epoch_start(self) -> None: self._reset_metric_accumulators() def on_validation_epoch_end(self) -> None: self._on_eval_epoch_end() def on_test_epoch_end(self) -> None: self._on_eval_epoch_end() def _reset_metric_accumulators(self) -> None: self._metric_device = next(self.validation_lpips_model.parameters()).device self._mse_sum = torch.tensor(0.0, device=self._metric_device) self._mse_count = torch.tensor(0.0, device=self._metric_device) self._psnr_sum = torch.tensor(0.0, device=self._metric_device) self._psnr_count = torch.tensor(0.0, device=self._metric_device) self._lpips_sum = torch.tensor(0.0, device=self._metric_device) self._lpips_count = torch.tensor(0.0, device=self._metric_device) self._frame_metrics_synced = False self._segment_metrics_synced = False if self.log_per_frame_metrics: self._frame_metric_eval_start = None for attr in ("mse_sum", "mse_count", "psnr_sum", "psnr_count", "lpips_sum", "lpips_count"): setattr(self, f"_frame_{attr}", torch.empty(0, device=self._metric_device)) if self.metric_report_segment > 0: self._segment_metric_eval_start = None self._segment_metric_frames = 0 for attr in ("mse_sum", "mse_count", "psnr_sum", "psnr_count", "lpips_sum", "lpips_count"): setattr(self, f"_segment_{attr}", torch.empty(0, device=self._metric_device)) if self.log_revisit_metrics: self._reset_revisit_metric_accumulators() def _eval_artifact_dir(self) -> Path: base = Path(self.local_save_dir) if self.local_save_dir is not None else Path.cwd() / "eval_artifacts" return base / "metrics" def _wandb_run_files_dir(self) -> Path | None: logger = getattr(self, "logger", None) experiment = None if logger is None else getattr(logger, "experiment", None) run_dir = None if experiment is None else getattr(experiment, "dir", None) return None if run_dir is None else Path(run_dir) def _segment_metric_artifact_dir(self) -> Path: run_dir = self._wandb_run_files_dir() if run_dir is not None: return run_dir / "metrics" return self._eval_artifact_dir() def _memory_sheet_artifact_dir(self) -> Path: base = Path(self.local_save_dir) if self.local_save_dir is not None else Path.cwd() / "eval_artifacts" return base / "memory_selection_sheets" def _memory_sheet_target_indices(self, eval_start: int, target_length: int) -> set[int]: return { int(eval_start) + int(frame) - 1 for frame in self.memory_sheet_generated_frames if int(frame) > 0 and int(eval_start) + int(frame) - 1 < int(target_length) } def _save_memory_selection_sheets( self, xs_pred: torch.Tensor, xs_gt: torch.Tensor, frame_indices: torch.Tensor, records: list[list[dict]], eval_start: int, batch_idx: int, ) -> None: if not self.log_memory_selection_sheet or not self._is_global_rank_zero(): return output_dir = self._memory_sheet_artifact_dir() output_dir.mkdir(parents=True, exist_ok=True) max_slots = {"anchor": 2, "dynamic": 4, "revisit": 2} tile_w, tile_h = 192, 108 label_h = 34 gap = 6 def decode_indices(source: torch.Tensor, batch_i: int, indices: list[int]) -> dict[int, Image.Image]: if not indices: return {} unique = sorted(set(int(i) for i in indices)) latents = torch.stack([source[i, batch_i] for i in unique], dim=0).unsqueeze(1) with torch.no_grad(): decoded = self.decode(latents.to(source.device))[:, 0].detach().cpu().float().clamp(0.0, 1.0) return {idx: TF.to_pil_image(decoded[row]).resize((tile_w, tile_h), Image.BILINEAR) for row, idx in enumerate(unique)} def make_tile(image: Image.Image | None, label: str, fill=(236, 236, 236)) -> Image.Image: tile = Image.new("RGB", (tile_w, tile_h + label_h), fill) if image is not None: tile.paste(image.convert("RGB"), (0, label_h)) draw = ImageDraw.Draw(tile) draw.rectangle((0, 0, tile_w, label_h), fill=(18, 18, 18)) for row, line in enumerate(label.split("\n")[:2]): draw.text((5, 4 + row * 14), line, fill=(255, 255, 255)) return tile summary_rows = [] for batch_i, batch_records in enumerate(records): if not batch_records: continue target_indices = [int(row["target_index"]) for row in batch_records] memory_gt = [] memory_pred = [] for row in batch_records: for stream in _DEMEMWM_STREAM_KEYS: for source_index in row["streams"][stream]: (memory_gt if source_index < eval_start else memory_pred).append(int(source_index)) target_gt_images = decode_indices(xs_gt, batch_i, target_indices) target_pred_images = decode_indices(xs_pred, batch_i, target_indices) memory_images = {} memory_images.update(decode_indices(xs_gt, batch_i, memory_gt)) memory_images.update(decode_indices(xs_pred, batch_i, memory_pred)) columns = [("target_gt", -1), ("target_pred", -1)] for stream in _DEMEMWM_STREAM_KEYS: columns.extend((stream, slot) for slot in range(max_slots[stream])) sheet = Image.new( "RGB", (len(columns) * tile_w + (len(columns) - 1) * gap, len(batch_records) * (tile_h + label_h) + (len(batch_records) - 1) * gap), (245, 245, 245), ) for record_row, row in enumerate(batch_records): y = record_row * (tile_h + label_h + gap) target_index = int(row["target_index"]) generated_frame = target_index - int(eval_start) + 1 raw_target = int(row["target_frame"]) for col, (stream, slot) in enumerate(columns): x = col * (tile_w + gap) if stream == "target_gt": image = target_gt_images.get(target_index) label = f"gen{generated_frame:03d} raw{raw_target}\ntarget GT" elif stream == "target_pred": image = target_pred_images.get(target_index) label = f"gen{generated_frame:03d} raw{raw_target}\ntarget pred" else: stream_indices = row["streams"][stream] source_index = int(stream_indices[slot]) if slot < len(stream_indices) else -1 image = memory_images.get(source_index) if source_index >= 0: raw_source = int(frame_indices[source_index, batch_i].detach().cpu().item()) source_type = "ctx" if source_index < eval_start else "pred" label = f"{stream}{slot} seq{source_index} raw{raw_source}\n{source_type} for gen{generated_frame:03d}" summary_rows.append( { "batch": batch_idx, "sample": batch_i, "generated_frame": generated_frame, "target_index": target_index, "target_raw_frame": raw_target, "stream": stream, "slot": slot, "source_index": source_index, "source_raw_frame": raw_source, "source_type": source_type, } ) else: label = f"{stream}{slot}\nempty" sheet.paste(make_tile(image, label), (x, y)) sheet.save(output_dir / f"batch{batch_idx:06d}_sample{batch_i:02d}_memory_streams.jpg", quality=95) if summary_rows: csv_path = output_dir / f"batch{batch_idx:06d}_memory_streams.csv" with csv_path.open("w", newline="", encoding="utf-8") as handle: writer = csv.DictWriter(handle, fieldnames=list(summary_rows[0].keys())) writer.writeheader() writer.writerows(summary_rows) def _is_global_rank_zero(self) -> bool: return not (dist.is_available() and dist.is_initialized()) or dist.get_rank() == 0 def _mean_or_nan_tensor(self, total: torch.Tensor, count: torch.Tensor) -> torch.Tensor: return torch.where(count > 0, total / count.clamp_min(1.0), torch.full_like(total, float("nan"))) def _pad_frame_metric_accumulators(self, length: int) -> None: current = int(self._frame_mse_sum.numel()) if current >= int(length): return pad = int(length) - current for attr in ("mse_sum", "mse_count", "psnr_sum", "psnr_count", "lpips_sum", "lpips_count"): name = f"_frame_{attr}" setattr(self, name, torch.cat([getattr(self, name), torch.zeros(pad, device=self._metric_device)])) def _sync_frame_metric_accumulators(self) -> None: if getattr(self, "_frame_metrics_synced", False): return if not (dist.is_available() and dist.is_initialized()): self._frame_metrics_synced = True return length = torch.tensor(int(self._frame_mse_sum.numel()), device=self._metric_device, dtype=torch.long) dist.all_reduce(length, op=dist.ReduceOp.MAX) self._pad_frame_metric_accumulators(int(length.item())) for attr in ("mse_sum", "mse_count", "psnr_sum", "psnr_count", "lpips_sum", "lpips_count"): dist.all_reduce(getattr(self, f"_frame_{attr}"), op=dist.ReduceOp.SUM) self._frame_metrics_synced = True def _update_per_frame_metric_accumulators( self, xs_pred: torch.Tensor, xs_gt: torch.Tensor, valid_mask: torch.Tensor | None, eval_start: int, ) -> None: if not self.log_per_frame_metrics or xs_pred.numel() == 0: return if not hasattr(self, "_frame_mse_sum"): self._reset_metric_accumulators() self._frame_metrics_synced = False frames = int(xs_pred.shape[0]) self._pad_frame_metric_accumulators(frames) if self._frame_metric_eval_start is None: self._frame_metric_eval_start = int(eval_start) pred = torch.clamp(xs_pred.to(self._metric_device).float(), 0.0, 1.0) gt = torch.clamp(xs_gt.to(self._metric_device).float(), 0.0, 1.0) mask = torch.ones(pred.shape[:2], device=self._metric_device, dtype=torch.bool) if valid_mask is not None: mask = valid_mask.to(device=self._metric_device, dtype=torch.bool) frame_mse = (pred - gt).square().flatten(start_dim=2).mean(dim=2) frame_psnr = 10.0 * torch.log10(1.0 / frame_mse.clamp_min(torch.finfo(frame_mse.dtype).eps)) counts = mask.sum(dim=1).to(dtype=frame_mse.dtype) mask_f = mask.to(dtype=frame_mse.dtype) self._frame_mse_sum[:frames] += (frame_mse * mask_f).sum(dim=1) self._frame_mse_count[:frames] += counts self._frame_psnr_sum[:frames] += (frame_psnr * mask_f).sum(dim=1) self._frame_psnr_count[:frames] += counts with torch.no_grad(): for frame_idx in range(frames): frame_mask = mask[frame_idx] count = int(frame_mask.sum().item()) if count == 0: continue self.validation_lpips_model.reset() self.validation_lpips_model.update(pred[frame_idx, frame_mask], gt[frame_idx, frame_mask]) value = self.validation_lpips_model.compute().detach().to(self._metric_device) self.validation_lpips_model.reset() self._frame_lpips_sum[frame_idx] += value * count self._frame_lpips_count[frame_idx] += count def _pad_segment_metric_accumulators(self, length: int) -> None: current = int(self._segment_mse_sum.numel()) if current >= int(length): return pad = int(length) - current for attr in ("mse_sum", "mse_count", "psnr_sum", "psnr_count", "lpips_sum", "lpips_count"): name = f"_segment_{attr}" setattr(self, name, torch.cat([getattr(self, name), torch.zeros(pad, device=self._metric_device)])) def _sync_segment_metric_accumulators(self) -> None: if getattr(self, "_segment_metrics_synced", False): return if not (dist.is_available() and dist.is_initialized()): self._segment_metrics_synced = True return length = torch.tensor(int(self._segment_mse_sum.numel()), device=self._metric_device, dtype=torch.long) frames = torch.tensor(int(getattr(self, "_segment_metric_frames", 0)), device=self._metric_device, dtype=torch.long) dist.all_reduce(length, op=dist.ReduceOp.MAX) dist.all_reduce(frames, op=dist.ReduceOp.MAX) self._segment_metric_frames = int(frames.item()) self._pad_segment_metric_accumulators(int(length.item())) for attr in ("mse_sum", "mse_count", "psnr_sum", "psnr_count", "lpips_sum", "lpips_count"): dist.all_reduce(getattr(self, f"_segment_{attr}"), op=dist.ReduceOp.SUM) self._segment_metrics_synced = True def _update_segment_metric_accumulators( self, xs_pred: torch.Tensor, xs_gt: torch.Tensor, valid_mask: torch.Tensor | None, eval_start: int, ) -> None: if self.metric_report_segment <= 0 or xs_pred.numel() == 0: return if not hasattr(self, "_segment_mse_sum"): self._reset_metric_accumulators() self._segment_metrics_synced = False frames = int(xs_pred.shape[0]) segment = int(self.metric_report_segment) num_segments = (frames + segment - 1) // segment self._pad_segment_metric_accumulators(num_segments) self._segment_metric_frames = max(int(getattr(self, "_segment_metric_frames", 0)), frames) if self._segment_metric_eval_start is None: self._segment_metric_eval_start = int(eval_start) pred = torch.clamp(xs_pred.to(self._metric_device).float(), 0.0, 1.0) gt = torch.clamp(xs_gt.to(self._metric_device).float(), 0.0, 1.0) mask = torch.ones(pred.shape[:2], device=self._metric_device, dtype=torch.bool) if valid_mask is not None: mask = valid_mask.to(device=self._metric_device, dtype=torch.bool) frame_mse = (pred - gt).square().flatten(start_dim=2).mean(dim=2) frame_psnr = 10.0 * torch.log10(1.0 / frame_mse.clamp_min(torch.finfo(frame_mse.dtype).eps)) lpips_batch_size = max(1, int(self.lpips_batch_size)) for segment_idx in range(num_segments): start = segment_idx * segment end = min(start + segment, frames) segment_mask = mask[start:end] count = int(segment_mask.sum().item()) if count == 0: continue mask_f = segment_mask.to(dtype=frame_mse.dtype) mse_sum = (frame_mse[start:end] * mask_f).sum() psnr_sum = (frame_psnr[start:end] * mask_f).sum() count_t = torch.tensor(float(count), device=self._metric_device) pred_segment = pred[start:end][segment_mask] gt_segment = gt[start:end][segment_mask] self.validation_lpips_model.reset() for batch_start in range(0, count, lpips_batch_size): batch_end = min(batch_start + lpips_batch_size, count) self.validation_lpips_model.update(pred_segment[batch_start:batch_end], gt_segment[batch_start:batch_end]) lpips = self.validation_lpips_model.compute().detach().to(self._metric_device) self.validation_lpips_model.reset() self._segment_mse_sum[segment_idx] += mse_sum self._segment_mse_count[segment_idx] += count_t self._segment_psnr_sum[segment_idx] += psnr_sum self._segment_psnr_count[segment_idx] += count_t self._segment_lpips_sum[segment_idx] += lpips * count_t self._segment_lpips_count[segment_idx] += count_t # Segment mode derives global metrics from this same pass to avoid # running LPIPS once for the whole rollout and again per segment. self._mse_sum += mse_sum self._mse_count += count_t self._psnr_sum += psnr_sum self._psnr_count += count_t self._lpips_sum += lpips * count_t self._lpips_count += count_t def _log_per_frame_metrics(self) -> None: if not hasattr(self, "_frame_mse_sum"): return self._sync_frame_metric_accumulators() valid = self._frame_mse_count > 0 if not bool(valid.any()): return eval_start = 0 if self._frame_metric_eval_start is None else int(self._frame_metric_eval_start) rows = [] log_dict = {} for frame_idx in torch.nonzero(valid, as_tuple=False).flatten().tolist(): mse = self._frame_mse_sum[frame_idx] / self._frame_mse_count[frame_idx].clamp_min(1.0) psnr = self._frame_psnr_sum[frame_idx] / self._frame_psnr_count[frame_idx].clamp_min(1.0) lpips = self._frame_lpips_sum[frame_idx] / self._frame_lpips_count[frame_idx].clamp_min(1.0) generated_frame = int(frame_idx) + 1 log_dict[f"per_frame/mse_{generated_frame:04d}"] = mse log_dict[f"per_frame/psnr_{generated_frame:04d}"] = psnr log_dict[f"per_frame/lpips_{generated_frame:04d}"] = lpips rows.append( { "generated_frame": generated_frame, "absolute_frame": eval_start + int(frame_idx), "mse": float(mse.detach().cpu().item()), "psnr": float(psnr.detach().cpu().item()), "lpips": float(lpips.detach().cpu().item()), "count": float(self._frame_mse_count[frame_idx].detach().cpu().item()), } ) self.log_dict(log_dict, sync_dist=False) if self._is_global_rank_zero(): output_path = self._eval_artifact_dir() / f"per_frame_metrics_step{int(getattr(self, 'global_step', 0)):08d}.csv" output_path.parent.mkdir(parents=True, exist_ok=True) with output_path.open("w", newline="", encoding="utf-8") as handle: writer = csv.DictWriter(handle, fieldnames=list(rows[0].keys())) writer.writeheader() writer.writerows(rows) def _log_metric_segments(self) -> None: if self.metric_report_segment <= 0 or not hasattr(self, "_segment_mse_sum"): return self._sync_segment_metric_accumulators() valid = self._segment_mse_count > 0 if not bool(valid.any()): return segment = int(self.metric_report_segment) eval_start = 0 if self._segment_metric_eval_start is None else int(self._segment_metric_eval_start) rows = [] log_dict = {} total_generated = int(getattr(self, "_segment_metric_frames", self._segment_mse_count.numel() * segment)) for segment_idx in torch.nonzero(valid, as_tuple=False).flatten().tolist(): start = int(segment_idx) * segment end = min(start + segment, total_generated) label = f"{start:04d}_{end:04d}" mse = self._mean_or_nan_tensor(self._segment_mse_sum[segment_idx], self._segment_mse_count[segment_idx]) psnr = self._mean_or_nan_tensor(self._segment_psnr_sum[segment_idx], self._segment_psnr_count[segment_idx]) lpips = self._mean_or_nan_tensor(self._segment_lpips_sum[segment_idx], self._segment_lpips_count[segment_idx]) generation_length = end - start log_dict[f"segment/mse_{label}"] = mse log_dict[f"segment/psnr_{label}"] = psnr log_dict[f"segment/lpips_{label}"] = lpips log_dict[f"segment/generation_length_{label}"] = torch.tensor(float(generation_length), device=self._metric_device) rows.append( { "generation_start": start, "generation_end": end, "generation_length": generation_length, "absolute_start": eval_start + start, "absolute_end": eval_start + end, "mse": float(mse.detach().cpu().item()), "psnr": float(psnr.detach().cpu().item()), "lpips": float(lpips.detach().cpu().item()), "count": float(self._segment_mse_count[segment_idx].detach().cpu().item()), } ) self.log_dict(log_dict, sync_dist=False) if self._is_global_rank_zero(): output_path = self._segment_metric_artifact_dir() / f"segment_metrics_step{int(getattr(self, 'global_step', 0)):08d}.csv" output_path.parent.mkdir(parents=True, exist_ok=True) with output_path.open("w", newline="", encoding="utf-8") as handle: writer = csv.DictWriter(handle, fieldnames=list(rows[0].keys())) writer.writeheader() writer.writerows(rows) logger = getattr(self, "logger", None) if logger is not None: import wandb columns = list(rows[0].keys()) experiment = logger.experiment experiment.log( { "segment_metrics/table": wandb.Table( columns=columns, data=[[row[column] for column in columns] for row in rows], ), "trainer/global_step": int(getattr(self, "global_step", 0)), } ) save_fn = getattr(experiment, "save", None) if callable(save_fn): save_fn(str(output_path), policy="now") def _revisit_gap_labels(self) -> list[tuple[str, int, int | None]]: edges = [self.revisit_metrics_min_gap] edges.extend(edge for edge in sorted(set(self.revisit_metrics_gap_bands)) if edge > self.revisit_metrics_min_gap) edges.append(None) return [ (f"{int(lo)}_{'inf' if hi is None else int(hi)}", int(lo), None if hi is None else int(hi)) for lo, hi in zip(edges[:-1], edges[1:]) ] def _reset_revisit_metric_accumulators(self) -> None: names = ( "num_pairs", "num_revisit_targets", "num_revisit_frames", "num_non_revisit_frames", "revisit_lpips_sum", "revisit_lpips_count", "non_revisit_lpips_sum", "non_revisit_lpips_count", "revisit_psnr_sum", "revisit_psnr_count", "non_revisit_psnr_sum", "non_revisit_psnr_count", "self_lpips_sum", "self_lpips_count", "self_psnr_sum", "self_psnr_count", ) for name in names: setattr(self, f"_revisit_{name}", torch.tensor(0.0, device=self._metric_device)) self._revisit_gap_stats = { label: { "count": torch.tensor(0.0, device=self._metric_device), "self_lpips_sum": torch.tensor(0.0, device=self._metric_device), "self_lpips_count": torch.tensor(0.0, device=self._metric_device), "self_psnr_sum": torch.tensor(0.0, device=self._metric_device), "self_psnr_count": torch.tensor(0.0, device=self._metric_device), "revisit_lpips_sum": torch.tensor(0.0, device=self._metric_device), "revisit_lpips_count": torch.tensor(0.0, device=self._metric_device), } for label, _, _ in self._revisit_gap_labels() } def _add_revisit_value(self, sum_attr: str, count_attr: str, value) -> None: try: value = float(value) except (TypeError, ValueError): return if not np.isfinite(value): return getattr(self, sum_attr).add_(value) getattr(self, count_attr).add_(1.0) def _add_revisit_value_for_dict(self, stats: dict[str, torch.Tensor], name: str, value) -> None: try: value = float(value) except (TypeError, ValueError): return if not np.isfinite(value): return stats[f"{name}_sum"] += value stats[f"{name}_count"] += 1.0 def _mine_revisit_pairs_with_validation_selector( self, poses: torch.Tensor, frame_indices: torch.Tensor, memory_selection_cfg, eval_start: int, split: str, clip_id: str, ) -> list[RevisitPair]: poses_np = poses.detach().cpu().numpy() frame_indices_cpu = frame_indices.detach().cpu() excluded = np.empty((0,), dtype=np.int64) pairs: list[RevisitPair] = [] for target_index in range(int(eval_start), int(poses.shape[0])): selected = _select_revisit( poses_np, np.asarray([target_index], dtype=np.int64), memory_selection_cfg, 1, excluded, split, ) if len(selected) == 0: continue source_index = int(selected[0]) source_frame = int(frame_indices_cpu[source_index].item()) target_frame = int(frame_indices_cpu[target_index].item()) if source_index >= target_index: continue source_pose = poses[source_index].detach().float().cpu() target_pose = poses[target_index].detach().float().cpu() position_distance = float(torch.linalg.vector_norm(target_pose[:3] - source_pose[:3]).item()) yaw_delta = float(abs(((target_pose[4] - source_pose[4] + 180.0) % 360.0) - 180.0).item()) pitch_delta = float(abs(((target_pose[3] - source_pose[3] + 180.0) % 360.0) - 180.0).item()) pairs.append( RevisitPair( clip_id=str(clip_id), source_index=source_index, target_index=target_index, source_frame=source_frame, target_frame=target_frame, gap=int(target_frame - source_frame), fov_overlap=1.0, plucker_overlap=0.0, position_distance=position_distance, yaw_delta_deg=yaw_delta, pitch_delta_deg=pitch_delta, positional=bool(position_distance <= 2.0), ) ) return pairs def _update_revisit_metric_accumulators( self, xs_pred: torch.Tensor, xs_gt: torch.Tensor, target_poses: torch.Tensor, target_frame_indices: torch.Tensor, valid_mask: torch.Tensor | None, eval_start: int, batch_idx: int, namespace: str, ) -> None: if not self.log_revisit_metrics or xs_pred.numel() == 0: return if not hasattr(self, "_revisit_num_pairs"): self._reset_revisit_metric_accumulators() valid_mask_cpu = None if valid_mask is None else valid_mask.detach().cpu() batch_size = int(xs_pred.shape[1]) for batch_i in range(batch_size): clip_id = f"batch{batch_idx:06d}_sample{batch_i:02d}" poses = target_poses[:, batch_i, :5].detach() frame_indices = target_frame_indices[:, batch_i].detach() pairs = self._mine_revisit_pairs_with_validation_selector( poses, frame_indices, _cfg_get(self.cfg, "memory_selection", {}), eval_start, namespace, clip_id, ) self._revisit_num_pairs += float(len(pairs)) self._revisit_num_revisit_targets += float(len({int(pair.target_index) for pair in pairs})) rows = compute_frame_metric_rows( branch="validation", clip_id=clip_id, pred=xs_pred[:, batch_i], gt=xs_gt[:, batch_i], poses=poses, pairs=pairs, context_frames=int(eval_start), lpips_model=self.validation_lpips_model, self_crop_fraction=self.revisit_metrics_self_crop_fraction, compute_plucker_similarity=False, ) for row in rows: output_index = int(row["output_index"]) if valid_mask_cpu is not None and not bool(valid_mask_cpu[output_index, batch_i].item()): continue if bool(row["is_revisit"]): self._revisit_num_revisit_frames += 1.0 self._add_revisit_value("_revisit_revisit_lpips_sum", "_revisit_revisit_lpips_count", row["gt_lpips"]) self._add_revisit_value("_revisit_revisit_psnr_sum", "_revisit_revisit_psnr_count", row["gt_psnr"]) self._add_revisit_value("_revisit_self_lpips_sum", "_revisit_self_lpips_count", row["self_lpips"]) self._add_revisit_value("_revisit_self_psnr_sum", "_revisit_self_psnr_count", row["self_psnr"]) try: gap = int(row["gap"]) except (TypeError, ValueError): gap = None if gap is not None: for label, lo, hi in self._revisit_gap_labels(): if gap >= lo and (hi is None or gap < hi): stats = self._revisit_gap_stats[label] stats["count"] += 1.0 self._add_revisit_value_for_dict(stats, "revisit_lpips", row["gt_lpips"]) self._add_revisit_value_for_dict(stats, "self_lpips", row["self_lpips"]) self._add_revisit_value_for_dict(stats, "self_psnr", row["self_psnr"]) break else: self._revisit_num_non_revisit_frames += 1.0 self._add_revisit_value("_revisit_non_revisit_lpips_sum", "_revisit_non_revisit_lpips_count", row["gt_lpips"]) self._add_revisit_value("_revisit_non_revisit_psnr_sum", "_revisit_non_revisit_psnr_count", row["gt_psnr"]) def _sync_revisit_metric_accumulators(self) -> None: if not (dist.is_available() and dist.is_initialized()): return attrs = ( "num_pairs", "num_revisit_targets", "num_revisit_frames", "num_non_revisit_frames", "revisit_lpips_sum", "revisit_lpips_count", "non_revisit_lpips_sum", "non_revisit_lpips_count", "revisit_psnr_sum", "revisit_psnr_count", "non_revisit_psnr_sum", "non_revisit_psnr_count", "self_lpips_sum", "self_lpips_count", "self_psnr_sum", "self_psnr_count", ) for attr in attrs: dist.all_reduce(getattr(self, f"_revisit_{attr}"), op=dist.ReduceOp.SUM) for stats in self._revisit_gap_stats.values(): for value in stats.values(): dist.all_reduce(value, op=dist.ReduceOp.SUM) def _log_revisit_metrics(self) -> None: if not hasattr(self, "_revisit_num_pairs"): return self._sync_revisit_metric_accumulators() revisit_lpips = self._mean_or_nan_tensor(self._revisit_revisit_lpips_sum, self._revisit_revisit_lpips_count) non_revisit_lpips = self._mean_or_nan_tensor(self._revisit_non_revisit_lpips_sum, self._revisit_non_revisit_lpips_count) log_dict = { "revisit_metrics/num_pairs": self._revisit_num_pairs, "revisit_metrics/num_revisit_targets": self._revisit_num_revisit_targets, "revisit_metrics/num_revisit_frames": self._revisit_num_revisit_frames, "revisit_metrics/num_non_revisit_frames": self._revisit_num_non_revisit_frames, "revisit_metrics/revisit_lpips": revisit_lpips, "revisit_metrics/non_revisit_lpips": non_revisit_lpips, "revisit_metrics/revisit_psnr": self._mean_or_nan_tensor(self._revisit_revisit_psnr_sum, self._revisit_revisit_psnr_count), "revisit_metrics/non_revisit_psnr": self._mean_or_nan_tensor(self._revisit_non_revisit_psnr_sum, self._revisit_non_revisit_psnr_count), "revisit_metrics/self_consistency_lpips": self._mean_or_nan_tensor(self._revisit_self_lpips_sum, self._revisit_self_lpips_count), "revisit_metrics/self_consistency_psnr": self._mean_or_nan_tensor(self._revisit_self_psnr_sum, self._revisit_self_psnr_count), "revisit_metrics/revisit_minus_non_revisit_lpips": revisit_lpips - non_revisit_lpips, } for label, stats in self._revisit_gap_stats.items(): log_dict[f"revisit_metrics/num_gap_{label}"] = stats["count"] log_dict[f"revisit_metrics/self_consistency_lpips_gap_{label}"] = self._mean_or_nan_tensor(stats["self_lpips_sum"], stats["self_lpips_count"]) log_dict[f"revisit_metrics/self_consistency_psnr_gap_{label}"] = self._mean_or_nan_tensor(stats["self_psnr_sum"], stats["self_psnr_count"]) log_dict[f"revisit_metrics/revisit_lpips_gap_{label}"] = self._mean_or_nan_tensor(stats["revisit_lpips_sum"], stats["revisit_lpips_count"]) self.log_dict(log_dict, sync_dist=False) if self._is_global_rank_zero(): rows = [{"metric": key, "value": float(value.detach().cpu().item())} for key, value in log_dict.items()] output_path = self._eval_artifact_dir() / f"revisit_metrics_step{int(getattr(self, 'global_step', 0)):08d}.csv" output_path.parent.mkdir(parents=True, exist_ok=True) with output_path.open("w", newline="", encoding="utf-8") as handle: writer = csv.DictWriter(handle, fieldnames=["metric", "value"]) writer.writeheader() writer.writerows(rows) def _update_metric_accumulators( self, xs_pred: torch.Tensor, xs_gt: torch.Tensor, valid_mask: torch.Tensor | None = None, eval_start: int = 0, ) -> None: if not hasattr(self, "_metric_device"): self._reset_metric_accumulators() if self.metric_report_segment > 0: self._update_segment_metric_accumulators(xs_pred, xs_gt, valid_mask, eval_start) return xs_pred_device = xs_pred.to(self._metric_device) xs_gt_device = xs_gt.to(self._metric_device) if valid_mask is not None: valid_mask = valid_mask.to(device=self._metric_device, dtype=torch.bool) if not bool(valid_mask.any()): return xs_pred_device = xs_pred_device[valid_mask].unsqueeze(1) xs_gt_device = xs_gt_device[valid_mask].unsqueeze(1) metric_dict = get_validation_metrics_for_videos( xs_pred_device, xs_gt_device, lpips_model=self.validation_lpips_model, lpips_batch_size=self.lpips_batch_size, ) mse_count = torch.tensor(float(xs_pred_device.numel()), device=self._metric_device) psnr_count = torch.tensor(float(xs_pred_device.shape[1]), device=self._metric_device) lpips_count = torch.tensor(float(xs_pred_device.shape[0] * xs_pred_device.shape[1]), device=self._metric_device) self._mse_sum += metric_dict["mse"].detach() * mse_count self._mse_count += mse_count self._psnr_sum += metric_dict["psnr"].detach() * psnr_count self._psnr_count += psnr_count self._lpips_sum += torch.tensor(float(metric_dict["lpips"]), device=self._metric_device) * lpips_count self._lpips_count += lpips_count del xs_pred_device, xs_gt_device def _on_eval_epoch_end(self) -> None: if not hasattr(self, "_metric_device"): return if dist.is_available() and dist.is_initialized(): for tensor in ( self._mse_sum, self._mse_count, self._psnr_sum, self._psnr_count, self._lpips_sum, self._lpips_count, ): dist.all_reduce(tensor, op=dist.ReduceOp.SUM) if self._mse_count.item() > 0: self.log_dict( { "mse": self._mse_sum / self._mse_count.clamp_min(1.0), "psnr": self._psnr_sum / self._psnr_count.clamp_min(1.0), "lpips": self._lpips_sum / self._lpips_count.clamp_min(1.0), }, sync_dist=False, ) if self.log_per_frame_metrics: self._log_per_frame_metrics() if self.metric_report_segment > 0: self._log_metric_segments() if self.log_revisit_metrics: self._log_revisit_metrics() self.validation_step_outputs.clear() def _preprocess_batch(self, batch): if not isinstance(batch, Mapping): raise TypeError( "DeMemWM requires the latent dict batch contract " "from video_minecraft_dememwm_latent; raw tuple batches are unsupported." ) return _preprocess_dememwm_latent_batch(batch) def decode(self, x): total_frames = x.shape[0] scaling_factor = 0.07843137255 x = rearrange(x, "t b c h w -> (t b) (h w) c") with torch.no_grad(): x = (self.vae.decode(x / scaling_factor) + 1) / 2 x = rearrange(x, "(t b) c h w-> t b c h w", t=total_frames) return x def validation_step(self, batch, batch_idx, namespace="validation") -> STEP_OUTPUT: """ Perform a single validation step. DeMemWM validation/test uses latent dict batches and builds typed online memory bundles for autoregressive rollout. Args: batch: Input batch of data containing frames, conditions, poses, etc. batch_idx: Index of the current batch. namespace: Namespace for logging (default: "validation"). Returns: torch.Tensor: Target-only latent MSE for the evaluated rollout frames. """ if not isinstance(batch, Mapping): raise TypeError( f"DeMemWM {namespace} requires the latent dict batch contract " "from video_minecraft_dememwm_latent; raw tuple batches are unsupported." ) preprocessed = self._preprocess_batch(batch) if isinstance(preprocessed, Mapping): target_tensors = preprocessed["target_tensors"] xs = target_tensors["latents"] target_length = preprocessed["target_length"] batch_size = xs.shape[1] image_hw = preprocessed["image_hw"].to(device=xs.device) target_mask = preprocessed["memory_masks"].get("target") target_mask = None if target_mask is None else target_mask.to(device=xs.device) target_poses = target_tensors["poses"].to(device=xs.device, dtype=xs.dtype) target_frame_indices = target_tensors["frame_indices"].to(device=xs.device) target_actions = target_tensors["actions"].to(device=xs.device) stream_lengths = preprocessed["stream_lengths"] memory_selection_cfg = _cfg_get(self.cfg, "memory_selection", {}) poses_np = target_poses.detach().cpu().numpy() xs_pred = xs.clone() n_context_frames = min(self.context_frames // self.frame_stack, target_length) anchor_count = int(stream_lengths["anchor"]) online_anchor_indices = [ _select_online_anchor_indices( n_context_frames, anchor_count, memory_selection_cfg, poses=target_poses[:n_context_frames, batch_i, :5], ) for batch_i in range(batch_size) ] dynamic_policy = _dynamic_policy(memory_selection_cfg) multiview_selector = _dynamic_multiview_selector(memory_selection_cfg) if dynamic_policy == "multiview" else None event_dynamic_caches = [_new_online_event_cache(target_length) for _ in range(batch_size)] if dynamic_policy == "event_triggered" else None memory_sheet_targets = self._memory_sheet_target_indices(n_context_frames, target_length) if self.log_memory_selection_sheet else set() memory_sheet_records = [[] for _ in range(batch_size)] if memory_sheet_targets else None curr_frame = n_context_frames generated_frames = 0 rollout_start_time = time.perf_counter() pbar = None if curr_frame < target_length: pbar = tqdm( total=target_length, initial=curr_frame, desc=f"{namespace} sampling[{batch_idx}]", unit="frame", dynamic_ncols=True, ) while curr_frame < target_length: chunk_start_time = time.perf_counter() horizon = min(target_length - curr_frame, self.chunk_size) if self.chunk_size > 0 else target_length - curr_frame target_slice = slice(curr_frame, curr_frame + horizon) xs_pred[target_slice] = torch.randn_like(xs_pred[target_slice]).clamp(-self.clip_noise, self.clip_noise) start_frame = max(0, curr_frame + horizon - self.n_tokens) local_slice = slice(start_frame, curr_frame + horizon) target_length_step = local_slice.stop - local_slice.start query_offset = curr_frame - start_frame stream_indices_by_key = { key: torch.full((batch_size, int(stream_lengths[key])), -1, device=xs.device, dtype=torch.long) for key in _DEMEMWM_STREAM_KEYS } stream_masks_by_key = {key: torch.zeros_like(indices, dtype=torch.bool) for key, indices in stream_indices_by_key.items()} target_positions = np.arange(target_slice.start, target_slice.stop, dtype=np.int64) source_latents = xs_pred[:curr_frame] source_actions = target_actions[:curr_frame] source_poses = target_poses[:curr_frame] source_frame_indices = target_frame_indices[:curr_frame] for batch_i in range(batch_size): anchor_count = int(stream_indices_by_key["anchor"].shape[1]) dynamic_count = int(stream_indices_by_key["dynamic"].shape[1]) revisit_count = int(stream_indices_by_key["revisit"].shape[1]) batch_poses = poses_np[:, batch_i, :5] selected = {"anchor": online_anchor_indices[batch_i]} if dynamic_policy == "recent": selected["dynamic"] = _select_online_dynamic_indices( start_frame, dynamic_count, memory_selection_cfg, ) else: selected["dynamic"] = np.empty((0,), dtype=np.int64) fov_pool = None if dynamic_policy == "multiview" and multiview_selector == "fov_greedy": fov_pool = _build_shared_fov_candidate_pool( batch_poses, target_positions, memory_selection_cfg, namespace, min_candidate_frame=0, dynamic_count=dynamic_count, revisit_count=revisit_count, ) revisit_kwargs = {"fov_pool": fov_pool} if fov_pool is not None else {} selected["revisit"] = _select_revisit( batch_poses, target_positions, memory_selection_cfg, revisit_count, np.empty((0,), dtype=np.int64), namespace, **revisit_kwargs, ) if dynamic_policy == "event_triggered": cache = event_dynamic_caches[batch_i] _extend_online_event_cache( cache, source_latents[:, batch_i], source_actions[:, batch_i], source_poses[:, batch_i, :5], curr_frame, memory_selection_cfg, ) dynamic_stop = max( 0, int(target_positions[0]) - max(0, int(_cfg_get(memory_selection_cfg, "local_context_exclusion_frames", self.n_tokens))), ) selected["dynamic"] = _select_online_event_dynamic_from_cache( cache, dynamic_stop, dynamic_count, memory_selection_cfg, reference_frames=selected["revisit"] if len(selected["revisit"]) > 0 else None, ) elif dynamic_policy == "multiview": selected["dynamic"] = _select_dynamic_by_policy( int(target_positions[0]), dynamic_count, memory_selection_cfg, poses=batch_poses, reference_frames=selected["revisit"], excluded=selected["revisit"], target_positions=target_positions, split=namespace, fov_pool=fov_pool, ) if memory_sheet_records is not None and int(target_positions[0]) in memory_sheet_targets: memory_sheet_records[batch_i].append( { "target_index": int(target_positions[0]), "target_frame": int(target_frame_indices[int(target_positions[0]), batch_i].detach().cpu().item()), "streams": { key: [int(v) for v in np.asarray(selected.get(key, []), dtype=np.int64) if 0 <= int(v) < curr_frame] for key in _DEMEMWM_STREAM_KEYS }, } ) selected_masks = {key: np.ones(len(value), dtype=bool) for key, value in selected.items()} for key, indices in stream_indices_by_key.items(): count = indices.shape[1] selected_key = np.asarray(selected.get(key, []), dtype=np.int64)[:count] mask_key = np.asarray(selected_masks.get(key, []), dtype=bool)[:count] indices[batch_i, : len(selected_key)] = torch.as_tensor(selected_key, device=xs.device, dtype=torch.long) stream_masks_by_key[key][batch_i, : len(mask_key)] = torch.as_tensor(mask_key, device=xs.device, dtype=torch.bool) frame_memory_masks = { "target": torch.ones((batch_size, target_length_step), device=xs.device, dtype=torch.bool) if target_mask is None else target_mask[:, local_slice], **stream_masks_by_key, } stream_latents_by_key = { key: _gather_online_memory_tensor(source_latents, stream_indices_by_key[key], stream_masks_by_key[key]) for key in _DEMEMWM_STREAM_KEYS } stream_poses_by_key = { key: _gather_online_memory_tensor(source_poses, stream_indices_by_key[key], stream_masks_by_key[key]) for key in _DEMEMWM_STREAM_KEYS } stream_frame_indices_by_key = { key: _gather_online_memory_tensor(source_frame_indices, stream_indices_by_key[key], stream_masks_by_key[key]) for key in _DEMEMWM_STREAM_KEYS } target_conditions = target_tensors["action_conditions"][local_slice].to(device=xs.device) target_latents_step = xs_pred[local_slice] full_frame_memory_segments = { "target": int(target_length_step), **{key: int(stream_masks_by_key[key].shape[1]) for key in _DEMEMWM_STREAM_KEYS}, } scheduling_matrix = self._generate_scheduling_matrix(horizon) for m in range(scheduling_matrix.shape[0] - 1): from_query = torch.as_tensor(scheduling_matrix[m], device=xs.device, dtype=torch.long)[:, None].repeat(1, batch_size) to_query = torch.as_tensor(scheduling_matrix[m + 1], device=xs.device, dtype=torch.long)[:, None].repeat(1, batch_size) context_levels = _stabilized_sampling_levels(from_query, query_offset) from_target = torch.cat([context_levels, from_query], dim=0) to_target = torch.cat([context_levels, to_query], dim=0) memory_noise_levels_by_key = _memory_noise_levels_for_streams( getattr(self, "cfg", None), self.diffusion_model, from_query, full_frame_memory_segments, mode=namespace, ) routed_frame_memory_masks = _apply_memory_route_masks( frame_memory_masks, from_query, getattr(self, "cfg", None), self.diffusion_model, mode=namespace, ) ( active_packed_latents, active_packed_conditions, active_frame_memory_pose, active_frame_indices, active_frame_memory_segments, active_frame_memory_masks, active_memory_noise_levels, _active_streams, _pruned_streams, ) = _pack_active_inference_memory_streams( target_latents_step, target_conditions, target_poses[local_slice], target_frame_indices[local_slice], frame_memory_masks["target"], stream_latents_by_key, stream_poses_by_key, stream_frame_indices_by_key, routed_frame_memory_masks, memory_noise_levels_by_key, ) sampled_target = self.diffusion_model.sample_step( active_packed_latents, active_packed_conditions, None, torch.cat([from_target, active_memory_noise_levels], dim=0), torch.cat([to_target, active_memory_noise_levels], dim=0), current_frame=curr_frame, mode=namespace, reference_length=0, frame_idx=active_frame_indices, frame_memory_segments=active_frame_memory_segments, frame_memory_masks=active_frame_memory_masks, frame_memory_pose=active_frame_memory_pose, image_hw=image_hw, ) target_latents_step = sampled_target[:target_length_step] xs_pred[local_slice] = target_latents_step chunk_seconds = max(time.perf_counter() - chunk_start_time, 1e-9) end_frame = curr_frame + horizon curr_frame = end_frame generated_frames += horizon if pbar is not None: rollout_seconds = max(time.perf_counter() - rollout_start_time, 1e-9) postfix = { "range": f"{start_frame}:{end_frame}", "sec/frame": f"{chunk_seconds / horizon:.3f}", "frames/s": f"{horizon / chunk_seconds:.2f}", "avg_frames/s": f"{generated_frames / rollout_seconds:.2f}", } vram = _cuda_vram_postfix(xs.device) if vram is not None: postfix["vram"] = vram pbar.update(horizon) pbar.set_postfix(postfix) if pbar is not None: pbar.close() eval_start = n_context_frames latent_loss = F.mse_loss(xs_pred[eval_start:target_length], xs[eval_start:target_length], reduction="none") if target_mask is not None: loss_mask = rearrange(target_mask[:, eval_start:target_length].to(dtype=latent_loss.dtype), "b t -> t b") loss_mask = loss_mask.view(*loss_mask.shape, *((1,) * (latent_loss.ndim - 2))) latent_loss = (latent_loss * loss_mask).sum() / loss_mask.expand_as(latent_loss).sum().clamp_min(1.0) elif latent_loss.numel(): latent_loss = latent_loss.mean() else: latent_loss = xs_pred.new_tensor(0.0) self.log(f"{namespace}/latent_mse", latent_loss.detach()) if memory_sheet_records is not None: self._save_memory_selection_sheets(xs_pred, xs, target_frame_indices, memory_sheet_records, eval_start, batch_idx) if eval_start < target_length: xs_pred_decode = self.decode(xs_pred[eval_start:target_length].to(target_poses.device)) xs_decode = self.decode(xs[eval_start:target_length].to(target_poses.device)) if self.logger and self.log_video: log_video( xs_pred_decode, xs_decode, step=getattr(self, "global_step", 0), namespace=namespace + "_vis", prefix=f"batch{batch_idx:06d}", context_frames=self.context_frames, logger=self.logger.experiment, save_local=self.save_local, local_save_dir=self.local_save_dir, ) metric_mask = None if target_mask is not None: metric_mask = rearrange( target_mask[:, eval_start:target_length].to(device=xs_pred_decode.device), "b t -> t b", ) self._update_metric_accumulators(xs_pred_decode, xs_decode, metric_mask, eval_start) self._update_per_frame_metric_accumulators(xs_pred_decode, xs_decode, metric_mask, eval_start) self._update_revisit_metric_accumulators( xs_pred_decode, xs_decode, target_poses, target_frame_indices, metric_mask, eval_start, batch_idx, namespace, ) return latent_loss def test_step(self, batch, batch_idx) -> STEP_OUTPUT: return self.validation_step(batch, batch_idx, namespace="test") @torch.no_grad() def interactive(self, first_frame, new_actions, first_pose, device, memory_latent_frames, memory_actions, memory_poses, memory_c2w, memory_frame_idx): raise NotImplementedError( "DeMemWM interactive generation is unsupported until it uses the packed " "[target][anchor][dynamic][revisit] frame-memory API." )