Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
run_replay_loop_two_chunk.py1946 linesDownload Raw Back to env
1#!/usr/bin/env python32"""3回环:两栏输出 [输入轨迹 | 生成视频]。context 一律放在右侧(suffix),I2V 首帧也放在最右作为 clean 去噪 target。4 5- 第一帧:从训练数据随机采样一帧作为首 chunk 的 context(I2V,1 帧放右侧)6- 之后每 chunk 的第一帧(context)= 上一 chunk 的 last n 帧,同样放在该段生成视频的右侧7- 右栏视频 = 每段 [生成 81 帧 | context 帧],context 始终在右8 9默认回环两场景(每 chunk 45°):1) 1左1右 2chunk  2) 2左2右 4chunk。--run_legacy_loop 时跑旧三组(roundtrip/onedir/replay)。多卡按 rank 划分样本。10 11多 chunk 时 RT(相对位姿)与训练一致:12- 训练:ref_rt = rt_list[0],target_actions = convert_rt_to_relative(rt_list, ref_rt),即每段内 action 相对本段首帧。13- 推理:每个 chunk 的 action 也是相对该 chunk 的首帧。chunk1 注入 0°→45° 表示本段内转 45°;14  chunk2 的首帧 = chunk1 的末帧(已 45°),chunk2 再注入 0°→45° 表示在 chunk2 局部再转 45°,世界系共 90°。15  因此是 45+45,不是第二个 chunk 直接 90°(直接 90° 表示相对 chunk2 首帧转 90°,世界系会变成 45+90=135°)。16 17注意:条件以 action + context 为主;prompt 使用 GT 帧对应视频的文案以补充场景描述,利于生成清晰度。18"""19 20import os21import sys22import json23import argparse24import random25import math26import io27from datetime import datetime28import torch29import numpy as np30from PIL import Image31 32_script_dir = os.path.dirname(os.path.abspath(__file__))33_repo_root = os.path.dirname(os.path.dirname(_script_dir))34_exp1_4_2_dir = os.path.join(_repo_root, "ab_study", "exp1_4_2_context_suffix_cam_rt_relative")35if _repo_root not in sys.path:36    sys.path.insert(0, _repo_root)37if _exp1_4_2_dir not in sys.path:38    sys.path.insert(0, _exp1_4_2_dir)39if _script_dir not in sys.path:40    sys.path.insert(0, _script_dir)41 42import loop_utils as irc43from diffsynth import save_video44 45from src.model_training.fov_retrieval import compute_rotation_list46from src.model_training.multichunk_sample_utils import (47    context_frames_for_next_chunk,48    replay_context_from_generated_frames,49)50 51 52def encode_context_frames_per_frame(pipe, pil_list, device, dtype=torch.bfloat16):53    """与训练 context_per_frame_vae 一致:每帧单独 VAE 编码再在时间维 concat,不做时序降采样。ctx=K → K 个 latent tokens(5/20 等均一帧一过 VAE)。"""54    if not pil_list:55        return None56    encoded = []57    for pil in pil_list:58        frame_video = pipe.preprocess_video([pil]).to(device=device)59        frame_sq = frame_video.squeeze(0) if frame_video.dim() == 5 else frame_video60        if frame_sq.dim() == 3:61            frame_sq = frame_sq.unsqueeze(0)62        lat_one = pipe.vae.encode([frame_sq], device=pipe.device, tiled=False, tile_size=None, tile_stride=None)63        encoded.append(lat_one)64    context_latents = torch.cat(encoded, dim=2).to(dtype=dtype, device=device)65    return context_latents66 67 68def _yaw_deg_from_rt(rt_list):69    """从 12 维 RT [t_x,t_y,t_z, R11,R12,...,R33] 提取 yaw(度)。R_z 时 yaw = atan2(R21, R11)."""70    if not rt_list or len(rt_list) < 12:71        return 0.072    R11, R21 = float(rt_list[3]), float(rt_list[6])73    return math.degrees(math.atan2(R21, R11))74 75 76def build_action_ccw_cw(deg: float, chunk_frames: int = 81):77    """R_z(yaw),相对本段首帧。CCW=0→+deg,CW=0→-deg。匀速。"""78    denom = max(1, chunk_frames - 1)79    actions_ccw = {}80    for i in range(chunk_frames):81        yaw = (i / denom) * deg82        actions_ccw[str(i)] = compute_rotation_list([0.0, 0.0, 0.0, yaw])83    actions_cw = {}84    for i in range(chunk_frames):85        yaw = -(i / denom) * deg86        actions_cw[str(i)] = compute_rotation_list([0.0, 0.0, 0.0, yaw])87    return actions_ccw, actions_cw88 89 90def build_action_cw_then_ccw(deg: float, chunk_frames: int = 81):91    """匀速:chunk1 顺时针 0→-deg,chunk2 逆时针 0→+deg(本段局部)。折返:先 CW 30° 再 CCW 30° 回起点。"""92    denom = max(1, chunk_frames - 1)93    ch1 = {}94    for i in range(chunk_frames):95        ch1[str(i)] = compute_rotation_list([0.0, 0.0, 0.0, -(i / denom) * deg])96    ch2 = {}97    for i in range(chunk_frames):98        ch2[str(i)] = compute_rotation_list([0.0, 0.0, 0.0, (i / denom) * deg])99    return ch1, ch2100 101 102def build_gt_trajectory_actions(dataset_base, video_name, start_frame, chunk_frames, json_file=None):103    """与 Replay 完全一致:从数据集 json 取同段 GT pose,转成 relative RT 作为 action。若帧不足则返回 None。"""104    if json_file is None:105        json_file = os.path.join(dataset_base, "jsons", f"{video_name}.json")106    if not os.path.isfile(json_file):107        return None108    try:109        rt_list = [irc.load_pose_rt(json_file, start_frame + i) for i in range(chunk_frames)]110        if not rt_list or any(r is None or len(r) < 12 for r in rt_list):111            return None112        ref_rt = rt_list[0]113        rel_actions = {str(i): irc.get_relative_rt(rt_list[i], ref_rt) for i in range(chunk_frames)}114        return rel_actions115    except Exception:116        return None117 118 119def build_action_from_gt_yaw_profile(dataset_base, video_name, start_frame, chunk_frames, target_total_deg, clockwise=False, json_file=None):120    """参考训练集方式:用同段 GT 轨迹的 yaw 曲线形状(时间分布)缩放为目标角度,使旋转节奏与训练集一致,而非简单线性。121    若该段无 GT 或总转角过小则返回 None,调用方回退线性合成。"""122    if json_file is None:123        json_file = os.path.join(dataset_base, "jsons", f"{video_name}.json")124    if not os.path.isfile(json_file):125        return None126    try:127        rt_list = [irc.load_pose_rt(json_file, start_frame + i) for i in range(chunk_frames)]128        if not rt_list or any(r is None or len(r) < 12 for r in rt_list):129            return None130        ref_rt = rt_list[0]131        # GT 相对首帧的 yaw 曲线(与训练一致)132        yaw_deg = [_yaw_deg_from_rt(irc.get_relative_rt(rt_list[i], ref_rt)) for i in range(chunk_frames)]133        total_gt_deg = yaw_deg[-1] - yaw_deg[0] if len(yaw_deg) > 1 else 0.0134        if abs(total_gt_deg) < 1e-4:135            return None  # 几乎无旋转,用线性更稳136        # 归一化曲线 t[i] in [0,1],再缩放到 target_total_deg,保持 GT 的时间节奏137        scale = target_total_deg / total_gt_deg138        sign = -1 if clockwise else 1139        actions = {}140        for i in range(chunk_frames):141            t = (yaw_deg[i] - yaw_deg[0]) / total_gt_deg  # 0 -> 1142            yaw = sign * t * target_total_deg143            actions[str(i)] = compute_rotation_list([0.0, 0.0, 0.0, yaw])144        return actions145    except Exception:146        return None147 148 149def build_action_chunk(deg: float, clockwise: bool, chunk_frames: int = 81):150    """单 chunk:0→deg(逆时针)或 0→-deg(顺时针)。相对本 chunk 首帧,与训练一致。151    多 chunk 时:chunk2 的首帧=chunk1 的末帧,所以 chunk2 注入 0→deg 表示在 chunk2 局部再转 deg;152    世界系转角 = chunk1 末帧 yaw + deg。例如 chunk1=45°,chunk2 再 45° 则注入 0→45°(不是 0→90°)。153    返回 (actions_dict, yaw_list_deg)."""154    denom = max(1, chunk_frames - 1)155    actions = {}156    yaw_list = []157    for i in range(chunk_frames):158        yaw = (i / denom) * (-deg if clockwise else deg)159        yaw_list.append(yaw)160        actions[str(i)] = compute_rotation_list([0.0, 0.0, 0.0, yaw])161    return actions, yaw_list162 163 164def build_action_translation_only(direction: str, translation_delta: float, chunk_frames: int = 81):165    """纯平移单 chunk:R=I,沿单轴线性位移。与训练一致:fov_retrieval.pose_to_rt(..., constrain_to_xy=True) 仅用 XY,tz=0。166    direction in ('forward','backward','left','right'):forward=+Y, backward=-Y, left=-X, right=+X(XY 平面,无 Z)。167    RT 格式:前 3 维 [tx,ty,tz],后 9 维 3x3 行优先单位阵;相对本段首帧(与 convert_rt_to_relative(ref=首帧) 一致)。"""168    # 与训练一致:2D 平面位移,z 恒为 0(见 fov_retrieval.pose_to_rt constrain_to_xy=True)169    identity_rot = [1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0]170    denom = max(1, chunk_frames - 1)171    actions = {}172    yaw_list = [0.0] * chunk_frames173    for i in range(chunk_frames):174        t = (i / denom) * translation_delta175        if direction == "forward":176            tx, ty, tz = 0.0, t, 0.0177        elif direction == "backward":178            tx, ty, tz = 0.0, -t, 0.0179        elif direction == "left":180            tx, ty, tz = -t, 0.0, 0.0181        elif direction == "right":182            tx, ty, tz = t, 0.0, 0.0183        else:184            tx, ty, tz = 0.0, 0.0, 0.0185        actions[str(i)] = [tx, ty, tz] + identity_rot186    return actions, yaw_list187 188 189def _draw_trajectory_pil(yaw_list, plot_width, plot_height, frame_index):190    """PIL-only fallback: draw trajectory 0..frame_index. Returns one PIL Image."""191    pad = 40192    w, h = plot_width - 2 * pad, plot_height - 2 * pad193    if w <= 0 or h <= 0:194        return Image.new("RGB", (plot_width, plot_height), (50, 50, 50))195    img = Image.new("RGB", (plot_width, plot_height), (45, 45, 48))196    end = min(frame_index + 1, len(yaw_list))197    if end == 0:198        return img199    ys = [float(yaw_list[i]) for i in range(end)]200    y_min, y_max = min(ys), max(ys)201    if y_max <= y_min:202        y_min, y_max = y_min - 10.0, y_max + 10.0203    from PIL import ImageDraw204    draw = ImageDraw.Draw(img)205    n_x = max(1, len(yaw_list))206    pts = []207    for i in range(end):208        x = pad + int((i / max(1, n_x - 1)) * w) if n_x > 1 else pad209        yy = (ys[i] - y_min) / max(1e-6, y_max - y_min)210        y = pad + int((1 - yy) * h)211        pts.append((x, y))212    if len(pts) >= 2:213        draw.line(pts, fill=(0, 200, 255), width=2)214    if pts:215        draw.ellipse([pts[-1][0] - 4, pts[-1][1] - 4, pts[-1][0] + 4, pts[-1][1] + 4], fill=(255, 220, 0), outline=(255, 255, 255))216    return img217 218 219def draw_trajectory_frames(yaw_history_deg, plot_width, plot_height, total_frames=None):220    """221    yaw_history_deg: list of cumulative yaw (one per frame).222    Returns list of PIL images (one per frame), each showing trajectory from 0 to current frame.223    Uses matplotlib; if output is too dark, falls back to PIL-only drawing.224    """225    yaw_list = [float(y) for y in (yaw_history_deg or [])]226    n = len(yaw_list)227    if total_frames is not None and total_frames > n:228        n = total_frames229    if n == 0:230        return [Image.new("RGB", (plot_width, plot_height), (45, 45, 48))]231 232    use_pil_fallback = False233    try:234        import matplotlib235        matplotlib.use("Agg")236        import matplotlib.pyplot as plt237    except ImportError:238        use_pil_fallback = True239 240    if not use_pil_fallback:241        y_min = min(yaw_list) if yaw_list else 0.0242        y_max = max(yaw_list) if yaw_list else 0.0243        if y_max <= y_min:244            y_min, y_max = y_min - 15.0, y_max + 15.0245        margin = max(5, (y_max - y_min) * 0.15)246        y_lo, y_hi = y_min - margin, y_max + margin247        out = []248        for t in range(n):249            fig, ax = plt.subplots(1, 1, figsize=(plot_width / 100.0, plot_height / 100.0), dpi=100)250            fig.patch.set_facecolor("#252525")251            ax.set_facecolor("#252525")252            ax.set_xlim(0, max(n, 1))253            ax.set_ylim(y_lo, y_hi)254            ax.tick_params(colors="0.9", labelsize=7)255            for spine in ax.spines.values():256                spine.set_color("0.6")257            end = min(t + 1, len(yaw_list))258            if end > 0:259                x = list(range(end))260                y = [yaw_list[i] for i in range(end)]261                ax.plot(x, y, color="cyan", linewidth=2.0, zorder=2)262                cur_y = yaw_list[min(t, len(yaw_list) - 1)] if t < len(yaw_list) else (yaw_list[-1] if yaw_list else 0)263                ax.scatter([t], [cur_y], color="yellow", s=35, zorder=3)264            ax.set_xlabel("Frame", color="0.9", fontsize=8)265            ax.set_ylabel("Yaw (deg)", color="0.9", fontsize=8)266            ax.set_title("Trajectory (yaw)", color="0.95", fontsize=9)267            fig.tight_layout(pad=0.5)268            buf = io.BytesIO()269            plt.savefig(buf, format="png", dpi=100, facecolor="#252525", edgecolor="none", bbox_inches="tight", pad_inches=0.2)270            plt.close(fig)271            buf.seek(0)272            img = Image.open(buf).convert("RGB")273            if img.size != (plot_width, plot_height):274                img = img.resize((plot_width, plot_height), Image.Resampling.LANCZOS)275            # If image is too dark (matplotlib failed in headless?), use PIL fallback for rest276            arr = np.array(img)277            if arr.mean() < 50:278                use_pil_fallback = True279                out = [_draw_trajectory_pil(yaw_list, plot_width, plot_height, t) for t in range(n)]280                break281            out.append(img)282        if not use_pil_fallback:283            return out284    if use_pil_fallback:285        return [_draw_trajectory_pil(yaw_list, plot_width, plot_height, t) for t in range(n)]286    return out287 288 289def composite_two_panel(trajectory_frames, right_panel_frames, w_traj=320, h_traj=352, w_gen=640, h_gen=352):290    """左:输入轨迹可视化;右:生成视频(每段后接 context,context 在右)。"""291    n = len(right_panel_frames)292    if not trajectory_frames or len(trajectory_frames) != n:293        traj_frames = draw_trajectory_frames([], w_traj, h_traj, total_frames=n)294    else:295        traj_frames = trajectory_frames296    out = []297    for i in range(n):298        traj = traj_frames[i] if i < len(traj_frames) else traj_frames[-1]299        if traj.size != (w_traj, h_traj):300            traj = traj.resize((w_traj, h_traj))301        gen = _frame_to_pil(right_panel_frames[i], w_gen, h_gen)302        canvas_w = w_traj + w_gen303        canvas_h = max(h_traj, h_gen)304        canvas = Image.new("RGB", (canvas_w, canvas_h), (40, 40, 40))305        canvas.paste(traj, (0, 0))306        canvas.paste(gen, (w_traj, 0))307        out.append(canvas)308    return out309 310 311def load_sample_first_frame(dataset_base, video_name, start_frame, w, h):312    """加载当前样本的首帧(与 Replay 一致:context = 同条轨迹同视频的首帧),避免 loop 用随机帧导致 prompt 与 context 场景不一致而糊。返回 PIL 或 None。"""313    if not video_name or start_frame is None:314        return None315    for fmt in (f"{int(start_frame):04d}.png", f"{int(start_frame)}.png"):316        img_p = os.path.join(dataset_base, "frames", str(video_name), fmt)317        if os.path.isfile(img_p):318            return Image.open(img_p).convert("RGB").resize((w, h))319    return None320 321 322def load_gt_frames_at_indices(dataset_base, video_name, frame_indices, w, h):323    """从数据集加载指定帧序号的 GT 图像,用于 chunk 间衔接的 context(与训练一致:context=真实帧,避免用生成帧再编码导致糊)。324    frame_indices: 从左到右使用(如 [80, 79] 表示最后一帧、倒数第二帧)。返回 list[PIL] 若全部存在,否则 None。"""325    if not video_name or not frame_indices:326        return None327    out = []328    base = os.path.join(dataset_base, "frames", str(video_name))329    for idx in frame_indices:330        for fmt in (f"{int(idx):04d}.png", f"{int(idx)}.png"):331            p = os.path.join(base, fmt)332            if os.path.isfile(p):333                out.append(Image.open(p).convert("RGB").resize((w, h)))334                break335        else:336            return None337    return out338 339 340def seed_for_sample(base_seed: int, video_name, start_frame) -> int:341    """与 (video_name, start_frame) 一一对应的确定性 seed,使同一样本在 loop_traj_fov_gen 与 evals_ep0 等不同调用下结果一致(避免因 idx/rank 不同导致偏移)。"""342    try:343        h = hash((str(video_name), int(start_frame)))344        return base_seed + (h & 0x7FFFFFFF) % 100000345    except Exception:346        return base_seed347 348 349def sample_random_frame_from_dataset(dataset_base, w, h, seed, metadata_path=None):350    """从训练数据中随机采样一帧(用于首 chunk 的 I2V context,放右侧)。返回 PIL 或 None。"""351    import csv352    candidates = []353    if metadata_path and os.path.isfile(metadata_path):354        try:355            with open(metadata_path, "r", encoding="utf-8") as f:356                reader = csv.DictReader(f)357                for row in reader:358                    vn = row.get("video_name", "").strip()359                    sf = row.get("start_frame", "")360                    if vn and str(sf).strip():361                        try:362                            candidates.append((vn, int(sf)))363                        except ValueError:364                            continue365        except Exception:366            pass367    if not candidates:368        frames_dir = os.path.join(dataset_base, "frames")369        if os.path.isdir(frames_dir):370            for vn in sorted(os.listdir(frames_dir)):371                vd = os.path.join(frames_dir, vn)372                if not os.path.isdir(vd):373                    continue374                pngs = [f for f in os.listdir(vd) if f.endswith(".png")]375                if not pngs:376                    continue377                for p in pngs[:3]:378                    try:379                        frame_idx = int(os.path.splitext(p)[0])380                        candidates.append((vn, frame_idx))381                    except ValueError:382                        continue383    if not candidates:384        return None385    rng = random.Random(seed)386    vn, frame_idx = rng.choice(candidates)387    img_p = os.path.join(dataset_base, "frames", vn, f"{frame_idx:04d}.png")388    if not os.path.isfile(img_p):389        img_p = os.path.join(dataset_base, "frames", vn, f"{frame_idx}.png")390    if os.path.isfile(img_p):391        return Image.open(img_p).convert("RGB").resize((w, h))392    return None393 394 395def _angle_distance_deg(a: float, b: float) -> float:396    """Smallest absolute yaw distance in degrees."""397    return abs((float(a) - float(b) + 180.0) % 360.0 - 180.0)398 399 400def fov_history_context_from_generated_frames(frames_list, yaw_list, K: int, target_yaw: float):401    """Select replay context by yaw/FOV proxy from generated history.402 403    The first context frame remains the most recent frame for short-term continuity.404    The remaining frames are selected from generated history by closest world-yaw405    distance to the next chunk's midpoint yaw. Context actions are handled by the406    caller and intentionally remain unchanged in this first version.407    """408    n = min(len(frames_list), len(yaw_list))409    if n <= 0 or K <= 0:410        return [], []411    if K == 1 or n == 1:412        return [frames_list[n - 1]], [n - 1]413 414    forced_idx = n - 1415    need = min(int(K) - 1, n - 1)416    scored = []417    for idx in range(n - 1):418        dist = _angle_distance_deg(yaw_list[idx], target_yaw)419        # Prefer recent frames as a tie-breaker while keeping FOV/yaw match primary.420        scored.append((dist, -idx, idx))421    scored.sort()422    selected = [idx for _dist, _neg_idx, idx in scored[:need]]423    return [frames_list[forced_idx]] + [frames_list[i] for i in selected], [forced_idx] + selected424 425 426def _world_yaw_rt(yaw: float):427    return compute_rotation_list([0.0, 0.0, 0.0, float(yaw)])428 429 430def context_actions_from_world_yaws(selected_yaws, ref_yaw: float):431    """Convert selected world-yaw RTs into RTs relative to the next chunk start."""432    from src.model_training.fov_retrieval import convert_rt_to_relative433 434    ref_rt = _world_yaw_rt(ref_yaw)435    selected_rts = [_world_yaw_rt(y) for y in selected_yaws]436    return convert_rt_to_relative(selected_rts, ref_rt)437 438 439def trim_continuation_first_frame(chunk_frames_list, yaw_history, chunk_frames=81):440    """与训练对齐:训练时 context[0]=target[0](同一帧),续段首帧=上一段末帧。拼接时去掉续段第 0 帧避免重复/跳帧。441    返回 (trimmed_chunk_frames_list, trimmed_yaw):chunk0 保留 81 帧,chunk1 及之后只保留 [1:81] 共 80 帧。"""442    if not chunk_frames_list or len(chunk_frames_list) <= 1:443        return chunk_frames_list, yaw_history444    trimmed_chunks = [list(chunk_frames_list[0])]445    trimmed_yaw = []446    offset = 0447    for i, gen_frames in enumerate(chunk_frames_list):448        n_use = len(gen_frames) if i == 0 else max(0, len(gen_frames) - 1)449        if i == 0:450            trimmed_yaw.extend(yaw_history[offset : offset + n_use])451        else:452            trimmed_yaw.extend(yaw_history[offset + 1 : offset + 1 + n_use])453        offset += len(gen_frames)454        if i > 0 and len(gen_frames) > 1:455            trimmed_chunks.append(list(gen_frames[1:]))456        elif i > 0:457            trimmed_chunks.append([])458    return trimmed_chunks, trimmed_yaw459 460 461def build_right_panel_and_yaw(chunk_frames_list, yaw_history, context_frames_per_chunk, chunk_frames=81, w=640, h=352, frame_to_pil_fn=None):462    """从每段生成帧 + 每段后的 context 帧拼出右栏;并扩展 yaw 与右栏帧数一致(context 帧沿用上一 yaw)。463    支持变长 chunk(如 trim 后首段 81、续段 80),按 yaw_history 顺序逐帧对齐。"""464    fp = frame_to_pil_fn or _frame_to_pil465    right_frames = []466    yaw_extended = []467    yaw_idx = 0468    for i, gen_frames in enumerate(chunk_frames_list):469        for j, f in enumerate(gen_frames):470            right_frames.append(f)471            yaw_extended.append(yaw_history[yaw_idx] if yaw_idx < len(yaw_history) else (yaw_history[-1] if yaw_history else 0.0))472            yaw_idx += 1473        ctx_list = context_frames_per_chunk[i] if i < len(context_frames_per_chunk) else []474        last_yaw = yaw_extended[-1] if yaw_extended else 0.0475        for ctx in ctx_list:476            right_frames.append(ctx if isinstance(ctx, Image.Image) else fp(ctx, w, h))477            yaw_extended.append(last_yaw)478    return right_frames, yaw_extended479 480 481def _frame_to_pil(f, tw, th):482    if hasattr(f, "convert") and hasattr(f, "resize"):483        return f.convert("RGB").resize((tw, th))484    if isinstance(f, np.ndarray):485        if f.dtype != np.uint8:486            f = (f * 255).astype(np.uint8) if f.max() <= 1.0 else f.astype(np.uint8)487        return Image.fromarray(f).convert("RGB").resize((tw, th))488    if isinstance(f, torch.Tensor):489        fn = f.cpu().numpy()490        if len(fn.shape) == 3 and fn.shape[0] == 3:491            fn = fn.transpose(1, 2, 0)492        fn = (fn * 255).clip(0, 255).astype(np.uint8) if fn.max() <= 1.0 else fn.clip(0, 255).astype(np.uint8)493        return Image.fromarray(fn).convert("RGB").resize((tw, th))494    return f495 496 497def run_one_chunk(498    pipe,499    prompt,500    use_negative_prompt,501    action_path=None,502    cam_pose_actions=None,503    context_latents=None,504    num_context_frames=1,505    context_actions_t=None,506    chunk_frames=81,507    h=352,508    w=640,509    seed=0,510    sigma_shift=15.0,511    num_inference_steps=50,512    cfg_scale=5.0,513    inference_noise_level=0.0,514    omit_context_actions=False,  # kept for backward compat, no longer used515    log_prefix="[Loop]",516    **_extra,517):518    """Generate one chunk. VWM-aligned action injection via cam_pose_actions (latent-frame-aligned 12-D RT)."""519    device = pipe.device520    kwargs_common = dict(521        prompt=prompt,522        negative_prompt=use_negative_prompt,523        height=h, width=w, num_frames=chunk_frames,524        num_inference_steps=num_inference_steps,525        seed=seed,526        cfg_scale=cfg_scale,527        sigma_shift=sigma_shift,528        denoising_strength=1.0,529    )530    if action_path is not None:531        kwargs_common["action_path"] = action_path532    elif cam_pose_actions is not None:533        kwargs_common["cam_pose_actions"] = cam_pose_actions534 535    has_action_mlp = hasattr(pipe.dit.blocks[0], 'action_mlp') if len(pipe.dit.blocks) > 0 else False536    if (action_path or cam_pose_actions is not None) and not has_action_mlp:537        print(f"{log_prefix} 警告: dit.blocks[0] 无 action_mlp,action 可能未注入(请确认 ckpt 含 action_mlp 权重)", flush=True)538    print(f"{log_prefix} run_one_chunk sigma_shift={sigma_shift} steps={num_inference_steps} (ctx={num_context_frames}) inference_noise={inference_noise_level} context_mem={context_latents is not None} action_mlp={has_action_mlp}")539    if context_latents is not None:540        pipe_kw = dict(541            **kwargs_common,542            enable_context_memory=True,543            context_latents=context_latents,544            num_context_frames=num_context_frames,545            context_position="suffix",546            cfg_target_only=True,547            inference_noise_level=inference_noise_level,548        )549        if context_actions_t is not None:550            pipe_kw["context_actions"] = context_actions_t551        with torch.no_grad():552            vid = pipe(**pipe_kw)553    else:554        with torch.no_grad():555            vid = pipe(**kwargs_common, enable_context_memory=False)556    return vid if isinstance(vid, list) else [vid]557 558 559def run_roundtrip_2chunk(560    pipe,561    dataset_base,562    output_dir,563    video_name,564    start_frame,565    deg=30.0,566    chunk_frames=81,567    context_frames=1,568    h=352,569    w=640,570    sigma_shift=15.0,571    num_inference_steps=50,572    cfg_scale=5.0,573    seed=42,574    inference_noise_level=0.0,575    metadata_path=None,576    keep_action_jsons=False,577    omit_context_actions=True,578):579    """2 chunk 折返:匀速 顺时针 deg° 再 逆时针 deg°(回起点)。"""580    print(f"[Roundtrip] 匀速 顺时针{deg}° then 逆时针{deg}° sigma_shift={sigma_shift} omit_ctx_act={omit_context_actions}")581    json_file = os.path.join(dataset_base, "jsons", f"{video_name}.json")582    prompt = irc.load_prompt_for_video(dataset_base, video_name) or "A scene."583    use_negative_prompt = getattr(irc, "DEFAULT_NEGATIVE_PROMPT", "oversaturated colors, overexposed, static, blurry details")584 585    actions_ch1, actions_ch2 = build_action_cw_then_ccw(deg, chunk_frames)586    subdir = os.path.join(output_dir, f"loop_{video_name}_start{start_frame}")587    os.makedirs(subdir, exist_ok=True)588    path_ch1 = os.path.join(subdir, "_ch1_cw.json")589    path_ch2 = os.path.join(subdir, "_ch2_ccw.json")590    with open(path_ch1, "w") as f:591        json.dump(actions_ch1, f, indent=2)592    with open(path_ch2, "w") as f:593        json.dump(actions_ch2, f, indent=2)594    print(f"[Roundtrip] Chunk1 CW 0->-{deg}° Chunk2 CCW 0->+{deg}° (匀速)")595 596    denom = max(1, chunk_frames - 1)597    yaw_history = [-(i / denom) * deg for i in range(chunk_frames)]  # chunk1: 0 -> -deg598    last_y1 = -deg599    yaw_history += [last_y1 + (i / denom) * deg for i in range(chunk_frames)]  # chunk2: -deg -> 0600 601    chunk_frames_list = []602    context_frames_per_chunk = []603 604    # Chunk 1: context = 当前样本首帧(与 Replay 一致,同视频同场景),避免随机帧导致糊;若无则回退随机帧605    ctx_pil_0 = load_sample_first_frame(dataset_base, video_name, start_frame, w, h)606    used_sample_first = ctx_pil_0 is not None607    if ctx_pil_0 is None:608        ctx_pil_0 = sample_random_frame_from_dataset(dataset_base, w, h, seed, metadata_path)609    if ctx_pil_0 is not None:610        print(f"[Roundtrip] Chunk1 context: {'sample first frame ' + str((video_name, start_frame)) if used_sample_first else 'random frame (fallback)'}")611    if ctx_pil_0 is not None:612        ctx_pil_0 = [ctx_pil_0]613        pipe.load_models_to_device(["vae"])614        with torch.no_grad():615            ctx_latents_0 = encode_context_frames_per_frame(pipe, ctx_pil_0, pipe.device)616        identity_rt = [0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0]617        ctx_actions_t_0 = torch.tensor([identity_rt], dtype=torch.float32)618        frames_ch1 = run_one_chunk(619            pipe, prompt, use_negative_prompt, path_ch1,620            context_latents=ctx_latents_0,621            num_context_frames=ctx_latents_0.shape[2],622            context_actions_t=ctx_actions_t_0,623            chunk_frames=chunk_frames, h=h, w=w, seed=seed,624            sigma_shift=sigma_shift, num_inference_steps=num_inference_steps,625            cfg_scale=cfg_scale, inference_noise_level=inference_noise_level,626            omit_context_actions=omit_context_actions,627        )628        context_frames_per_chunk.append([ctx_pil_0[0]])629    else:630        frames_ch1 = run_one_chunk(631            pipe, prompt, use_negative_prompt, path_ch1,632            chunk_frames=chunk_frames, h=h, w=w, seed=seed,633            sigma_shift=sigma_shift, num_inference_steps=num_inference_steps, cfg_scale=cfg_scale,634            omit_context_actions=omit_context_actions,635        )636        context_frames_per_chunk.append([_frame_to_pil(frames_ch1[-1], w, h)])637    chunk_frames_list.append(frames_ch1)638 639    # Chunk 2: condition = 上一 chunk,顺序与训练一致 [last_frame, ctx1, ...] 紧挨噪声后640    n_ctx = min(context_frames, len(frames_ch1))641    prev_frames = replay_context_from_generated_frames(frames_ch1, n_ctx)642    prev_pil = [_frame_to_pil(f, w, h) for f in prev_frames]643    identity_rt = [0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0]644    pipe.load_models_to_device(["vae"])645    with torch.no_grad():646        context_latents = encode_context_frames_per_frame(pipe, prev_pil, pipe.device)647    num_ctx_tokens = context_latents.shape[2]648    context_actions_t = torch.tensor([identity_rt] * num_ctx_tokens, dtype=torch.float32)649    print(f"[Loop] continuation chunk: len(prev_pil)={len(prev_pil)} num_ctx_tokens={num_ctx_tokens} (应等于 ctx)")650    frames_ch2 = run_one_chunk(651        pipe, prompt, use_negative_prompt, path_ch2,652        context_latents=context_latents,653        num_context_frames=num_ctx_tokens,654        context_actions_t=context_actions_t,655        chunk_frames=chunk_frames, h=h, w=w, seed=seed + 1,656        sigma_shift=sigma_shift, num_inference_steps=num_inference_steps,657        cfg_scale=cfg_scale, inference_noise_level=inference_noise_level,658        omit_context_actions=omit_context_actions,659    )660    chunk_frames_list.append(frames_ch2)661    context_frames_per_chunk.append(list(prev_pil))662 663    if not keep_action_jsons:664        for p in [path_ch1, path_ch2]:665            if os.path.exists(p):666                try:667                    os.remove(p)668                except Exception:669                    pass670    return chunk_frames_list, yaw_history, context_frames_per_chunk671 672 673def run_single_chunk_rotation(674    pipe,675    dataset_base,676    output_dir,677    video_name,678    start_frame,679    deg=45.0,680    clockwise=False,681    chunk_frames=81,682    h=352,683    w=640,684    sigma_shift=15.0,685    num_inference_steps=50,686    cfg_scale=5.0,687    seed=42,688    inference_noise_level=0.0,689    metadata_path=None,690    keep_action_jsons=False,691    sampling_action_dir=None,692    omit_context_actions=True,693    context_image_path=None,694):695    """旋转原子操作:单 chunk 左转(CCW)或右转(CW) deg°。优先使用与训练采样相同的 action JSON。omit_context_actions 与训练 ctx 设置对齐。context_image_path:泛化检查时用指定图片作为首帧 context。"""696    label = "右转(CW)" if clockwise else "左转(CCW)"697    print(f"[SingleChunk] 1 chunk {label} {deg}° sigma_shift={sigma_shift}")698    prompt = irc.load_prompt_for_video(dataset_base, video_name) or "A scene."699    use_negative_prompt = getattr(irc, "DEFAULT_NEGATIVE_PROMPT", "oversaturated colors, overexposed, static, blurry details")700 701    # 与采样对齐:45°、81 帧时优先用采样同款 JSON(generate_rotation_actions.py 生成),注入方式一致702    path_ch = None703    if deg == 45.0 and chunk_frames == 81:704        action_dir = sampling_action_dir if sampling_action_dir else _script_dir705        if clockwise:706            candidate = os.path.join(action_dir, "action_rotation_right_45.json")707        else:708            candidate = os.path.join(action_dir, "action_rotation_left_45.json")709        if os.path.isfile(candidate):710            path_ch = candidate711            print(f"[SingleChunk] 使用采样同款 action: {os.path.basename(path_ch)} (与训练采样注入一致)", flush=True)712 713    if path_ch is None:714        actions, yaw_chunk = build_action_chunk(deg, clockwise, chunk_frames)715        yaw0 = _yaw_deg_from_rt(actions["0"])716        yaw_last = _yaw_deg_from_rt(actions[str(chunk_frames - 1)])717        print(f"[SingleChunk] action 首帧 yaw={yaw0:.1f}° 末帧 yaw={yaw_last:.1f}° (预期: 左转 0→+{deg}° 右转 0→-{deg}°)", flush=True)718        subdir = os.path.join(output_dir, f"loop_{video_name}_start{start_frame}")719        os.makedirs(subdir, exist_ok=True)720        suffix = "cw_right" if clockwise else "ccw_left"721        path_ch = os.path.join(subdir, f"_single_{suffix}.json")722        with open(path_ch, "w") as f:723            json.dump(actions, f, indent=2)724        yaw_history = list(yaw_chunk)725        delete_path_after = not keep_action_jsons726    else:727        with open(path_ch, "r") as f:728            actions = json.load(f)729        yaw_history = [_yaw_deg_from_rt(actions.get(str(i), [0] * 12)) for i in range(min(chunk_frames, len(actions)))]730        if len(yaw_history) < chunk_frames:731            yaw_history += [yaw_history[-1]] * (chunk_frames - len(yaw_history))732        delete_path_after = False733 734    subdir = os.path.join(output_dir, f"loop_{video_name}_start{start_frame}")735    os.makedirs(subdir, exist_ok=True)736    yaw_history = yaw_history[:chunk_frames] if len(yaw_history) > chunk_frames else yaw_history737    chunk_frames_list = []738    context_frames_per_chunk = []739 740    if context_image_path and os.path.isfile(context_image_path):741        ctx_pil_0 = Image.open(context_image_path).convert("RGB").resize((w, h))742        print(f"[SingleChunk] 使用指定 context 图: {context_image_path}", flush=True)743    else:744        ctx_pil_0 = load_sample_first_frame(dataset_base, video_name, start_frame, w, h)745    if ctx_pil_0 is None:746        ctx_pil_0 = sample_random_frame_from_dataset(dataset_base, w, h, seed, metadata_path)747    if ctx_pil_0 is not None:748        ctx_pil_0 = [ctx_pil_0]749        pipe.load_models_to_device(["vae"])750        with torch.no_grad():751            ctx_latents_0 = encode_context_frames_per_frame(pipe, ctx_pil_0, pipe.device)752        identity_rt = [0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0]753        ctx_actions_t_0 = torch.tensor([identity_rt], dtype=torch.float32)754        frames_ch = run_one_chunk(755            pipe, prompt, use_negative_prompt, path_ch,756            context_latents=ctx_latents_0,757            num_context_frames=ctx_latents_0.shape[2],758            context_actions_t=ctx_actions_t_0,759            chunk_frames=chunk_frames, h=h, w=w, seed=seed,760            sigma_shift=sigma_shift, num_inference_steps=num_inference_steps,761            cfg_scale=cfg_scale, inference_noise_level=inference_noise_level,762            omit_context_actions=omit_context_actions,763        )764        context_frames_per_chunk.append([ctx_pil_0[0]])765    else:766        frames_ch = run_one_chunk(767            pipe, prompt, use_negative_prompt, path_ch,768            chunk_frames=chunk_frames, h=h, w=w, seed=seed,769            sigma_shift=sigma_shift, num_inference_steps=num_inference_steps, cfg_scale=cfg_scale,770            omit_context_actions=omit_context_actions,771        )772        context_frames_per_chunk.append([_frame_to_pil(frames_ch[-1], w, h)])773    chunk_frames_list.append(frames_ch)774 775    if delete_path_after and os.path.exists(path_ch):776        try:777            os.remove(path_ch)778        except Exception:779            pass780    return chunk_frames_list, yaw_history, context_frames_per_chunk781 782 783def run_left_right_2chunk(784    pipe,785    dataset_base,786    output_dir,787    video_name,788    start_frame,789    deg=45.0,790    chunk_frames=81,791    context_frames=1,792    h=352,793    w=640,794    sigma_shift=15.0,795    num_inference_steps=50,796    cfg_scale=5.0,797    seed=42,798    inference_noise_level=0.0,799    metadata_path=None,800    keep_action_jsons=False,801    sampling_action_dir=None,802    omit_context_actions=True,803):804    """回环:1 chunk 左转(CCW) deg°,1 chunk 右转(CW) deg°。45°×81 时优先用采样同款 left_45/right_45.json。"""805    print(f"[LeftRight 2chunk] 左转(CCW){deg}° then 右转(CW){deg}° sigma_shift={sigma_shift} omit_ctx_act={omit_context_actions} context_frames={context_frames} (续段 ctx 数)")806    prompt = irc.load_prompt_for_video(dataset_base, video_name) or "A scene."807    use_negative_prompt = getattr(irc, "DEFAULT_NEGATIVE_PROMPT", "oversaturated colors, overexposed, static, blurry details")808 809    subdir = os.path.join(output_dir, f"loop_{video_name}_start{start_frame}")810    os.makedirs(subdir, exist_ok=True)811    action_dir = sampling_action_dir if sampling_action_dir else _script_dir812    path_ch1 = path_ch2 = None813    delete_ch1 = delete_ch2 = True814    if deg == 45.0 and chunk_frames == 81:815        p_left = os.path.join(action_dir, "action_rotation_left_45.json")816        p_right = os.path.join(action_dir, "action_rotation_right_45.json")817        if os.path.isfile(p_left) and os.path.isfile(p_right):818            path_ch1, path_ch2 = p_left, p_right819            delete_ch1 = delete_ch2 = False820            print(f"[LeftRight 2chunk] 使用采样同款 action: left_45 + right_45 (与训练采样注入一致)", flush=True)821    if path_ch1 is None:822        actions_ccw, actions_cw = build_action_ccw_cw(deg, chunk_frames)823        path_ch1 = os.path.join(subdir, "_ch1_ccw_left.json")824        path_ch2 = os.path.join(subdir, "_ch2_cw_right.json")825        with open(path_ch1, "w") as f:826            json.dump(actions_ccw, f, indent=2)827        with open(path_ch2, "w") as f:828            json.dump(actions_cw, f, indent=2)829        delete_ch1 = delete_ch2 = not keep_action_jsons830 831    if not delete_ch1 and path_ch1 and path_ch2:832        with open(path_ch1, "r") as f:833            a1 = json.load(f)834        with open(path_ch2, "r") as f:835            a2 = json.load(f)836        yaw_left = [_yaw_deg_from_rt(a1.get(str(i), [0] * 12)) for i in range(min(chunk_frames, len(a1)))]837        yaw_right = [_yaw_deg_from_rt(a2.get(str(i), [0] * 12)) for i in range(min(chunk_frames, len(a2)))]838        if len(yaw_left) < chunk_frames:839            yaw_left += [yaw_left[-1]] * (chunk_frames - len(yaw_left))840        if len(yaw_right) < chunk_frames:841            yaw_right += [yaw_right[-1]] * (chunk_frames - len(yaw_right))842        base = yaw_left[-1] if yaw_left else 45.0843        yaw_history = yaw_left[:chunk_frames] + [base + yaw_right[i] for i in range(chunk_frames)]844    else:845        denom = max(1, chunk_frames - 1)846        yaw_history = [(i / denom) * deg for i in range(chunk_frames)]847        yaw_history += [deg - (i / denom) * deg for i in range(chunk_frames)]848 849    chunk_frames_list = []850    context_frames_per_chunk = []851 852    ctx_pil_0 = load_sample_first_frame(dataset_base, video_name, start_frame, w, h)853    used_sample_first = ctx_pil_0 is not None854    if ctx_pil_0 is None:855        ctx_pil_0 = sample_random_frame_from_dataset(dataset_base, w, h, seed, metadata_path)856    if ctx_pil_0 is not None:857        ctx_pil_0 = [ctx_pil_0]858        pipe.load_models_to_device(["vae"])859        with torch.no_grad():860            ctx_latents_0 = encode_context_frames_per_frame(pipe, ctx_pil_0, pipe.device)861        identity_rt = [0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0]862        ctx_actions_t_0 = torch.tensor([identity_rt], dtype=torch.float32)863        frames_ch1 = run_one_chunk(864            pipe, prompt, use_negative_prompt, path_ch1,865            context_latents=ctx_latents_0,866            num_context_frames=ctx_latents_0.shape[2],867            context_actions_t=ctx_actions_t_0,868            chunk_frames=chunk_frames, h=h, w=w, seed=seed,869            sigma_shift=sigma_shift, num_inference_steps=num_inference_steps,870            cfg_scale=cfg_scale, inference_noise_level=inference_noise_level,871            omit_context_actions=omit_context_actions,872        )873        context_frames_per_chunk.append([ctx_pil_0[0]])874    else:875        frames_ch1 = run_one_chunk(876            pipe, prompt, use_negative_prompt, path_ch1,877            chunk_frames=chunk_frames, h=h, w=w, seed=seed,878            sigma_shift=sigma_shift, num_inference_steps=num_inference_steps, cfg_scale=cfg_scale,879            omit_context_actions=omit_context_actions,880        )881        context_frames_per_chunk.append([_frame_to_pil(frames_ch1[-1], w, h)])882    chunk_frames_list.append(frames_ch1)883 884    n_ctx = min(context_frames, len(frames_ch1))885    prev_frames = replay_context_from_generated_frames(frames_ch1, n_ctx)886    prev_pil = [_frame_to_pil(f, w, h) for f in prev_frames]887    identity_rt = [0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0]888    pipe.load_models_to_device(["vae"])889    with torch.no_grad():890        context_latents = encode_context_frames_per_frame(pipe, prev_pil, pipe.device)891    num_ctx_tokens = context_latents.shape[2]892    context_actions_t = torch.tensor([identity_rt] * num_ctx_tokens, dtype=torch.float32)893    print(f"[Loop] continuation chunk: len(prev_pil)={len(prev_pil)} num_ctx_tokens={num_ctx_tokens} (应等于 ctx)")894    frames_ch2 = run_one_chunk(895        pipe, prompt, use_negative_prompt, path_ch2,896        context_latents=context_latents,897        num_context_frames=num_ctx_tokens,898        context_actions_t=context_actions_t,899        chunk_frames=chunk_frames, h=h, w=w, seed=seed + 1,900        sigma_shift=sigma_shift, num_inference_steps=num_inference_steps,901        cfg_scale=cfg_scale, inference_noise_level=inference_noise_level,902        omit_context_actions=omit_context_actions,903    )904    chunk_frames_list.append(frames_ch2)905    context_frames_per_chunk.append(list(prev_pil))906 907    if delete_ch1 and path_ch1 and os.path.exists(path_ch1):908        try:909            os.remove(path_ch1)910        except Exception:911            pass912    if delete_ch2 and path_ch2 and os.path.exists(path_ch2):913        try:914            os.remove(path_ch2)915        except Exception:916            pass917    return chunk_frames_list, yaw_history, context_frames_per_chunk918 919 920def run_translation_4chunk(921    pipe,922    dataset_base,923    output_dir,924    video_name,925    start_frame,926    direction="forward",927    translation_delta=0.1,928    chunk_frames=81,929    context_frames=1,930    h=352,931    w=640,932    sigma_shift=15.0,933    num_inference_steps=50,934    cfg_scale=5.0,935    seed=42,936    inference_noise_level=0.0,937    metadata_path=None,938    omit_context_actions=True,939):940    """仅纯平移 4chunk:同一方向(forward/backward/left/right)连续 4 段,R=I,无旋转。仅用于 --run_atomic_translation_4chunk。"""941    _prefix = "[atomic_translation_4chunk]"942    print(f"{_prefix} 纯平移 direction={direction} delta={translation_delta} sigma_shift={sigma_shift} omit_ctx_act={omit_context_actions} context_frames={context_frames} (续段 ctx 数)")943    prompt = irc.load_prompt_for_video(dataset_base, video_name) or "A scene."944    use_negative_prompt = getattr(irc, "DEFAULT_NEGATIVE_PROMPT", "oversaturated colors, overexposed, static, blurry details")945 946    yaw_history = []947    chunk_frames_list = []948    context_frames_per_chunk = []949    subdir = os.path.join(output_dir, f"loop_{video_name}_start{start_frame}")950    os.makedirs(subdir, exist_ok=True)951    identity_rt = [0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0]952    context_latents = None953    context_actions_t = None954 955    for ch in range(4):956        actions, yaw_chunk = build_action_translation_only(direction, translation_delta, chunk_frames)957        path_ch = os.path.join(subdir, f"_ch{ch}_trans_{direction}.json")958        with open(path_ch, "w") as f:959            json.dump(actions, f, indent=2)960        try:961            if ch == 0:962                ctx_pil_0 = load_sample_first_frame(dataset_base, video_name, start_frame, w, h)963                if ctx_pil_0 is None:964                    ctx_pil_0 = sample_random_frame_from_dataset(dataset_base, w, h, seed, metadata_path)965                if ctx_pil_0 is not None:966                    ctx_pil_0 = [ctx_pil_0]967                    pipe.load_models_to_device(["vae"])968                    with torch.no_grad():969                        context_latents = encode_context_frames_per_frame(pipe, ctx_pil_0, pipe.device)970                    context_actions_t = torch.tensor([identity_rt], dtype=torch.float32)971                    context_frames_per_chunk.append([ctx_pil_0[0]])972                else:973                    context_frames_per_chunk.append([])974 975            if context_latents is not None and context_actions_t is not None:976                frames_ch = run_one_chunk(977                    pipe, prompt, use_negative_prompt, path_ch,978                    context_latents=context_latents,979                    num_context_frames=context_latents.shape[2],980                    context_actions_t=context_actions_t,981                    chunk_frames=chunk_frames, h=h, w=w, seed=seed + ch,982                    sigma_shift=sigma_shift, num_inference_steps=num_inference_steps,983                    cfg_scale=cfg_scale, inference_noise_level=inference_noise_level,984                    omit_context_actions=omit_context_actions,985                    log_prefix=_prefix,986                )987            else:988                frames_ch = run_one_chunk(989                    pipe, prompt, use_negative_prompt, path_ch,990                    chunk_frames=chunk_frames, h=h, w=w, seed=seed + ch,991                    sigma_shift=sigma_shift, num_inference_steps=num_inference_steps, cfg_scale=cfg_scale,992                    omit_context_actions=omit_context_actions,993                    log_prefix=_prefix,994                )995                if ch == 0 and not context_frames_per_chunk[0]:996                    context_frames_per_chunk[0] = [_frame_to_pil(frames_ch[-1], w, h)]997 998            for y in yaw_chunk:999                yaw_history.append(y)1000            chunk_frames_list.append(frames_ch)1001 1002            if ch < 3:1003                n_ctx = min(context_frames, len(frames_ch))1004                prev_frames = context_frames_for_next_chunk(frames_ch, n_ctx) if n_ctx else [frames_ch[-1]]1005                prev_pil = [_frame_to_pil(f, w, h) for f in prev_frames]1006                context_frames_per_chunk.append(list(prev_pil))1007                pipe.load_models_to_device(["vae"])1008                with torch.no_grad():1009                    context_latents = encode_context_frames_per_frame(pipe, prev_pil, pipe.device)1010                num_ctx_tokens = context_latents.shape[2]1011                context_actions_t = torch.tensor([identity_rt] * num_ctx_tokens, dtype=torch.float32)1012                print(f"{_prefix} continuation chunk ch={ch}: len(prev_pil)={len(prev_pil)} num_ctx_tokens={num_ctx_tokens} (应等于 ctx)")1013        finally:1014            if os.path.exists(path_ch):1015                try:1016                    os.remove(path_ch)1017                except Exception:1018                    pass1019 1020    return chunk_frames_list, yaw_history, context_frames_per_chunk1021 1022 1023def run_left2_right2_4chunk(1024    pipe,1025    dataset_base,1026    output_dir,1027    video_name,1028    start_frame,1029    deg_per_chunk=45.0,1030    chunk_frames=81,1031    context_frames=1,1032    h=352,1033    w=640,1034    sigma_shift=15.0,1035    num_inference_steps=50,1036    cfg_scale=5.0,1037    seed=42,1038    inference_noise_level=0.0,1039    metadata_path=None,1040    sampling_action_dir=None,1041    omit_context_actions=True,1042    multi_ctx_all_history=False,1043    fov_history_context=False,1044    fov_last_target=False,1045    fov_context_rt=False,1046):1047    """回环:先左转 2 chunk(各 deg° CCW),再右转 2 chunk(各 deg° CW)。1048 1049    默认行为:与 2chunk 回环一致,续段 context 仅来自上一 chunk 的末尾若干帧。1050    当 multi_ctx_all_history=True 时:续段 context 第一帧必须为上一 chunk 的最后一帧;1051    其余 (ctx-1) 帧从「前面所有已生成帧」均匀采样(保持时间序),用于记忆机制对比。1052    当 fov_history_context=True 时:续段 context 第一帧仍为上一 chunk 的最后一帧;1053    其余 (ctx-1) 帧按下一 chunk 中点 world-yaw,从所有已生成历史帧中检索最接近的帧。1054    当 fov_last_target=True 时:仅第 4 个 chunk 的检索 target 改用该 chunk 末尾 world-yaw。1055    当 fov_context_rt=True 时:FOV/yaw 检索出的 context actions 转成相对下一 chunk 起始 world-yaw 的 RT。1056    """1057    print(f"[Left2Right2 4chunk] 左转2chunk then 右转2chunk sigma_shift={sigma_shift} omit_ctx_act={omit_context_actions} context_frames={context_frames} (续段 ctx 数)")1058    prompt = irc.load_prompt_for_video(dataset_base, video_name) or "A scene."1059    use_negative_prompt = getattr(irc, "DEFAULT_NEGATIVE_PROMPT", "oversaturated colors, overexposed, static, blurry details")1060 1061    yaw_history = []1062    chunk_frames_list = []1063    context_frames_per_chunk = []1064    all_prev_frames = []1065    subdir = os.path.join(output_dir, f"loop_{video_name}_start{start_frame}")1066    os.makedirs(subdir, exist_ok=True)1067    identity_rt = [0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0]1068    cumulative_yaw = 0.01069    context_latents = None1070    context_actions_t = None1071 1072    # 45°×81 时直接使用采样同款 JSON,多 chunk 多次导入同一文件即可1073    action_dir = sampling_action_dir if sampling_action_dir else _script_dir1074    path_left = os.path.join(action_dir, "action_rotation_left_45.json")1075    path_right = os.path.join(action_dir, "action_rotation_right_45.json")1076    use_sampling_files = (deg_per_chunk == 45.0 and chunk_frames == 81 and1077                          os.path.isfile(path_left) and os.path.isfile(path_right))1078    if use_sampling_files:1079        print(f"[Left2Right2 4chunk] 使用采样同款 action: left_45 x2 + right_45 x2 (与训练采样注入一致)", flush=True)1080        with open(path_left, "r") as f:1081            _actions_left = json.load(f)1082        with open(path_right, "r") as f:1083            _actions_right = json.load(f)1084        _yaw_left = [_yaw_deg_from_rt(_actions_left.get(str(i), [0] * 12)) for i in range(min(chunk_frames, len(_actions_left)))]1085        _yaw_right = [_yaw_deg_from_rt(_actions_right.get(str(i), [0] * 12)) for i in range(min(chunk_frames, len(_actions_right)))]1086        if len(_yaw_left) < chunk_frames:1087            _yaw_left += [_yaw_left[-1]] * (chunk_frames - len(_yaw_left))1088        if len(_yaw_right) < chunk_frames:1089            _yaw_right += [_yaw_right[-1]] * (chunk_frames - len(_yaw_right))1090 1091    for ch in range(4):1092        clockwise = ch >= 21093        if use_sampling_files:1094            path_ch = path_right if clockwise else path_left1095        else:1096            actions, yaw_chunk = build_action_chunk(deg_per_chunk, clockwise, chunk_frames)1097            path_ch = os.path.join(subdir, f"_ch{ch}_left2right2.json")1098            with open(path_ch, "w") as f:1099                json.dump(actions, f, indent=2)1100 1101        # 与 2chunk 一致:ch=0 用首帧 context;ch>=1 用上一 chunk 输出做 context(上一轮末尾已写好 context_latents/context_actions_t)1102        if ch == 0:1103            ctx_pil_0 = load_sample_first_frame(dataset_base, video_name, start_frame, w, h)1104            if ctx_pil_0 is None:1105                ctx_pil_0 = sample_random_frame_from_dataset(dataset_base, w, h, seed, metadata_path)1106            if ctx_pil_0 is not None:1107                ctx_pil_0 = [ctx_pil_0]1108                pipe.load_models_to_device(["vae"])1109                with torch.no_grad():1110                    context_latents = encode_context_frames_per_frame(pipe, ctx_pil_0, pipe.device)1111                context_actions_t = torch.tensor([identity_rt], dtype=torch.float32)1112                context_frames_per_chunk.append([ctx_pil_0[0]])1113            else:1114                context_frames_per_chunk.append([])1115 1116        if context_latents is not None and context_actions_t is not None:1117            frames_ch = run_one_chunk(1118                pipe, prompt, use_negative_prompt, path_ch,1119                context_latents=context_latents,1120                num_context_frames=context_latents.shape[2],1121                context_actions_t=context_actions_t,1122                chunk_frames=chunk_frames, h=h, w=w, seed=seed + ch,1123                sigma_shift=sigma_shift, num_inference_steps=num_inference_steps,1124                cfg_scale=cfg_scale, inference_noise_level=inference_noise_level,1125                omit_context_actions=omit_context_actions,1126            )1127        else:1128            frames_ch = run_one_chunk(1129                pipe, prompt, use_negative_prompt, path_ch,1130                chunk_frames=chunk_frames, h=h, w=w, seed=seed + ch,1131                sigma_shift=sigma_shift, num_inference_steps=num_inference_steps, cfg_scale=cfg_scale,1132                omit_context_actions=omit_context_actions,1133            )1134            if ch == 0 and not context_frames_per_chunk[0]:1135                context_frames_per_chunk[0] = [_frame_to_pil(frames_ch[-1], w, h)]1136 1137        if use_sampling_files:1138            yaw_chunk = _yaw_right if clockwise else _yaw_left1139        else:1140            yaw_chunk = list(yaw_chunk)1141        for y in yaw_chunk:1142            yaw_history.append(cumulative_yaw + y)1143        cumulative_yaw = yaw_history[-1]1144        chunk_frames_list.append(frames_ch)1145        all_prev_frames.extend(frames_ch)1146 1147        # 续段 context:1148        # - 默认:与 2chunk 一致,仅用上一 chunk 的 last n 帧1149        # - fov_history_context=True:第一帧必须为上一 chunk 最后一帧;其余 (ctx-1) 帧按下一 chunk 中点 world-yaw 做 FOV/yaw proxy 检索。1150        # - multi_ctx_all_history=True:第一帧必须为上一 chunk 最后一帧;其余 (ctx-1) 帧从「前面所有帧」均匀采样,保持时间序。1151        if ch < 3:1152            next_context_actions = None1153            if context_frames <= 0:1154                prev_frames = [frames_ch[-1]]1155            elif fov_history_context and len(all_prev_frames) > 0:1156                next_ch = ch + 11157                next_clockwise = next_ch >= 21158                if use_sampling_files:1159                    next_yaw_chunk = _yaw_right if next_clockwise else _yaw_left1160                else:1161                    _unused_actions, next_yaw_chunk = build_action_chunk(deg_per_chunk, next_clockwise, chunk_frames)1162                target_idx = min(max(0, chunk_frames // 2), len(next_yaw_chunk) - 1)1163                if fov_last_target and next_ch == 3:1164                    target_idx = len(next_yaw_chunk) - 11165                target_mid_yaw = cumulative_yaw + float(next_yaw_chunk[target_idx])1166                prev_frames, picked_indices = fov_history_context_from_generated_frames(1167                    all_prev_frames,1168                    yaw_history,1169                    context_frames,1170                    target_mid_yaw,1171                )1172                picked_yaws = [float(yaw_history[i]) for i in picked_indices if i < len(yaw_history)]1173                if fov_context_rt and picked_yaws:1174                    next_context_actions = context_actions_from_world_yaws(picked_yaws, cumulative_yaw)1175                print(1176                    f"[Loop] fov_history_context ch={ch}->next={next_ch}: "1177                    f"target_idx={target_idx} target_yaw={target_mid_yaw:.2f} picked_indices={picked_indices} "1178                    f"picked_yaws={[round(y, 2) for y in picked_yaws]} "1179                    f"context_rt={bool(next_context_actions)} ref_yaw={cumulative_yaw:.2f}",1180                    flush=True,1181                )1182            elif multi_ctx_all_history and len(all_prev_frames) > 0:1183                total = len(all_prev_frames)1184                if total == 1 or context_frames == 1:1185                    prev_frames = [all_prev_frames[-1]]1186                else:1187                    # 第一帧 = 上一 chunk 最后一帧 (all_prev_frames[-1]);其余 (ctx-1) 帧从 all_prev_frames[0..total-2] 均匀采样1188                    n_pick = min(context_frames - 1, total - 1)1189                    if n_pick <= 0:1190                        prev_frames = [all_prev_frames[-1]]1191                    else:1192                        pool_size = total - 1  # 可选下标 0..total-2(前面所有帧,不含已单独占位的最后一帧)1193                        indices = []1194                        for i in range(n_pick):1195                            pos = int(round(i * (pool_size - 1) / max(n_pick - 1, 1)))1196                            indices.append(max(0, min(pool_size - 1, pos)))1197                        uniq_indices = sorted(set(indices))1198                        while len(uniq_indices) < n_pick:1199                            uniq_indices.append(pool_size - 1 if pool_size > 0 else 0)1200                        uniq_indices = sorted(uniq_indices)[:n_pick]

Showing the first 1,200 of 1946 lines. Download the file for the rest.