hugging-apps/echo-memory
0
1"""2推理工具:供 run_replay_loop_two_chunk 及评估脚本使用。3提供:load_pipeline_and_ckpt、load_prompt_for_video、sample_trajectory_samples_from_dataset。4 5VWM-style 简化版:使用 DiTBlock_w_Action + MLP_CamPose(block 内 action_mlp),6去除 CameraEncoder / camera_encoder_shallow 等冗余路径。7"""8 9import os10import re11import sys12import csv13import random14 15_script_dir = os.path.dirname(os.path.abspath(__file__))16_repo_root = os.path.dirname(os.path.dirname(_script_dir))17if _repo_root not in sys.path:18 sys.path.insert(0, _repo_root)19 20import torch21import torch.nn as nn22from safetensors.torch import load_file as safe_load_file23from diffsynth.pipelines.wan_video_new import WanVideoPipeline, ModelConfig24from diffsynth.models.wan_video_dit import SelfAttention, CrossAttention, GateModule, modulate25from diffsynth.models.memory.block_wise_ssm import BlockWiseStateSpaceMemory26from diffsynth.models.memory.videossm_hybrid import HybridStateSpaceMemory27 28DEFAULT_NEGATIVE_PROMPT = "oversaturated colors, overexposed, static, blurry details"29 30 31# ── MLP_CamPose + DiTBlock_w_Action(与训练侧 train.py 完全一致)──────────32class MLP_CamPose(nn.Module):33 def __init__(self, out_dim, pose_dim=12):34 super().__init__()35 self.proj = nn.Linear(pose_dim, out_dim)36 nn.init.zeros_(self.proj.weight)37 nn.init.zeros_(self.proj.bias)38 39 def forward(self, x):40 return self.proj(x)41 42 43class DiTBlock_w_Action(nn.Module):44 def __init__(self, has_image_input, dim, num_heads, ffn_dim, eps=1e-6,45 add_action_attn=False, action_use_temporal_attention=True,46 use_cam_pose=False, use_block_wise_ssm=False, use_videossm_hybrid=False,47 videossm_kernel_size=3, videossm_expand=2):48 super().__init__()49 self.dim = dim50 self.num_heads = num_heads51 self.ffn_dim = ffn_dim52 if add_action_attn:53 self.self_attn_with_action = SelfAttention(dim, num_heads, eps)54 nn.init.zeros_(self.self_attn_with_action.o.weight)55 nn.init.zeros_(self.self_attn_with_action.o.bias)56 if use_cam_pose:57 self.action_mlp = MLP_CamPose(dim)58 else:59 self.action_mlp = MLP_CamPose(dim)60 self.self_attn = SelfAttention(dim, num_heads, eps)61 self.cross_attn = CrossAttention(dim, num_heads, eps, has_image_input=has_image_input)62 self.norm1 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)63 self.norm2 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)64 self.norm3 = nn.LayerNorm(dim, eps=eps)65 self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'), nn.Linear(ffn_dim, dim))66 self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)67 self.gate = GateModule()68 self.action_use_temporal_attention = action_use_temporal_attention69 self.use_block_wise_ssm = bool(use_block_wise_ssm)70 self.use_videossm_hybrid = bool(use_videossm_hybrid)71 if use_block_wise_ssm:72 self.block_wise_ssm = BlockWiseStateSpaceMemory(dim)73 if use_videossm_hybrid:74 self.videossm_hybrid = HybridStateSpaceMemory(75 dim, kernel_size=videossm_kernel_size, expand=videossm_expand76 )77 78 def forward(self, x, context, t_mod, freqs, actions=None):79 has_seq = len(t_mod.shape) == 480 chunk_dim = 2 if has_seq else 181 shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (82 self.modulation.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod).chunk(6, dim=chunk_dim)83 if has_seq:84 shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (85 shift_msa.squeeze(2), scale_msa.squeeze(2), gate_msa.squeeze(2),86 shift_mlp.squeeze(2), scale_mlp.squeeze(2), gate_mlp.squeeze(2),87 )88 num_frames = None89 if actions is not None:90 original_x = x91 actions = self.action_mlp(actions.to(x.dtype)).to(x.dtype)92 bs, num_frames, dim = actions.shape93 actions = actions.reshape(bs, num_frames, 1, dim)94 x = x.reshape(bs, num_frames, -1, dim)95 x = x + actions96 if hasattr(self, "self_attn_with_action"):97 if not self.action_use_temporal_attention:98 x = x.reshape(bs, -1, dim)99 x = original_x + self.self_attn_with_action(x, freqs)100 else:101 from einops import rearrange102 x = rearrange(x, "b f p d -> (b p) f d")103 attn_out = self.self_attn_with_action(x)104 attn_out = rearrange(attn_out, "(b p) f d -> b f p d", b=bs)105 x = original_x + attn_out.reshape(bs, -1, dim)106 else:107 x = x.reshape(bs, -1, dim)108 input_x = modulate(self.norm1(x), shift_msa, scale_msa)109 x = self.gate(x, gate_msa, self.self_attn(input_x, freqs))110 if num_frames is not None:111 if hasattr(self, "block_wise_ssm"):112 x = self.block_wise_ssm(x, f=num_frames)113 if hasattr(self, "videossm_hybrid"):114 spatial = x.shape[1] // int(num_frames) if int(num_frames) > 0 else 0115 x = self.videossm_hybrid(x, f=num_frames, h=1, w=spatial)116 x = x + self.cross_attn(self.norm3(x), context)117 input_x = modulate(self.norm2(x), shift_mlp, scale_mlp)118 x = self.gate(x, gate_mlp, self.ffn(input_x))119 return x120 121 122# ── Utility functions ─────────────────────────────────────────────────────123 124def load_pose_rt(json_file, frame_idx):125 """从数据集 camera json 读取单帧 12 维 RT。"""126 from src.model_training.fov_retrieval import load_camera_pose, pose_to_rt127 pose = load_camera_pose(json_file, int(frame_idx))128 if pose is None:129 return None130 return pose_to_rt(pose, constrain_to_xy=True)131 132 133def get_relative_rt(rt, ref_rt):134 """单帧相对位姿。"""135 from src.model_training.fov_retrieval import convert_rt_to_relative136 if rt is None or ref_rt is None or len(rt) < 12 or len(ref_rt) < 12:137 return None138 out = convert_rt_to_relative([rt], ref_rt)139 return out[0] if out else None140 141 142def load_prompt_for_video(dataset_base, video_name):143 """从 dataset 目录下的 metadata CSV 读取该视频的 prompt。"""144 if not dataset_base or not video_name:145 return None146 vn = str(video_name).replace(".mp4", "").replace(".avi", "").strip()147 for name in ("metadata_full.csv", "metadata.csv", "prompts.csv"):148 path = os.path.join(dataset_base, name)149 if not os.path.isfile(path):150 continue151 try:152 with open(path, "r", encoding="utf-8") as f:153 for row in csv.DictReader(f):154 if row.get("video_name", "").strip() == vn:155 p = row.get("prompt", "").strip()156 if p:157 return p158 except Exception:159 pass160 return None161 162 163def sample_trajectory_samples_from_dataset(dataset_base, num_samples=4, num_frames=81, seed=42):164 """从 dataset 枚举 (video_name, start_frame)。"""165 frames_dir = os.path.join(dataset_base, "frames")166 if not os.path.isdir(frames_dir):167 return []168 candidates = []169 for vn in sorted(os.listdir(frames_dir)):170 vd = os.path.join(frames_dir, vn)171 if not os.path.isdir(vd):172 continue173 try:174 names = [f for f in os.listdir(vd) if f.endswith(".png")]175 indices = sorted({int(os.path.splitext(n)[0]) for n in names if n[:-4].isdigit()})176 if not indices:177 continue178 max_idx = max(indices)179 for start in indices:180 if start + num_frames - 1 <= max_idx:181 candidates.append((vn, start))182 except Exception:183 continue184 if not candidates:185 return []186 rng = random.Random(seed)187 if len(candidates) <= num_samples:188 return candidates189 return [candidates[i] for i in rng.sample(range(len(candidates)), num_samples)]190 191 192# ── Pipeline loading (VWM-style) ──────────────────────────────────────────193 194def _build_action_blocks(195 pipe,196 add_action_attn=False,197 action_use_temporal_attention=True,198 block_wise_block_ids=None,199 videossm_block_ids=None,200):201 """Replace DiT blocks with DiTBlock_w_Action (VWM cam_infer.py style)."""202 dit = pipe.dit203 old_blocks = dit.blocks204 has_image_input = getattr(dit, "has_image_input", False)205 dim = dit.dim206 num_heads = getattr(dit, "num_heads", None) or getattr(old_blocks[0], "num_heads", None)207 ffn_dim = getattr(dit, "ffn_dim", None) or getattr(old_blocks[0], "ffn_dim", None)208 eps = getattr(dit, "eps", 1e-6)209 210 block_dtype = next(old_blocks[0].parameters()).dtype211 block_device = next(old_blocks[0].parameters()).device212 213 block_wise_block_ids = set(block_wise_block_ids or [])214 videossm_block_ids = set(videossm_block_ids or [])215 216 new_blocks = torch.nn.ModuleList()217 for block_id, old_block in enumerate(old_blocks):218 new_block = DiTBlock_w_Action(219 has_image_input=has_image_input,220 dim=dim, num_heads=num_heads, ffn_dim=ffn_dim, eps=eps,221 add_action_attn=add_action_attn,222 action_use_temporal_attention=action_use_temporal_attention,223 use_cam_pose=True,224 use_block_wise_ssm=block_id in block_wise_block_ids,225 use_videossm_hybrid=block_id in videossm_block_ids,226 )227 new_block = new_block.to(dtype=block_dtype, device=block_device)228 for attr in ("self_attn", "cross_attn", "norm1", "norm2", "norm3", "ffn"):229 if hasattr(old_block, attr) and hasattr(new_block, attr):230 getattr(new_block, attr).load_state_dict(getattr(old_block, attr).state_dict())231 if hasattr(old_block, "modulation") and hasattr(new_block, "modulation"):232 with torch.no_grad():233 new_block.modulation.copy_(old_block.modulation.to(dtype=block_dtype))234 new_blocks.append(new_block)235 236 dit.blocks = new_blocks237 print(f"[loop_utils] Replaced {len(new_blocks)} blocks with DiTBlock_w_Action (MLP_CamPose)")238 if block_wise_block_ids:239 print(f"[loop_utils] Loaded Block-wise SSM slots on blocks: {sorted(block_wise_block_ids)[:8]}{'...' if len(block_wise_block_ids) > 8 else ''}")240 if videossm_block_ids:241 print(f"[loop_utils] Loaded VideoSSM hybrid slots on blocks: {sorted(videossm_block_ids)[:8]}{'...' if len(videossm_block_ids) > 8 else ''}")242 243 244def load_pipeline_and_ckpt(245 ckpt_path,246 dit_path,247 text_encoder_path,248 vae_path,249 device="cuda",250 add_action_attn=False,251 action_use_temporal_attention=True,252 tokenizer_path=None,253 # Legacy kwargs accepted but ignored (CameraEncoder removed)254 **kwargs,255):256 """Load WanVideoPipeline, replace blocks with DiTBlock_w_Action, load ckpt (strict=False).257 258 VWM-style: no CameraEncoder, no complex inference logic. Action is injected259 via MLP_CamPose (nn.Linear(12, dim), zero-init) inside each DiTBlock_w_Action.260 """261 print(f"[loop_utils] Loading pipeline (DiT -> {device})")262 if not tokenizer_path:263 import os as _os264 _base = _os.path.dirname(dit_path)265 _cand = _os.path.join(_base, "google", "umt5-xxl")266 if _os.path.isdir(_cand):267 tokenizer_path = _cand268 print(f"[loop_utils] Auto-detected tokenizer at {tokenizer_path}")269 model_configs = [270 ModelConfig(path=dit_path, offload_device=device),271 ModelConfig(path=text_encoder_path, offload_device="cpu"),272 ModelConfig(path=vae_path, offload_device="cpu"),273 ]274 pipe = WanVideoPipeline.from_pretrained(275 torch_dtype=torch.bfloat16,276 device=device,277 model_configs=model_configs,278 tokenizer_config=ModelConfig(path=tokenizer_path) if tokenizer_path else None,279 )280 281 ckpt = None282 block_wise_block_ids = set()283 videossm_block_ids = set()284 action_attn_block_ids = set()285 if ckpt_path and os.path.isfile(ckpt_path):286 ckpt = safe_load_file(ckpt_path)287 for key in ckpt.keys():288 m = re.match(r"blocks\.(\d+)\.block_wise_ssm\.", key)289 if m:290 block_wise_block_ids.add(int(m.group(1)))291 m = re.match(r"blocks\.(\d+)\.videossm_hybrid\.", key)292 if m:293 videossm_block_ids.add(int(m.group(1)))294 m = re.match(r"blocks\.(\d+)\.self_attn_with_action\.", key)295 if m:296 action_attn_block_ids.add(int(m.group(1)))297 298 if action_attn_block_ids and not add_action_attn:299 add_action_attn = True300 print("[loop_utils] Detected action-attention weights in checkpoint; enabling self_attn_with_action")301 302 # Replace blocks with DiTBlock_w_Action, including memory slots implied by ckpt keys.303 _build_action_blocks(304 pipe,305 add_action_attn=add_action_attn,306 action_use_temporal_attention=action_use_temporal_attention,307 block_wise_block_ids=block_wise_block_ids,308 videossm_block_ids=videossm_block_ids,309 )310 311 # Load ckpt (strict=False: base model keys match, action_mlp keys are extra)312 if ckpt_path and not os.path.isfile(ckpt_path):313 print(f"[loop_utils] WARNING: ckpt not found: {ckpt_path} — running with base model weights only!")314 if ckpt_path and os.path.isfile(ckpt_path):315 if ckpt is None:316 ckpt = safe_load_file(ckpt_path)317 missing, unexpected = pipe.dit.load_state_dict(ckpt, strict=False)318 action_keys = [k for k in ckpt if "action_mlp" in k]319 if not missing and not unexpected:320 print(f"[loop_utils] Ckpt loaded: {len(ckpt)} keys, perfect match")321 else:322 print(f"[loop_utils] Ckpt loaded: {len(ckpt)} keys, "323 f"missing={len(missing)}, unexpected={len(unexpected)}, "324 f"action_mlp_keys={len(action_keys)}")325 if missing:326 for k in sorted(missing)[:5]:327 print(f" missing: {k}")328 if unexpected:329 for k in sorted(unexpected)[:5]:330 print(f" unexpected: {k}")331 332 # Optional: load SpatialGridMemory if present in ckpt333 _smsd = {334 k.replace("spatial_memory_module.", "", 1): v335 for k, v in ckpt.items()336 if k.startswith("spatial_memory_module.")337 }338 if _smsd:339 try:340 from diffsynth.models.memory.spatial_grid_memory import SpatialGridMemory341 except ImportError:342 SpatialGridMemory = None343 if SpatialGridMemory is not None:344 dim = pipe.dit.dim345 w = _smsd.get("spatial_to_tokens")346 if w is not None:347 g2, num_tok = int(w.shape[0]), int(w.shape[1])348 gsz = int(round(g2 ** 0.5))349 if gsz * gsz != g2:350 gsz = 8351 sm = SpatialGridMemory(dim, grid_size=gsz, num_tokens=num_tok)352 sm.load_state_dict(_smsd, strict=False)353 sm = sm.to(dtype=next(pipe.dit.parameters()).dtype, device=next(pipe.dit.parameters()).device)354 pipe.spatial_memory_module = sm355 pipe.use_spatial_memory_legacy = False356 print(f"[loop_utils] Loaded spatial_memory_module (grid={gsz}, tokens={num_tok})")357 358 _srmsd = {359 k.replace("spatial_memory_readout_module.", "", 1): v360 for k, v in ckpt.items()361 if k.startswith("spatial_memory_readout_module.")362 }363 if _srmsd:364 try:365 from diffsynth.models.memory.spatial_grid_memory import SpatialCrossAttnReadout366 except ImportError:367 SpatialCrossAttnReadout = None368 if SpatialCrossAttnReadout is not None:369 dim = pipe.dit.dim370 readout = SpatialCrossAttnReadout(dim=dim, num_heads=8)371 readout.load_state_dict(_srmsd, strict=False)372 readout = readout.to(dtype=next(pipe.dit.parameters()).dtype, device=next(pipe.dit.parameters()).device)373 pipe.spatial_memory_readout_module = readout374 print("[loop_utils] Loaded spatial_memory_readout_module")375 376 if getattr(pipe, "enable_vram_management", None):377 pipe.enable_vram_management()378 return pipe379 