Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
loop_utils.py379 linesDownload Raw Back to env
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