hugging-apps/echo-memory
0
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]