| import hashlib |
| import json |
| import logging |
| import os |
| import random |
| from typing import Any, Dict, Optional |
|
|
| import torch |
| from safetensors.torch import load_file as safe_load_file |
|
|
| from diffsynth.pipelines.wan_video_new import WanVideoPipeline, ModelConfig |
| from src.model_training.transformers_compat import patch_transformers_hybrid_cache |
|
|
| patch_transformers_hybrid_cache() |
| from diffsynth.trainers.utils import DiffusionTrainingModule |
| from diffsynth.models.memory.geometry_spatial_memory import GeometrySpatialMemory |
| from diffsynth.models.memory.mixture_of_contexts import MixtureOfContexts |
| from diffsynth.models.memory.spatial_grid_memory import SpatialCrossAttnReadout, SpatialGridMemory |
| from src.model_training.fov_retrieval import flip_yaw_rt_list |
|
|
| logger = logging.getLogger(__name__) |
|
|
|
|
| class WanTrainingModule(DiffusionTrainingModule): |
| def __init__( |
| self, |
| model_paths=None, model_id_with_origin_paths=None, |
| tokenizer_path=None, |
| trainable_models=None, |
| lora_base_model=None, lora_target_modules="q,k,v,o,ffn.0,ffn.2", lora_rank=32, |
| use_gradient_checkpointing=True, |
| use_gradient_checkpointing_offload=False, |
| extra_inputs=None, |
| timestep_shift=1.0, |
| resume_from_checkpoint=None, |
| dataset_base_path: Optional[str] = None, |
| enable_context_memory=False, |
| context_memory_frames=8, |
| training_mode="context", |
| context_drop_prob: float = 0.0, |
| context_drop_seed: int = 42, |
| omit_context_actions: bool = False, |
| context_noise_prob=0.0, |
| context_noise_std=0.02, |
| context_fixed_noise_std=None, |
| teacher_forcing_prob=0.0, |
| yaw_flip_aug: bool = False, |
| context_per_frame_vae: bool = False, |
| context_source: str = "fov", |
| use_framepack_memory: bool = False, |
| context_temporal_decay: float = 1.0, |
| context_attention_weight: float = 1.0, |
| use_framepack_length_compress: bool = False, |
| framepack_ratio: int = 2, |
| framepack_length_strategy: str = "distance_merge", |
| framepack_recent_keep_ratio: float = 0.5, |
| framepack_multiscale_w2: float = 0.25, |
| framepack_multiscale_w4: float = 0.15, |
| use_spatial_memory: bool = False, |
| use_spatial_memory_legacy: bool = False, |
| spatial_memory_tokens: int = 64, |
| spatial_memory_grid: int = 8, |
| spatial_memory_inject_mode: str = "concat_text", |
| use_geometry_spatial_memory: bool = False, |
| geometry_spatial_memory_tokens: int = 64, |
| geometry_spatial_memory_grid: int = 8, |
| geometry_spatial_memory_temporal_bins: int = 4, |
| geometry_spatial_memory_inject_mode: str = "concat_text", |
| use_moc: bool = False, |
| moc_temperature: float = 1.0, |
| moc_top_k: int = 0, |
| |
| ): |
| super().__init__() |
| |
| model_configs = [] |
| if model_paths is not None: |
| model_paths = json.loads(model_paths) |
| model_configs += [ModelConfig(path=path) for path in model_paths] |
| if model_id_with_origin_paths is not None: |
| model_id_with_origin_paths = model_id_with_origin_paths.split(",") |
| model_configs += [ModelConfig(model_id=i.split(":")[0], origin_file_pattern=i.split(":")[1]) for i in model_id_with_origin_paths] |
| from_pretrained_kw = {"torch_dtype": torch.bfloat16, "device": "cpu", "model_configs": model_configs} |
| if tokenizer_path: |
| from_pretrained_kw["tokenizer_config"] = ModelConfig(path=tokenizer_path) |
| self.pipe = WanVideoPipeline.from_pretrained(**from_pretrained_kw) |
| |
| |
| self.timestep_shift = timestep_shift |
| |
| |
| self.pipe.scheduler.set_timesteps(1000, training=True, shift=timestep_shift) |
| |
| |
| self.pipe.freeze_except([] if trainable_models is None else trainable_models.split(",")) |
| |
| |
| if lora_base_model is not None: |
| model = self.add_lora_to_model( |
| getattr(self.pipe, lora_base_model), |
| target_modules=lora_target_modules.split(","), |
| lora_rank=lora_rank |
| ) |
| setattr(self.pipe, lora_base_model, model) |
| |
| |
| if resume_from_checkpoint is not None: |
| logger.info(f"Loading LoRA checkpoint from: {resume_from_checkpoint}") |
| if not os.path.exists(resume_from_checkpoint): |
| raise FileNotFoundError(f"Checkpoint file not found: {resume_from_checkpoint}") |
| checkpoint_state_dict = safe_load_file(resume_from_checkpoint) |
| logger.info(f"Checkpoint contains {len(checkpoint_state_dict)} parameters") |
| |
| |
| |
| missing_keys, unexpected_keys = model.load_state_dict(checkpoint_state_dict, strict=False) |
| if missing_keys: |
| logger.warning(f"{len(missing_keys)} keys were missing when loading checkpoint") |
| if len(missing_keys) <= 10: |
| logger.debug(f"Missing keys: {missing_keys}") |
| if unexpected_keys: |
| logger.warning(f"{len(unexpected_keys)} unexpected keys in checkpoint (will be ignored)") |
| if len(unexpected_keys) <= 10: |
| logger.debug(f"Unexpected keys: {unexpected_keys}") |
| loaded_count = len(checkpoint_state_dict) - len(missing_keys) - len(unexpected_keys) |
| logger.info(f"Successfully loaded {loaded_count} parameters from checkpoint!") |
| |
| |
| self.use_gradient_checkpointing = use_gradient_checkpointing |
| self.use_gradient_checkpointing_offload = use_gradient_checkpointing_offload |
| self.extra_inputs = extra_inputs.split(",") if extra_inputs is not None else [] |
| self.dataset_base_path = dataset_base_path |
| |
| |
| self.enable_context_memory = enable_context_memory |
| self.context_memory_frames = context_memory_frames |
| self.training_mode = training_mode |
| self.context_drop_prob = float(context_drop_prob or 0.0) |
| self.context_drop_seed = int(context_drop_seed or 42) |
| self.omit_context_actions = bool(omit_context_actions) |
| self.context_per_frame_vae = bool(context_per_frame_vae) |
| self.context_source = (context_source or "fov").strip().lower() |
| if self.context_source not in ("fov", "replay", "prev_chunk_tail"): |
| self.context_source = "fov" |
| self.context_noise_prob = context_noise_prob |
| self.context_noise_std = context_noise_std |
| self.context_fixed_noise_std = context_fixed_noise_std |
| self.teacher_forcing_prob = teacher_forcing_prob |
| self.teacher_forcing_enabled = teacher_forcing_prob > 0.0 |
| self.yaw_flip_aug = bool(yaw_flip_aug) |
| |
| self.use_framepack_memory = bool(use_framepack_memory) |
| self.context_temporal_decay = float(context_temporal_decay or 1.0) |
| self.context_attention_weight = float(context_attention_weight or 1.0) |
| self.use_framepack_length_compress = bool(use_framepack_length_compress) |
| self.framepack_ratio = int(framepack_ratio or 2) |
| self.framepack_length_strategy = str(framepack_length_strategy or "distance_merge").lower() |
| self.framepack_recent_keep_ratio = float(framepack_recent_keep_ratio or 0.5) |
| self.framepack_multiscale_w2 = float(framepack_multiscale_w2 or 0.25) |
| self.framepack_multiscale_w4 = float(framepack_multiscale_w4 or 0.15) |
| |
| self.pipe.use_framepack_memory = self.use_framepack_memory |
| self.pipe.context_temporal_decay = self.context_temporal_decay |
| self.pipe.context_attention_weight = self.context_attention_weight |
| self.pipe.use_framepack_length_compress = self.use_framepack_length_compress |
| self.pipe.framepack_ratio = self.framepack_ratio |
| self.pipe.framepack_length_strategy = self.framepack_length_strategy |
| self.pipe.framepack_recent_keep_ratio = self.framepack_recent_keep_ratio |
| self.pipe.framepack_multiscale_w2 = self.framepack_multiscale_w2 |
| self.pipe.framepack_multiscale_w4 = self.framepack_multiscale_w4 |
| self.use_moc = bool(use_moc) |
| self.moc_temperature = float(moc_temperature or 1.0) |
| self.moc_top_k = int(moc_top_k or 0) |
| self.moc_module = MixtureOfContexts( |
| temperature=self.moc_temperature, |
| top_k=self.moc_top_k, |
| ) if self.use_moc else None |
| self.pipe.use_moc = self.use_moc |
| self.pipe.moc_module = self.moc_module |
| self.pipe.use_spatial_memory = bool(use_spatial_memory) |
| self.pipe.use_spatial_memory_legacy = bool(use_spatial_memory_legacy) |
| self.pipe.spatial_memory_tokens = int(spatial_memory_tokens or 64) |
| self.pipe.spatial_memory_inject_mode = str(spatial_memory_inject_mode or "concat_text") |
| self.spatial_memory_module = None |
| self.spatial_memory_readout_module = None |
| if self.pipe.use_spatial_memory and not self.pipe.use_spatial_memory_legacy: |
| dim = int(getattr(self.pipe.dit, "dim")) |
| grid_size = int(spatial_memory_grid or 8) |
| self.pipe.spatial_memory_grid = grid_size |
| self.spatial_memory_module = SpatialGridMemory( |
| dim=dim, |
| grid_size=grid_size, |
| num_tokens=self.pipe.spatial_memory_tokens, |
| ) |
| self.pipe.spatial_memory_module = self.spatial_memory_module |
| if self.pipe.spatial_memory_inject_mode == "cross_attn_readout": |
| self.spatial_memory_readout_module = SpatialCrossAttnReadout(dim=dim, num_heads=8) |
| self.pipe.spatial_memory_readout_module = self.spatial_memory_readout_module |
| else: |
| self.pipe.spatial_memory_module = None |
| self.pipe.spatial_memory_readout_module = None |
| self.use_geometry_spatial_memory = bool(use_geometry_spatial_memory) |
| self.geometry_spatial_memory_module = None |
| self.geometry_spatial_memory_readout_module = None |
| self.pipe.use_geometry_spatial_memory = self.use_geometry_spatial_memory |
| self.pipe.geometry_spatial_memory_inject_mode = str( |
| geometry_spatial_memory_inject_mode or "concat_text" |
| ) |
| if self.use_geometry_spatial_memory: |
| dim = int(getattr(self.pipe.dit, "dim")) |
| self.geometry_spatial_memory_module = GeometrySpatialMemory( |
| dim=dim, |
| latent_channels=int(getattr(self.pipe.dit, "in_dim", 16)), |
| patch_size=tuple(getattr(self.pipe.dit, "patch_size", (1, 2, 2))), |
| grid_size=int(geometry_spatial_memory_grid or 8), |
| temporal_bins=int(geometry_spatial_memory_temporal_bins or 4), |
| num_tokens=int(geometry_spatial_memory_tokens or 64), |
| ) |
| self.geometry_spatial_memory_module.initialize_from_dit_patch_embedding( |
| self.pipe.dit.patch_embedding |
| ) |
| dit_parameter = next(self.pipe.dit.parameters()) |
| self.geometry_spatial_memory_module = self.geometry_spatial_memory_module.to( |
| device=dit_parameter.device, |
| dtype=dit_parameter.dtype, |
| ) |
| self.pipe.geometry_spatial_memory_module = self.geometry_spatial_memory_module |
| if self.pipe.geometry_spatial_memory_inject_mode == "cross_attn_readout": |
| self.geometry_spatial_memory_readout_module = SpatialCrossAttnReadout( |
| dim=dim, |
| num_heads=8, |
| ).to(device=dit_parameter.device, dtype=dit_parameter.dtype) |
| self.pipe.geometry_spatial_memory_readout_module = ( |
| self.geometry_spatial_memory_readout_module |
| ) |
| else: |
| self.pipe.geometry_spatial_memory_module = None |
| self.pipe.geometry_spatial_memory_readout_module = None |
| |
| self.current_step = 0 |
| |
| def _forward_preprocess_batch(self, samples: list) -> dict: |
| """Batch preprocessing for Stage 1 Interactive (no context). data is list of sample dicts.""" |
| if not samples: |
| raise ValueError("samples cannot be empty in _forward_preprocess_batch") |
| batch_size = len(samples) |
| prompts = [] |
| video_frames_list = [] |
| actions_list = [] |
| for s in samples: |
| p = s.get("prompt") |
| if p is None: |
| raise ValueError("sample['prompt'] is missing or None") |
| prompts.append(str(p) if not isinstance(p, str) else p) |
| video_frames_list.append(s["video"]) |
| if "actions" in s and s["actions"] is not None: |
| acts = s["actions"] |
| if getattr(self, 'yaw_flip_aug', False) and isinstance(acts, list) and len(acts) > 0 and isinstance(acts[0], (list, tuple)) and len(acts[0]) >= 12 and random.random() < 0.5: |
| acts = flip_yaw_rt_list(acts) |
| if isinstance(acts, torch.Tensor): |
| actions_list.append(acts) |
| elif isinstance(acts, list) and len(acts) > 0: |
| actions_list.append(torch.tensor(acts, dtype=torch.float32)) |
| else: |
| actions_list.append(None) |
| else: |
| actions_list.append(None) |
| |
| |
| input_video = video_frames_list |
| first = samples[0] |
| h, w = first["video"][0].size[1], first["video"][0].size[0] |
| num_frames = len(first["video"]) |
| |
| inputs_posi = {"prompt": prompts} |
| inputs_nega = {} |
| inputs_shared = { |
| "input_video": input_video, |
| "height": h, |
| "width": w, |
| "num_frames": num_frames, |
| "batch_size": batch_size, |
| "cfg_scale": 1, |
| "tiled": False, |
| "rand_device": self.pipe.device, |
| "use_gradient_checkpointing": self.use_gradient_checkpointing, |
| "use_gradient_checkpointing_offload": self.use_gradient_checkpointing_offload, |
| "cfg_merge": False, |
| "vace_scale": 1, |
| } |
| |
| ref_action = next((a for a in actions_list if a is not None), None) |
| if ref_action is not None and batch_size == 1: |
| inputs_shared["actions"] = ref_action.detach().cpu().tolist() if isinstance(ref_action, torch.Tensor) else ref_action |
| elif ref_action is not None: |
| device = self.pipe.device |
| dtype = ref_action.dtype |
| stacked = [] |
| for a in actions_list: |
| if a is not None: |
| stacked.append(a.to(device=device)) |
| else: |
| stacked.append(torch.zeros_like(ref_action, device=device, dtype=dtype)) |
| inputs_shared["actions"] = torch.stack(stacked) |
| else: |
| inputs_shared["actions"] = None |
| |
| for unit in self.pipe.units: |
| inputs_shared, inputs_posi, inputs_nega = self.pipe.unit_runner(unit, self.pipe, inputs_shared, inputs_posi, inputs_nega) |
| return {**inputs_shared, **inputs_posi} |
| |
| def _build_context_with_anchor(self, context_frames, context_actions=None, expected_k=None): |
| """Training-side anchor helper: keep last frame as mandatory anchor and keep action length aligned.""" |
| frames = list(context_frames or []) |
| actions = list(context_actions or []) if context_actions is not None else [] |
| if not frames or not getattr(self, "use_anchor_frame", False): |
| return frames, actions |
| k = int(expected_k) if (expected_k is not None and int(expected_k) > 0) else len(frames) |
| if len(frames) > k: |
| frames = frames[-k:] |
| if actions: |
| actions = actions[-k:] |
| if actions: |
| if len(actions) < len(frames): |
| actions = actions + [actions[-1]] * (len(frames) - len(actions)) |
| elif len(actions) > len(frames): |
| actions = actions[:len(frames)] |
| return frames, actions |
|
|
| def _forward_preprocess_batch_context(self, samples: list) -> dict: |
| """Batch preprocessing for Stage 2 Context Memory. Batch-level drop: if drop, all samples get no context.""" |
| if not samples: |
| raise ValueError("samples cannot be empty in _forward_preprocess_batch_context") |
| batch_size = len(samples) |
| first = samples[0] |
| |
| def _should_drop_context(_data) -> bool: |
| p = float(getattr(self, "context_drop_prob", 0.0) or 0.0) |
| if p <= 0.0: |
| return False |
| if p >= 1.0: |
| return True |
| vn = str(_data.get("video_name", "")) |
| sf = str(_data.get("start_frame", "")) |
| key = f"{int(getattr(self, 'context_drop_seed', 42))}|{vn}|{sf}" |
| h = hashlib.md5(key.encode("utf-8")).hexdigest() |
| u = int(h[:8], 16) / 0xFFFFFFFF |
| return u < p |
| |
| |
| dropped_context = _should_drop_context(first) |
| |
| |
| |
| |
| try: |
| import torch.distributed as dist |
| if dist.is_available() and dist.is_initialized(): |
| flag = torch.tensor([1 if dropped_context else 0], device=self.pipe.device, dtype=torch.int64) |
| dist.broadcast(flag, src=0) |
| dropped_context = bool(int(flag.item())) |
| except Exception: |
| pass |
| |
| prompts = [] |
| video_frames_list = [] |
| actions_list = [] |
| context_latents_list = [] |
| context_actions_list = [] |
| geometry_memory_latents_list = [] |
| expected_k = self.context_memory_frames |
| training_mode = getattr(self, 'training_mode', 'context') |
| |
| target_h = first["video"][0].size[1] |
| target_w = first["video"][0].size[0] |
| num_frames = len(first["video"]) |
| |
| from PIL import Image |
| |
| for s in samples: |
| p = s.get("prompt") |
| if p is None: |
| raise ValueError("sample['prompt'] is missing or None") |
| prompts.append(str(p) if not isinstance(p, str) else p) |
| video_frames_list.append(s["video"]) |
| |
| if "actions" in s and s["actions"] is not None: |
| acts = s["actions"] |
| if getattr(self, 'yaw_flip_aug', False) and isinstance(acts, list) and len(acts) > 0 and isinstance(acts[0], (list, tuple)) and len(acts[0]) >= 12 and random.random() < 0.5: |
| acts = flip_yaw_rt_list(acts) |
| if isinstance(acts, torch.Tensor): |
| actions_list.append(acts) |
| elif isinstance(acts, list) and len(acts) > 0: |
| actions_list.append(torch.tensor(acts, dtype=torch.float32)) |
| else: |
| actions_list.append(None) |
| else: |
| actions_list.append(None) |
|
|
| geometry_frames = s.get("geometry_memory_frames") or [] |
| if self.use_geometry_spatial_memory: |
| if not geometry_frames: |
| raise ValueError( |
| "Geometry-grounded Spatial Memory requires sample['geometry_memory_frames']. " |
| "Provide TSDF/point-cloud renders through the configured metadata column." |
| ) |
| resized_geometry = [] |
| for frame in geometry_frames: |
| if hasattr(frame, "resize") and hasattr(frame, "size"): |
| gw, gh = frame.size |
| if gh != target_h or gw != target_w: |
| frame = frame.resize((target_w, target_h), Image.Resampling.LANCZOS) |
| resized_geometry.append(frame) |
| with torch.no_grad(): |
| geometry_video = self.pipe.preprocess_video(resized_geometry) |
| if geometry_video.dim() == 4: |
| geometry_video = geometry_video.unsqueeze(0) |
| geometry_latents = self.pipe.vae.encode( |
| [geometry_video[i] for i in range(geometry_video.shape[0])], |
| device=self.pipe.device, |
| tiled=False, |
| tile_size=None, |
| tile_stride=None, |
| ) |
| geometry_memory_latents_list.append( |
| geometry_latents.to(dtype=self.pipe.torch_dtype, device=self.pipe.device) |
| ) |
| else: |
| geometry_memory_latents_list.append(None) |
| |
| if dropped_context: |
| context_latents_list.append(None) |
| context_actions_list.append(None) |
| continue |
| |
| ctx_frames = s.get("context_frames") or [] |
| ctx_actions = [] if getattr(self, "omit_context_actions", False) else (s.get("context_actions") or []) |
| context_indices = s.get("context_frame_indices", []) |
| start_frame = s.get("start_frame", None) |
| end_frame = s.get("end_frame", None) |
| |
| if ctx_frames and context_indices and start_frame is not None and end_frame is not None: |
| filtered_frames, filtered_actions = [ctx_frames[0]], [] |
| if ctx_actions: |
| filtered_actions.append(ctx_actions[0]) |
| for i in range(1, len(ctx_frames)): |
| idx = context_indices[i] if i < len(context_indices) else None |
| if idx is None or idx < start_frame or idx > end_frame: |
| filtered_frames.append(ctx_frames[i]) |
| if ctx_actions and i < len(ctx_actions): |
| filtered_actions.append(ctx_actions[i]) |
| ctx_frames, ctx_actions = filtered_frames, filtered_actions if filtered_actions else ctx_actions |
| |
| if not ctx_frames and len(s["video"]) > expected_k: |
| ctx_frames = s["video"][:expected_k] |
| if s.get("actions") and len(s["actions"]) >= expected_k: |
| ctx_actions = s["actions"][:expected_k] |
| |
| if not ctx_frames: |
| context_latents_list.append(None) |
| context_actions_list.append(None) |
| continue |
| |
| resized = [] |
| for f in ctx_frames: |
| if hasattr(f, 'resize') and hasattr(f, 'size'): |
| w, h = f.size |
| if h != target_h or w != target_w: |
| f = f.resize((target_w, target_h), Image.Resampling.LANCZOS) |
| resized.append(f) |
| ctx_frames = resized |
| |
| if len(ctx_frames) < expected_k: |
| last = ctx_frames[-1] if ctx_frames else Image.new('RGB', (target_w, target_h), (0, 0, 0)) |
| ctx_frames = ctx_frames + [last] * (expected_k - len(ctx_frames)) |
| if ctx_actions: |
| ctx_actions = ctx_actions + [ctx_actions[-1]] * (expected_k - len(ctx_actions)) |
| elif len(ctx_frames) > expected_k: |
| ctx_frames = ctx_frames[:expected_k] |
| ctx_actions = ctx_actions[:expected_k] if ctx_actions else [] |
|
|
| ctx_frames, ctx_actions = self._build_context_with_anchor( |
| ctx_frames, |
| context_actions=ctx_actions, |
| expected_k=expected_k, |
| ) |
| |
| with torch.no_grad(): |
| if getattr(self, "context_per_frame_vae", False): |
| |
| context_latents_per_sample = [] |
| for f in ctx_frames: |
| frame_video = self.pipe.preprocess_video([f]) |
| frame_sq = frame_video.squeeze(0) |
| lat_one = self.pipe.vae.encode([frame_sq], device=self.pipe.device, tiled=False, tile_size=None, tile_stride=None) |
| context_latents_per_sample.append(lat_one) |
| lat = torch.cat(context_latents_per_sample, dim=2) |
| else: |
| ctx_video = self.pipe.preprocess_video(ctx_frames) |
| if ctx_video.dim() == 4: |
| ctx_video = ctx_video.unsqueeze(0) |
| lat = self.pipe.vae.encode([ctx_video[i] for i in range(ctx_video.shape[0])], device=self.pipe.device, tiled=False, tile_size=None, tile_stride=None) |
| context_latents_list.append(lat.to(dtype=self.pipe.torch_dtype, device=self.pipe.device)) |
| |
| if ctx_actions: |
| if isinstance(ctx_actions[0], (list, tuple)): |
| context_actions_list.append(torch.tensor(ctx_actions, dtype=torch.float32)) |
| else: |
| context_actions_list.append(torch.tensor(ctx_actions, dtype=torch.float32)) |
| else: |
| context_actions_list.append(None) |
| |
| input_video = video_frames_list |
| inputs_posi = {"prompt": prompts} |
| inputs_nega = {} |
| inputs_shared = { |
| "input_video": input_video, |
| "height": target_h, |
| "width": target_w, |
| "num_frames": num_frames, |
| "batch_size": batch_size, |
| "cfg_scale": 1, |
| "tiled": False, |
| "rand_device": self.pipe.device, |
| "use_gradient_checkpointing": self.use_gradient_checkpointing, |
| "use_gradient_checkpointing_offload": self.use_gradient_checkpointing_offload, |
| "cfg_merge": False, |
| "vace_scale": 1, |
| } |
| |
| |
| |
| has_context_step = (not dropped_context) and any(x is not None for x in context_latents_list) |
| try: |
| import torch.distributed as dist |
| if dist.is_available() and dist.is_initialized(): |
| flag = torch.tensor([1 if has_context_step else 0], device=self.pipe.device, dtype=torch.int64) |
| dist.all_reduce(flag, op=dist.ReduceOp.MIN) |
| has_context_step = bool(int(flag.item())) |
| except Exception: |
| pass |
| if not has_context_step: |
| dropped_context = True |
|
|
| if not dropped_context and any(x is not None for x in context_latents_list): |
| valid = [x for x in context_latents_list if x is not None] |
| if valid: |
| ref = valid[0] |
| device, dtype = self.pipe.device, ref.dtype |
| stacked_ctx = [] |
| for x in context_latents_list: |
| if x is not None: |
| stacked_ctx.append(x.to(device=device)) |
| else: |
| stacked_ctx.append(torch.zeros_like(ref, device=device, dtype=dtype)) |
| inputs_shared["context_latents"] = torch.cat(stacked_ctx, dim=0) |
| inputs_shared["num_context_frames"] = ref.shape[2] |
| inputs_shared["training_mode"] = training_mode |
| inputs_shared["context_noise_prob"] = getattr(self, 'context_noise_prob', 0.0) |
| inputs_shared["context_noise_std"] = getattr(self, 'context_noise_std', 0.02) |
| if self.context_fixed_noise_std is not None: |
| inputs_shared["context_fixed_noise_std"] = self.context_fixed_noise_std |
| inputs_shared["context_position"] = os.environ.get("CONTEXT_POSITION", "suffix") |
| inputs_shared["omit_context_actions"] = getattr(self, "omit_context_actions", False) |
| inputs_shared["context_attention_weight"] = getattr(self, "context_attention_weight", 1.0) |
| inputs_shared["use_anchor_frame"] = getattr(self, "use_anchor_frame", False) |
| inputs_shared["context_temporal_decay"] = getattr(self, "context_temporal_decay", 1.0) |
| inputs_shared["use_spatial_memory"] = getattr(self.pipe, "use_spatial_memory", False) |
| inputs_shared["spatial_memory_tokens"] = int(getattr(self.pipe, "spatial_memory_tokens", 64) or 64) |
| inputs_shared["use_spatial_memory_legacy"] = bool(getattr(self.pipe, "use_spatial_memory_legacy", False)) |
| inputs_shared["spatial_memory_module"] = getattr(self.pipe, "spatial_memory_module", None) |
| inputs_shared["spatial_memory_inject_mode"] = getattr(self.pipe, "spatial_memory_inject_mode", "concat_text") |
| inputs_shared["spatial_memory_readout_module"] = getattr(self.pipe, "spatial_memory_readout_module", None) |
| inputs_shared["use_framepack_memory"] = bool(getattr(self, "use_framepack_memory", False)) |
| if self.use_moc and self.moc_module is not None: |
| inputs_shared["use_moc"] = True |
| inputs_shared["moc_module"] = self.moc_module |
| nf_list = [s.get("non_fov_frames") or [] for s in samples] |
| if any(nf for nf in nf_list): |
| inputs_shared["non_fov_frames_list"] = nf_list |
|
|
| if self.use_geometry_spatial_memory: |
| if not all(x is not None for x in geometry_memory_latents_list): |
| raise ValueError("Geometry memory is missing for one or more samples in the batch.") |
| inputs_shared["geometry_memory_latents"] = torch.cat( |
| geometry_memory_latents_list, |
| dim=0, |
| ) |
| inputs_shared["use_geometry_spatial_memory"] = True |
| inputs_shared["geometry_spatial_memory_module"] = ( |
| self.geometry_spatial_memory_module |
| ) |
| inputs_shared["geometry_spatial_memory_inject_mode"] = ( |
| self.pipe.geometry_spatial_memory_inject_mode |
| ) |
| inputs_shared["geometry_spatial_memory_readout_module"] = ( |
| self.geometry_spatial_memory_readout_module |
| ) |
|
|
| ctx_acts_valid = [a for a in context_actions_list if a is not None] |
| if not getattr(self, "omit_context_actions", False) and ctx_acts_valid: |
| ref_act = ctx_acts_valid[0] |
| target_len = ref_act.shape[0] |
| stacked_ca = [] |
| for a in context_actions_list: |
| if a is not None: |
| a = a.to(device=device) |
| if a.shape[0] != target_len: |
| if a.shape[0] > target_len: |
| a = a[:target_len] |
| else: |
| pad = a.new_zeros(target_len - a.shape[0], a.shape[-1]) |
| a = torch.cat([a, pad], dim=0) |
| stacked_ca.append(a) |
| else: |
| stacked_ca.append(torch.zeros_like(ref_act, device=device, dtype=ref_act.dtype)) |
| inputs_shared["context_actions"] = torch.stack(stacked_ca) |
| |
| ref_action = next((a for a in actions_list if a is not None), None) |
| if ref_action is not None and batch_size == 1: |
| inputs_shared["actions"] = ref_action.detach().cpu().tolist() if isinstance(ref_action, torch.Tensor) else ref_action |
| elif ref_action is not None: |
| device = self.pipe.device |
| dtype = ref_action.dtype |
| stacked = [] |
| for a in actions_list: |
| if a is not None: |
| stacked.append(a.to(device=device)) |
| else: |
| stacked.append(torch.zeros_like(ref_action, device=device, dtype=dtype)) |
| inputs_shared["actions"] = torch.stack(stacked) |
| else: |
| inputs_shared["actions"] = None |
| |
| for unit in self.pipe.units: |
| inputs_shared, inputs_posi, inputs_nega = self.pipe.unit_runner(unit, self.pipe, inputs_shared, inputs_posi, inputs_nega) |
| return {**inputs_shared, **inputs_posi} |
| |
| @staticmethod |
| def _translate_condition_keys(d): |
| """Map VWM CamVideoDataset condition_* keys to context-memory keys.""" |
| if not isinstance(d, dict): |
| return d |
| if "condition_frames" in d and "context_frames" not in d: |
| d["context_frames"] = d.pop("condition_frames") |
| if "condition_actions" in d and "context_actions" not in d: |
| d["context_actions"] = d.pop("condition_actions") |
| if "condition_frame_indices" in d and "context_frame_indices" not in d: |
| d["context_frame_indices"] = d.pop("condition_frame_indices") |
| if "use_condition_context_frames" in d: |
| d.pop("use_condition_context_frames") |
| if "condition_source" in d: |
| d.pop("condition_source", None) |
| return d |
|
|
| def forward_preprocess(self, data): |
| if data is None: |
| raise ValueError("data cannot be None in forward_preprocess") |
| samples = data if isinstance(data, list) else [data] |
| samples = [self._translate_condition_keys(d) for d in samples] |
| if self.enable_context_memory: |
| return self._forward_preprocess_batch_context(samples) |
| return self._forward_preprocess_batch(samples) |
|
|
| def _ensure_input_latents(self, inputs: Dict[str, Any], *, strict: bool = False) -> Dict[str, Any]: |
| if "input_latents" in inputs: |
| return inputs |
| import warnings |
| video_obj = inputs.get("input_video", None) |
| if video_obj is None: |
| video_obj = inputs.get("video", None) |
| vae = getattr(self.pipe, "vae", None) |
| rebuild_err = None |
| if video_obj is not None and vae is not None and hasattr(vae, "encode"): |
| try: |
| if isinstance(video_obj, list): |
| video_tensor = self.pipe.preprocess_video(video_obj) |
| else: |
| video_tensor = video_obj |
| if hasattr(video_tensor, "dim"): |
| video_sq = video_tensor.squeeze(0) if video_tensor.dim() == 5 else video_tensor |
| with torch.no_grad(): |
| try: |
| lat = vae.encode(video_tensor, device=self.pipe.device, tiled=False, tile_size=None, tile_stride=None) |
| except Exception as e_first: |
| |
| |
| |
| |
| try: |
| lat = vae.encode([video_sq], device=self.pipe.device, tiled=False, tile_size=None, tile_stride=None) |
| except Exception as e_retry: |
| raise RuntimeError( |
| f"VAE encode failed -- tensor form: {e_first!r}; list form: {e_retry!r}" |
| ) from e_retry |
| if isinstance(lat, (list, tuple)): |
| lat = lat[0] |
| if hasattr(lat, "dim") and lat.dim() == 4: |
| lat = lat.unsqueeze(0) |
| inputs["input_latents"] = lat.to(dtype=torch.bfloat16, device=self.pipe.device) |
| return inputs |
| except Exception as e: |
| rebuild_err = e |
| warnings.warn(f"Failed to rebuild input_latents: {e}") |
| msg = ( |
| "input_latents missing and auto-rebuild failed" |
| + (f" (rebuild error: {rebuild_err!r})" if rebuild_err |
| else " (no input_video/video or vae unavailable)") |
| + f". available input keys={sorted(list(inputs.keys()))}" |
| ) |
| if strict: |
| raise KeyError(msg) |
| warnings.warn(msg) |
| return inputs |
|
|
| def restore_after_sampling(self): |
| """Restore training-time pipe state clobbered by the periodic sampling |
| monitor (``pipe.__call__``) and release its GPU cache. Called by the |
| ModelLogger after every paper-process sampling step so the next training |
| step is unaffected. |
| |
| - Scheduler: sampling runs ``set_timesteps(num_inference_steps, |
| training=False)``; ``training_loss`` reads ``self.scheduler.timesteps`` |
| directly, so we must re-apply the training schedule (1000 steps, |
| ``training=True``) -- otherwise every later step silently samples from |
| inference timesteps/sigmas (wrong loss). ``self.timestep_shift`` is |
| stored in ``__init__`` for exactly this. |
| - Cache: ~50 denoise steps + an 81-frame VAE decode leave the main rank's |
| GPU fragmented; releasing the cache prevents the next step's target |
| VAE encode (which auto-rebuilds ``input_latents``) from OOMing. |
| """ |
| self.pipe.scheduler.set_timesteps(1000, training=True, shift=self.timestep_shift) |
| import gc |
| gc.collect() |
| if torch.cuda.is_available(): |
| torch.cuda.empty_cache() |
|
|
| def dump_input_video_debug(self, samples, step, out_dir, fps=15): |
| """Cache the raw dataset video + a VAE encode→decode roundtrip + color |
| stats to diagnose train/infer color consistency (e.g. color inversion). |
| |
| Writes, for ``samples[0]``, to ``<out_dir>/step_{N:07d}_*``: |
| * ``_raw.mp4`` — the raw input frames (``d["video"]``) |
| * ``_vae_roundtrip.mp4`` — ``preprocess_video → vae.encode → vae.decode`` |
| (mirrors the exact encode/decode calls used in training's |
| ``_ensure_input_latents`` and the pipeline's decode path, so a color |
| artifact here implicates the VAE / preprocessing, not the DiT/CGLA) |
| * ``_stats.json`` — per-channel mean RGB (raw vs roundtrip), |
| latent mean/std/min/max, and inversion / R-B-swap flags |
| |
| All under ``torch.no_grad`` and wrapped so a failure never aborts training |
| (returns the error string). No VRAM management in training → |
| ``load_models_to_device`` is a no-op, so the VAE stays on the train device. |
| """ |
| import json as _json |
| import os as _os |
| import numpy as _np |
| from diffsynth import save_video |
|
|
| if not samples: |
| return "no samples" |
| d = samples[0] |
| frames = d.get("video") or [] |
| if not frames: |
| return "no video frames in sample" |
| vae = getattr(self.pipe, "vae", None) |
| if vae is None or not hasattr(vae, "encode") or not hasattr(vae, "decode"): |
| return "vae unavailable" |
|
|
| _os.makedirs(out_dir, exist_ok=True) |
| tag = f"step_{int(step):07d}" |
| vn = str(d.get("video_name", "")) |
| sf = int(d.get("start_frame", 0) or 0) |
|
|
| raw_path = _os.path.join(out_dir, f"{tag}_raw.mp4") |
| roundtrip_path = _os.path.join(out_dir, f"{tag}_vae_roundtrip.mp4") |
| stats_path = _os.path.join(out_dir, f"{tag}_stats.json") |
|
|
| |
| save_video(list(frames), raw_path, fps=fps, quality=5) |
|
|
| stats = {"step": int(step), "video_name": vn, "start_frame": sf, |
| "num_frames": len(frames), "vae_dtype": str(next(vae.parameters()).dtype)} |
|
|
| |
| |
| |
| rec_frames = None |
| try: |
| video_tensor = self.pipe.preprocess_video(list(frames)) |
| with torch.no_grad(): |
| lat = vae.encode(video_tensor, device=self.pipe.device, tiled=False, |
| tile_size=None, tile_stride=None) |
| if isinstance(lat, (list, tuple)): |
| lat = lat[0] |
| rec = vae.decode(lat, device=self.pipe.device, tiled=False, |
| tile_size=None, tile_stride=None) |
| rec_frames = self.pipe.vae_output_to_video(rec) |
| save_video(list(rec_frames), roundtrip_path, fps=fps, quality=5) |
| except Exception as e: |
| stats["roundtrip_error"] = repr(e) |
|
|
| |
| def _mean_rgb(pil_list): |
| arr = _np.stack([_np.asarray(f.convert("RGB"), dtype=_np.float32) for f in pil_list]) |
| return arr.reshape(-1, 3).mean(0).tolist() |
|
|
| raw_mean = _mean_rgb(frames) |
| stats["raw_mean_rgb"] = raw_mean |
| if rec_frames is not None: |
| rec_mean = _mean_rgb(rec_frames) |
| stats["roundtrip_mean_rgb"] = rec_mean |
| r_raw, g_raw, b_raw = raw_mean |
| r_rec, g_rec, b_rec = rec_mean |
| stats["color_inversion_flag"] = bool( |
| abs((255 - r_raw) - r_rec) < abs(r_raw - r_rec) or |
| abs((255 - g_raw) - g_rec) < abs(g_raw - g_rec) |
| ) |
| stats["rb_swap_flag"] = bool(abs(r_raw - b_rec) < abs(r_raw - r_rec)) |
| if torch.is_tensor(lat): |
| _lat = lat.detach().float() |
| stats["latent_mean"] = float(_lat.mean().item()) |
| stats["latent_std"] = float(_lat.std().item()) |
| stats["latent_min"] = float(_lat.min().item()) |
| stats["latent_max"] = float(_lat.max().item()) |
|
|
| with open(stats_path, "w", encoding="utf-8") as f: |
| _json.dump(stats, f, ensure_ascii=False, indent=2) |
| return f"saved {tag} (raw_mean_rgb={raw_mean})" |
|
|
| def forward(self, data, inputs=None): |
| if inputs is None: |
| inputs = self.forward_preprocess(data) |
| models = {name: getattr(self.pipe, name) for name in self.pipe.in_iteration_models} |
| if self.enable_context_memory and "context_latents" in inputs: |
| return self._training_loss_with_context(**models, **inputs) |
| inputs = self._ensure_input_latents(inputs, strict=True) |
| return self.pipe.training_loss(**models, **inputs) |
| |
| def _training_loss_with_context(self, **kwargs): |
| context_latents = kwargs.pop("context_latents", None) |
| num_context_frames = kwargs.pop("num_context_frames", 0) |
| models = {k: v for k, v in kwargs.items() if k in self.pipe.in_iteration_models} |
| inputs = {k: v for k, v in kwargs.items() if k not in self.pipe.in_iteration_models} |
| if context_latents is not None: |
| inputs.update({ |
| "context_latents": context_latents, |
| "num_context_frames": num_context_frames, |
| "context_noise_prob": self.context_noise_prob, |
| "context_noise_std": self.context_noise_std, |
| "context_attention_weight": getattr(self, "context_attention_weight", 1.0), |
| "use_anchor_frame": getattr(self, "use_anchor_frame", False), |
| "context_temporal_decay": getattr(self, "context_temporal_decay", 1.0), |
| "use_spatial_memory": getattr(self.pipe, "use_spatial_memory", False), |
| "spatial_memory_tokens": int(getattr(self.pipe, "spatial_memory_tokens", 64) or 64), |
| "use_spatial_memory_legacy": bool(getattr(self.pipe, "use_spatial_memory_legacy", False)), |
| "spatial_memory_module": getattr(self.pipe, "spatial_memory_module", None), |
| "spatial_memory_inject_mode": getattr(self.pipe, "spatial_memory_inject_mode", "concat_text"), |
| "spatial_memory_readout_module": getattr(self.pipe, "spatial_memory_readout_module", None), |
| "use_framepack_memory": bool(getattr(self, "use_framepack_memory", False)), |
| }) |
| if self.use_moc and self.moc_module is not None: |
| inputs["use_moc"] = True |
| inputs["moc_module"] = self.moc_module |
| if self.context_fixed_noise_std is not None: |
| inputs["context_fixed_noise_std"] = self.context_fixed_noise_std |
| inputs = self._ensure_input_latents(inputs, strict=True) |
| return self.pipe.training_loss(**models, **inputs) |
|
|
|
|