hugging-apps/echo-memory
0
1import imageio, os, torch, warnings, torchvision, argparse, json, random2from peft import LoraConfig, inject_adapter_in_model3from PIL import Image4import pandas as pd5from tqdm import tqdm6from accelerate import Accelerator7 8 9 10class ImageDataset(torch.utils.data.Dataset):11 def __init__(12 self,13 base_path=None, metadata_path=None,14 max_pixels=1920*1080, height=None, width=None,15 height_division_factor=16, width_division_factor=16,16 data_file_keys=("image",),17 image_file_extension=("jpg", "jpeg", "png", "webp"),18 repeat=1,19 args=None,20 ):21 if args is not None:22 base_path = args.dataset_base_path23 metadata_path = args.dataset_metadata_path24 height = args.height25 width = args.width26 max_pixels = args.max_pixels27 data_file_keys = args.data_file_keys.split(",")28 repeat = args.dataset_repeat29 30 self.base_path = base_path31 self.max_pixels = max_pixels32 self.height = height33 self.width = width34 self.height_division_factor = height_division_factor35 self.width_division_factor = width_division_factor36 self.data_file_keys = data_file_keys37 self.image_file_extension = image_file_extension38 self.repeat = repeat39 40 if height is not None and width is not None:41 print("Height and width are fixed. Setting `dynamic_resolution` to False.")42 self.dynamic_resolution = False43 elif height is None and width is None:44 print("Height and width are none. Setting `dynamic_resolution` to True.")45 self.dynamic_resolution = True46 47 if metadata_path is None:48 print("No metadata. Trying to generate it.")49 metadata = self.generate_metadata(base_path)50 print(f"{len(metadata)} lines in metadata.")51 self.data = [metadata.iloc[i].to_dict() for i in range(len(metadata))]52 elif metadata_path.endswith(".json"):53 with open(metadata_path, "r") as f:54 metadata = json.load(f)55 self.data = metadata56 else:57 metadata = pd.read_csv(metadata_path)58 # Ensure prompt column is string type to avoid float conversion for NaN values59 if 'prompt' in metadata.columns:60 metadata['prompt'] = metadata['prompt'].astype(str)61 # Replace 'nan' string (from NaN) with empty string62 metadata['prompt'] = metadata['prompt'].replace('nan', '')63 self.data = [metadata.iloc[i].to_dict() for i in range(len(metadata))]64 65 66 def generate_metadata(self, folder):67 image_list, prompt_list = [], []68 file_set = set(os.listdir(folder))69 for file_name in file_set:70 if "." not in file_name:71 continue72 file_ext_name = file_name.split(".")[-1].lower()73 file_base_name = file_name[:-len(file_ext_name)-1]74 if file_ext_name not in self.image_file_extension:75 continue76 prompt_file_name = file_base_name + ".txt"77 if prompt_file_name not in file_set:78 continue79 with open(os.path.join(folder, prompt_file_name), "r", encoding="utf-8") as f:80 prompt = f.read().strip()81 image_list.append(file_name)82 prompt_list.append(prompt)83 metadata = pd.DataFrame()84 metadata["image"] = image_list85 metadata["prompt"] = prompt_list86 return metadata87 88 89 def crop_and_resize(self, image, target_height, target_width):90 width, height = image.size91 scale = max(target_width / width, target_height / height)92 image = torchvision.transforms.functional.resize(93 image,94 (round(height*scale), round(width*scale)),95 interpolation=torchvision.transforms.InterpolationMode.BILINEAR96 )97 image = torchvision.transforms.functional.center_crop(image, (target_height, target_width))98 return image99 100 101 def get_height_width(self, image):102 if self.dynamic_resolution:103 width, height = image.size104 if width * height > self.max_pixels:105 scale = (width * height / self.max_pixels) ** 0.5106 height, width = int(height / scale), int(width / scale)107 height = height // self.height_division_factor * self.height_division_factor108 width = width // self.width_division_factor * self.width_division_factor109 else:110 height, width = self.height, self.width111 return height, width112 113 114 def load_image(self, file_path):115 image = Image.open(file_path).convert("RGB")116 image = self.crop_and_resize(image, *self.get_height_width(image))117 return image118 119 120 def load_data(self, file_path):121 return self.load_image(file_path)122 123 124 def __getitem__(self, data_id):125 data = self.data[data_id % len(self.data)].copy()126 for key in self.data_file_keys:127 if key in data:128 path = os.path.join(self.base_path, data[key])129 data[key] = self.load_data(path)130 if data[key] is None:131 warnings.warn(f"cannot load file {data[key]}.")132 return None133 return data134 135 136 def __len__(self):137 return len(self.data) * self.repeat138 139 140 141class VideoDataset(torch.utils.data.Dataset):142 def __init__(143 self,144 base_path=None, metadata_path=None,145 num_frames=81,146 time_division_factor=4, time_division_remainder=1,147 max_pixels=1920*1080, height=None, width=None,148 height_division_factor=16, width_division_factor=16,149 data_file_keys=("video",),150 image_file_extension=("jpg", "jpeg", "png", "webp"),151 video_file_extension=("mp4", "avi", "mov", "wmv", "mkv", "flv", "webm"),152 repeat=1,153 args=None,154 action_base_path=None,155 enable_icl=False,156 icl_num_examples=2,157 icl_context_frames=8,158 ):159 if args is not None:160 base_path = args.dataset_base_path161 metadata_path = args.dataset_metadata_path162 height = args.height163 width = args.width164 max_pixels = args.max_pixels165 num_frames = args.num_frames166 data_file_keys = args.data_file_keys.split(",")167 repeat = args.dataset_repeat168 # In-context learning parameters169 if hasattr(args, 'enable_icl'):170 enable_icl = args.enable_icl171 if hasattr(args, 'icl_num_examples'):172 icl_num_examples = args.icl_num_examples173 if hasattr(args, 'icl_context_frames'):174 icl_context_frames = args.icl_context_frames175 176 self.base_path = base_path177 self.num_frames = num_frames178 self.time_division_factor = time_division_factor179 self.time_division_remainder = time_division_remainder180 self.max_pixels = max_pixels181 self.height = height182 self.width = width183 self.height_division_factor = height_division_factor184 self.width_division_factor = width_division_factor185 self.data_file_keys = data_file_keys186 self.image_file_extension = image_file_extension187 self.video_file_extension = video_file_extension188 self.repeat = repeat189 190 # In-context learning parameters191 self.enable_icl = enable_icl192 self.icl_num_examples = icl_num_examples193 self.icl_context_frames = icl_context_frames194 195 if height is not None and width is not None:196 print("Height and width are fixed. Setting `dynamic_resolution` to False.")197 self.dynamic_resolution = False198 elif height is None and width is None:199 print("Height and width are none. Setting `dynamic_resolution` to True.")200 self.dynamic_resolution = True201 202 if metadata_path is None:203 print("No metadata. Trying to generate it.")204 metadata = self.generate_metadata(base_path)205 print(f"{len(metadata)} lines in metadata.")206 self.data = [metadata.iloc[i].to_dict() for i in range(len(metadata))]207 elif metadata_path.endswith(".json"):208 with open(metadata_path, "r") as f:209 metadata = json.load(f)210 self.data = metadata211 else:212 metadata = pd.read_csv(metadata_path)213 # Ensure prompt column is string type to avoid float conversion for NaN values214 if 'prompt' in metadata.columns:215 metadata['prompt'] = metadata['prompt'].astype(str)216 # Replace 'nan' string (from NaN) with empty string217 metadata['prompt'] = metadata['prompt'].replace('nan', '')218 219 # CRITICAL FIX: Clean prompt - remove video path prefix if present220 # Some CSV prompts start with "video_name.mp4 " prefix, which should be removed221 def clean_prompt(prompt_str):222 if not isinstance(prompt_str, str) or not prompt_str:223 return prompt_str224 # Check if prompt starts with a video path (contains .mp4 or /)225 # Pattern: "VideoName/1234_5678.mp4 " or "VideoName.mp4 "226 import re227 # Match pattern: word/word.mp4 or word.mp4 at the start, followed by space228 pattern = r'^[A-Za-z0-9_]+(/[A-Za-z0-9_]+)?\.mp4\s+'229 cleaned = re.sub(pattern, '', prompt_str)230 # Also handle truncated prompts ending with "..."231 if cleaned.endswith('...'):232 cleaned = cleaned[:-3].rstrip()233 return cleaned.strip()234 235 metadata['prompt'] = metadata['prompt'].apply(clean_prompt)236 self.data = [metadata.iloc[i].to_dict() for i in range(len(metadata))]237 238 self.action_base_path = action_base_path239 240 if self.enable_icl:241 print(f"In-context learning enabled: {icl_num_examples} examples, {icl_context_frames} context frames each")242 243 244 def generate_metadata(self, folder):245 video_list, prompt_list = [], []246 file_set = set(os.listdir(folder))247 for file_name in file_set:248 if "." not in file_name:249 continue250 file_ext_name = file_name.split(".")[-1].lower()251 file_base_name = file_name[:-len(file_ext_name)-1]252 if file_ext_name not in self.image_file_extension and file_ext_name not in self.video_file_extension:253 continue254 prompt_file_name = file_base_name + ".txt"255 if prompt_file_name not in file_set:256 continue257 with open(os.path.join(folder, prompt_file_name), "r", encoding="utf-8") as f:258 prompt = f.read().strip()259 video_list.append(file_name)260 prompt_list.append(prompt)261 metadata = pd.DataFrame()262 metadata["video"] = video_list263 metadata["prompt"] = prompt_list264 return metadata265 266 267 def crop_and_resize(self, image, target_height, target_width):268 width, height = image.size269 scale = max(target_width / width, target_height / height)270 image = torchvision.transforms.functional.resize(271 image,272 (round(height*scale), round(width*scale)),273 interpolation=torchvision.transforms.InterpolationMode.BILINEAR274 )275 image = torchvision.transforms.functional.center_crop(image, (target_height, target_width))276 return image277 278 279 def get_height_width(self, image):280 if self.dynamic_resolution:281 width, height = image.size282 if width * height > self.max_pixels:283 scale = (width * height / self.max_pixels) ** 0.5284 height, width = int(height / scale), int(width / scale)285 height = height // self.height_division_factor * self.height_division_factor286 width = width // self.width_division_factor * self.width_division_factor287 else:288 height, width = self.height, self.width289 return height, width290 291 292 def get_num_frames(self, reader):293 num_frames = self.num_frames294 if int(reader.count_frames()) < num_frames:295 num_frames = int(reader.count_frames())296 while num_frames > 1 and num_frames % self.time_division_factor != self.time_division_remainder:297 num_frames -= 1298 return num_frames299 300 301 def load_video(self, file_path):302 reader = imageio.get_reader(file_path)303 num_frames = self.get_num_frames(reader)304 frames = []305 for frame_id in range(num_frames):306 frame = reader.get_data(frame_id)307 frame = Image.fromarray(frame)308 frame = self.crop_and_resize(frame, *self.get_height_width(frame))309 frames.append(frame)310 reader.close()311 return frames312 313 314 def load_image(self, file_path):315 image = Image.open(file_path).convert("RGB")316 image = self.crop_and_resize(image, *self.get_height_width(image))317 frames = [image]318 return frames319 320 321 def is_image(self, file_path):322 file_ext_name = file_path.split(".")[-1]323 return file_ext_name.lower() in self.image_file_extension324 325 326 def is_video(self, file_path):327 file_ext_name = file_path.split(".")[-1]328 return file_ext_name.lower() in self.video_file_extension329 330 331 def load_data(self, file_path):332 # Handle multiple frame paths separated by '|' (for frame sequences)333 if '|' in str(file_path):334 # Split the path by '|' to get individual frame paths335 frame_paths = str(file_path).split('|')336 frames = []337 338 # Get base_path (dataset root)339 if not hasattr(self, 'base_path') or not self.base_path:340 warnings.warn(f"Cannot determine base directory for frame sequence: {file_path}")341 return None342 343 base_dir = self.base_path # This is the dataset root344 345 # Check the first path to determine the format346 first_frame = frame_paths[0].strip() if frame_paths else ""347 348 # If first frame is already an absolute path (from __getitem__ joining),349 # extract the base directory from it350 if os.path.isabs(first_frame):351 # Extract base directory from first frame path352 # First frame format: /path/to/dataset/frames/video_name/frame.png353 # We need to get /path/to/dataset354 parts = first_frame.split(os.sep)355 # Find 'frames' in the path and get everything before it356 if 'frames' in parts:357 frames_idx = parts.index('frames')358 base_dir = os.sep.join(parts[:frames_idx])359 else:360 # Fallback: use self.base_path361 base_dir = self.base_path362 363 for frame_path in frame_paths:364 frame_path = frame_path.strip()365 if not frame_path:366 continue367 368 # Construct full path369 if os.path.isabs(frame_path):370 # Already absolute path (from __getitem__)371 full_frame_path = frame_path372 else:373 # Relative path - need to construct full path374 # Remove 'frames/' prefix if present (we'll add it consistently)375 if frame_path.startswith('frames/'):376 frame_path = frame_path[7:] # Remove 'frames/' prefix377 378 # Always join with base_dir + 'frames/' since base_dir is dataset root379 full_frame_path = os.path.join(base_dir, 'frames', frame_path)380 381 # Load individual frame382 if os.path.exists(full_frame_path):383 if self.is_image(full_frame_path):384 frame_data = self.load_image(full_frame_path)385 if frame_data:386 frames.extend(frame_data)387 else:388 warnings.warn(f"Frame is not an image: {full_frame_path}")389 else:390 warnings.warn(f"Frame not found: {full_frame_path}")391 392 if frames:393 return frames394 else:395 warnings.warn(f"No frames loaded from sequence: {file_path}")396 return None397 398 # Handle single file (image or video)399 if self.is_image(file_path):400 return self.load_image(file_path)401 elif self.is_video(file_path):402 return self.load_video(file_path)403 else:404 return None405 406 407 def __getitem__(self, data_id):408 data = self.data[data_id % len(self.data)].copy()409 for key in self.data_file_keys:410 if key in ["video_name", "start_frame", "end_frame"]:411 if "actions" in data:412 continue413 try:414 video_name = data.get("video_name")415 if video_name is None:416 warnings.warn(f"video_name is missing in metadata for data_id {data_id}. Skipping action loading.")417 continue418 419 if video_name.endswith(".mp4"):420 video_name = ".".join(video_name.split(".")[:-1])421 if "_" in video_name:422 video_name = "_".join(video_name.split("_")[:4])423 424 import json425 json_path = os.path.join(self.action_base_path, video_name + ".json")426 427 # Check if action file exists428 if not os.path.exists(json_path):429 warnings.warn(f"Action file does not exist: {json_path}. Skipping action loading for data_id {data_id}.")430 continue431 432 start_frame = data.get("start_frame")433 end_frame = data.get("end_frame")434 if start_frame is None or end_frame is None:435 warnings.warn(f"start_frame or end_frame is missing in metadata for data_id {data_id}. Skipping action loading.")436 continue437 438 json_data = json.load(open(json_path, "r"))['actions']439 actions = []440 current_yaw = 0.0441 for frame_id in range(start_frame+1, end_frame+1):442 frame_str = str(frame_id)443 if frame_str not in json_data:444 warnings.warn(f"Frame {frame_id} not found in action file {json_path}. Skipping this frame.")445 continue446 447 action = json_data[frame_str]448 new_action = [0.0] * (2 + 2 + 3 + 1 + 2)449 if action['ws'] == 1:450 new_action[0] = 1451 elif action['ws'] == 2:452 new_action[1] = 1453 454 if action['ad'] == 1:455 new_action[2] = 1456 elif action['ad'] == 2:457 new_action[3] = 1458 459 if action['scs'] == 1 and action.get("jump_invalid", 0) == 0:460 new_action[4] = 1461 elif action['scs'] == 2:462 new_action[5] = 1463 elif action['scs'] == 3:464 new_action[6] = 1465 466 if action.get('collision', 0) == 1:467 new_action[7] = 1468 new_action[0] = 0469 new_action[1] = 0470 new_action[2] = 0471 new_action[3] = 0472 473 pre_pitch = action.get('pre_pitch', 0.0)474 current_pitch = pre_pitch + action.get('pitch_delta', 0.0) * 15.0475 current_yaw += action.get('yaw_delta', 0.0) * 15.0476 new_action[8] = current_pitch477 new_action[9] = current_yaw478 479 actions.append(new_action)480 data["actions"] = actions481 except Exception as e:482 warnings.warn(f"Exception while loading actions for data_id {data_id}: {e}. Continuing without actions.")483 # Don't return None, just continue without actions484 continue485 elif key == "video":486 # Check if data[key] exists and is not None487 if key not in data or data[key] is None:488 warnings.warn(f"Video key '{key}' is missing or None in metadata for data_id {data_id}. Skipping this sample.")489 return None490 491 # Handle frame sequences (paths with '|' separator)492 video_path_str = str(data[key])493 if '|' in video_path_str:494 # For frame sequences, pass the full path string to load_data495 # load_data will handle splitting and loading individual frames496 path = os.path.join(self.base_path, video_path_str)497 # Don't check path existence here for frame sequences498 # load_data will handle individual frame loading499 else:500 path = os.path.join(self.base_path, data[key])501 # Check if path exists (only for single files)502 if not os.path.exists(path):503 warnings.warn(f"Video file does not exist: {path}. Skipping this sample.")504 return None505 try:506 data[key] = self.load_data(path)507 if data[key] is None:508 warnings.warn(f"Failed to load video file: {path}. load_data returned None.")509 return None510 except Exception as e:511 warnings.warn(f"Exception while loading video file {path}: {e}. Skipping this sample.")512 return None513 514 # In-context learning: sample context examples from dataset515 if self.enable_icl and len(self.data) > 1:516 context_frames_list = []517 context_actions_list = []518 519 # Sample random examples from dataset (excluding current one)520 current_idx = data_id % len(self.data)521 candidate_indices = [i for i in range(len(self.data)) if i != current_idx]522 if len(candidate_indices) > 0:523 num_samples = min(self.icl_num_examples, len(candidate_indices))524 sampled_indices = random.sample(candidate_indices, num_samples)525 526 for sample_idx in sampled_indices:527 sample_data = self.data[sample_idx].copy()528 # Load video for context529 if "video" in self.data_file_keys and "video" in sample_data:530 video_path = os.path.join(self.base_path, sample_data["video"])531 sample_video = self.load_data(video_path)532 if sample_video is not None and len(sample_video) >= self.icl_context_frames:533 # Sample context_frames from the video534 start_idx = random.randint(0, max(0, len(sample_video) - self.icl_context_frames))535 context_frames = sample_video[start_idx:start_idx + self.icl_context_frames]536 context_frames_list.extend(context_frames)537 538 # Load corresponding actions if available539 if self.action_base_path is not None and "video_name" in sample_data:540 try:541 sample_video_name = sample_data["video_name"]542 if sample_video_name.endswith(".mp4"):543 sample_video_name = ".".join(sample_video_name.split(".")[:-1])544 if "_" in sample_video_name:545 sample_video_name = "_".join(sample_video_name.split("_")[:4])546 sample_json_path = os.path.join(self.action_base_path, sample_video_name + ".json")547 if os.path.exists(sample_json_path):548 sample_json_data = json.load(open(sample_json_path, "r"))['actions']549 sample_start_frame = sample_data.get("start_frame", 0)550 sample_end_frame = sample_data.get("end_frame", len(sample_video))551 552 # Get actions for the context frames553 context_actions = []554 context_yaw = 0.0555 for frame_idx in range(sample_start_frame + start_idx + 1, 556 min(sample_start_frame + start_idx + self.icl_context_frames + 1, sample_end_frame + 1)):557 if str(frame_idx) in sample_json_data:558 action = sample_json_data[str(frame_idx)]559 new_action = [0.0] * (2 + 2 + 3 + 1 + 2)560 if action['ws'] == 1:561 new_action[0] = 1562 elif action['ws'] == 2:563 new_action[1] = 1564 if action['ad'] == 1:565 new_action[2] = 1566 elif action['ad'] == 2:567 new_action[3] = 1568 if action['scs'] == 1 and action.get("jump_invalid", 0) == 0:569 new_action[4] = 1570 elif action['scs'] == 2:571 new_action[5] = 1572 elif action['scs'] == 3:573 new_action[6] = 1574 if action.get('collision', 0) == 1:575 new_action[7] = 1576 new_action[0] = 0577 new_action[1] = 0578 new_action[2] = 0579 new_action[3] = 0580 pre_pitch = action.get('pre_pitch', 0.0)581 current_pitch = pre_pitch + action.get('pitch_delta', 0.0) * 15.0582 context_yaw += action.get('yaw_delta', 0.0) * 15.0583 new_action[8] = current_pitch584 new_action[9] = context_yaw585 context_actions.append(new_action)586 context_actions_list.extend(context_actions[:len(context_frames)])587 except Exception as e:588 # If loading actions fails, just skip589 pass590 591 if context_frames_list:592 data["context_frames"] = context_frames_list593 if context_actions_list and len(context_actions_list) == len(context_frames_list):594 data["context_actions"] = context_actions_list595 596 return data597 598 599 def __len__(self):600 return len(self.data) * self.repeat601 602 @staticmethod603 def get_one_hot(action, range=2):604 one_hot = [0] * (range + 1)605 one_hot[action] = 1606 return one_hot607 608 609 610import numpy as np611 612 613class CamVideoDataset(torch.utils.data.Dataset):614 """Dataset for Context-as-Memory camera pose conditioned training (ported from VWM).615 616 Loads 81 PNG frames from UE scenes with random temporal cropping and extracts617 corresponding camera poses as 12-dim relative RT vectors subsampled to match618 the 21 latent frames.619 """620 def __init__(621 self,622 base_path=None,623 num_frames=81,624 height=None, width=None,625 max_pixels=1920*1080,626 height_division_factor=16, width_division_factor=16,627 repeat=1,628 args=None,629 cam_position_scale=None,630 ):631 if args is not None:632 base_path = args.dataset_base_path633 height = args.height634 width = args.width635 max_pixels = args.max_pixels636 num_frames = args.num_frames637 repeat = args.dataset_repeat638 cam_position_scale = getattr(args, "cam_position_scale", 0.01)639 self.use_condition_context_frames = getattr(args, "use_condition_context_frames", False)640 self.condition_first_frame = getattr(args, "condition_first_frame", False)641 self.condition_history_keyframes = getattr(args, "condition_history_keyframes", False)642 self.condition_use_camera_pose = getattr(args, "condition_use_camera_pose", True)643 self.num_condition_frames = getattr(args, "num_condition_frames", 1)644 self.condition_frame_mode = getattr(args, "condition_frame_mode", "first_frame_only")645 self.overlap_labels_root = getattr(args, "overlap_labels_root", None)646 self.condition_t2v_ratio = getattr(args, "condition_t2v_ratio", 0.10)647 self.condition_i2v_ratio = getattr(args, "condition_i2v_ratio", 0.10)648 else:649 self.use_condition_context_frames = False650 self.condition_first_frame = False651 self.condition_history_keyframes = False652 self.condition_use_camera_pose = True653 self.num_condition_frames = 1654 self.condition_frame_mode = "first_frame_only"655 self.overlap_labels_root = None656 self.condition_t2v_ratio = 0.10657 self.condition_i2v_ratio = 0.10658 659 if cam_position_scale is None:660 cam_position_scale = 0.01661 self.cam_position_scale = float(cam_position_scale)662 663 self.base_path = base_path664 self.frames_dir = os.path.join(base_path, "frames")665 self.jsons_dir = os.path.join(base_path, "jsons")666 self.num_frames = num_frames667 self.max_pixels = max_pixels668 self.height = height669 self.width = width670 self.height_division_factor = height_division_factor671 self.width_division_factor = width_division_factor672 self.repeat = repeat673 674 if height is not None and width is not None:675 self.dynamic_resolution = False676 else:677 self.dynamic_resolution = True678 679 captions_path = os.path.join(base_path, "captions.txt")680 self.scene_captions = {}681 with open(captions_path, "r") as f:682 for line in f:683 parts = line.strip().split("\t", 1)684 if len(parts) < 2:685 continue686 clip_path, caption = parts687 scene_name = "/".join(clip_path.split("/")[:-1])688 fname = clip_path.split("/")[-1].replace(".mp4", "")689 clip_start = int(fname.split("_")[0])690 if scene_name not in self.scene_captions:691 self.scene_captions[scene_name] = []692 self.scene_captions[scene_name].append((clip_start, caption))693 694 for scene_name in self.scene_captions:695 self.scene_captions[scene_name].sort(key=lambda x: x[0])696 697 self.scene_names = sorted(self.scene_captions.keys())698 self.pose_cache = {}699 self.overlap_cache = {}700 self.invalid_scenes = set()701 self.overlap_labels_root = self._resolve_overlap_labels_root(base_path, self.overlap_labels_root)702 self._validate_condition_config()703 704 total_scenes = len(self.scene_names)705 total_captions = sum(len(v) for v in self.scene_captions.values())706 print(f"CamVideoDataset: {total_scenes} scenes, {total_captions} captions, "707 f"repeat={repeat}, cam_position_scale={self.cam_position_scale}, "708 f"effective length={total_scenes * repeat}")709 710 def _resolve_overlap_labels_root(self, base_path, overlap_labels_root):711 candidate_roots = []712 if overlap_labels_root is not None:713 candidate_roots.append(overlap_labels_root)714 if base_path is not None:715 candidate_roots.append(os.path.join(base_path, "overlap_labels"))716 for root in candidate_roots:717 if root is not None and os.path.isdir(root):718 return root719 return overlap_labels_root720 721 def _validate_condition_config(self):722 if self.condition_t2v_ratio < 0 or self.condition_i2v_ratio < 0:723 raise ValueError("Condition sampling ratios must be non-negative.")724 if self.condition_t2v_ratio + self.condition_i2v_ratio >= 1.0:725 raise ValueError("condition_t2v_ratio + condition_i2v_ratio must be < 1.0.")726 needs_overlap = (727 self.use_condition_context_frames728 and self.condition_frame_mode == "first_plus_overlap"729 and self.condition_history_keyframes730 and self.num_condition_frames > 1731 )732 if needs_overlap and (self.overlap_labels_root is None or not os.path.isdir(self.overlap_labels_root)):733 raise FileNotFoundError(734 "K-frame condition mode requires overlap_labels_root. "735 "Pass --overlap_labels_root or keep overlap_labels under dataset_base_path/overlap_labels."736 )737 738 def _load_scene_poses(self, scene_name):739 if scene_name not in self.pose_cache:740 json_path = os.path.join(self.jsons_dir, scene_name + ".json")741 try:742 with open(json_path, "r") as f:743 data = json.load(f)744 except (FileNotFoundError, json.JSONDecodeError) as e:745 raise ValueError(f"Pose JSON for scene '{scene_name}' is missing or corrupt: {e}")746 if not isinstance(data, dict) or "CineCameraActor" not in data:747 raise ValueError(748 f"Pose JSON for scene '{scene_name}' lacks 'CineCameraActor' key "749 f"(found keys: {list(data.keys()) if isinstance(data, dict) else type(data).__name__})."750 )751 cine = data["CineCameraActor"]752 if not isinstance(cine, dict) or len(cine) == 0:753 raise ValueError(f"Pose JSON for scene '{scene_name}' has empty 'CineCameraActor' entries.")754 self.pose_cache[scene_name] = cine755 return self.pose_cache[scene_name]756 757 def _find_nearest_caption(self, scene_name, start_frame):758 captions = self.scene_captions[scene_name]759 best_idx = 0760 best_dist = abs(captions[0][0] - start_frame)761 for i, (clip_start, _) in enumerate(captions):762 dist = abs(clip_start - start_frame)763 if dist < best_dist:764 best_dist = dist765 best_idx = i766 return captions[best_idx][1]767 768 @staticmethod769 def _compute_rt(position, rotation):770 x, y, z = position771 yaw_rad = np.radians(rotation[2])772 cos_y, sin_y = np.cos(yaw_rad), np.sin(yaw_rad)773 R = np.array([[cos_y, -sin_y, 0], [sin_y, cos_y, 0], [0, 0, 1]])774 return [x, y, z] + R.flatten().tolist()775 776 @staticmethod777 def _to_relative_rt(rt_list, ref_rt):778 R_ref = np.array(ref_rt[3:]).reshape(3, 3)779 T_ref = np.array(ref_rt[:3]).reshape(3, 1)780 R_ref_inv = R_ref.T781 T_ref_inv = -R_ref_inv @ T_ref782 result = []783 for rt in rt_list:784 R_i = np.array(rt[3:]).reshape(3, 3)785 T_i = np.array(rt[:3]).reshape(3, 1)786 R_new = R_ref_inv @ R_i787 T_new = R_ref_inv @ T_i + T_ref_inv788 result.append(T_new.flatten().tolist() + R_new.flatten().tolist())789 return result790 791 def crop_and_resize(self, image, target_height, target_width):792 width, height = image.size793 scale = max(target_width / width, target_height / height)794 image = torchvision.transforms.functional.resize(795 image,796 (round(height * scale), round(width * scale)),797 interpolation=torchvision.transforms.InterpolationMode.BILINEAR798 )799 image = torchvision.transforms.functional.center_crop(image, (target_height, target_width))800 return image801 802 def get_height_width(self, image):803 if self.dynamic_resolution:804 width, height = image.size805 if width * height > self.max_pixels:806 scale = (width * height / self.max_pixels) ** 0.5807 height, width = int(height / scale), int(width / scale)808 height = height // self.height_division_factor * self.height_division_factor809 width = width // self.width_division_factor * self.width_division_factor810 else:811 height, width = self.height, self.width812 return height, width813 814 def _load_resized_frame(self, scene_name, frame_index, target_height, target_width):815 frame_path = os.path.join(self.frames_dir, scene_name, f"{frame_index:04d}.png")816 img = Image.open(frame_path).convert("RGB")817 return self.crop_and_resize(img, target_height, target_width)818 819 def _load_overlap_frames(self, scene_name, frame_index):820 if self.overlap_labels_root is None:821 return []822 cache_key = (scene_name, int(frame_index))823 if cache_key not in self.overlap_cache:824 overlap_path = os.path.join(self.overlap_labels_root, scene_name, f"{int(frame_index)}.json")825 if not os.path.exists(overlap_path):826 self.overlap_cache[cache_key] = []827 else:828 with open(overlap_path, "r") as f:829 overlap_data = json.load(f)830 overlaps = overlap_data.get("overlapping_frames", [])831 self.overlap_cache[cache_key] = [int(idx) for idx in overlaps]832 return self.overlap_cache[cache_key]833 834 def _compute_scene_rt(self, scene_name, frame_index):835 frame_data = self._load_scene_poses(scene_name)[str(int(frame_index))]836 raw_pos = frame_data["position"]837 pos = [float(p) * self.cam_position_scale for p in raw_pos]838 return self._compute_rt(pos, frame_data["rotation"])839 840 def _sample_condition_mode(self):841 if not self.use_condition_context_frames:842 return "disabled"843 if (844 self.condition_frame_mode != "first_plus_overlap"845 or not self.condition_history_keyframes846 or self.num_condition_frames <= 1847 ):848 return "first_frame_only"849 sample = random.random()850 if sample < self.condition_t2v_ratio:851 return "text_only"852 if sample < self.condition_t2v_ratio + self.condition_i2v_ratio:853 return "first_frame_only"854 return "first_plus_overlap"855 856 def _sample_overlap_conditions(self, scene_name, start_frame, ref_rt, target_height, target_width, num_extra_conditions):857 if num_extra_conditions <= 0:858 return [], [], []859 window_indices = set(range(start_frame, start_frame + self.num_frames))860 target_candidates = list(range(start_frame + 1, start_frame + self.num_frames))861 sampled_target_frames = random.sample(target_candidates, k=min(num_extra_conditions, len(target_candidates)))862 overlap_frames = []863 overlap_indices = []864 overlap_actions = []865 used_condition_indices = set()866 for target_frame_idx in sampled_target_frames:867 candidate_indices = [868 idx for idx in self._load_overlap_frames(scene_name, target_frame_idx)869 if idx not in window_indices and idx != target_frame_idx and idx not in used_condition_indices870 ]871 if len(candidate_indices) == 0:872 return None873 chosen_idx = random.choice(candidate_indices)874 used_condition_indices.add(chosen_idx)875 overlap_indices.append(chosen_idx)876 overlap_frames.append(self._load_resized_frame(scene_name, chosen_idx, target_height, target_width))877 if self.condition_use_camera_pose:878 overlap_rt = self._compute_scene_rt(scene_name, chosen_idx)879 overlap_actions.append(self._to_relative_rt([overlap_rt], ref_rt)[0])880 if len(overlap_frames) != num_extra_conditions:881 return None882 return overlap_frames, overlap_indices, overlap_actions883 884 def _try_get_sample(self, scene_name):885 cam_data = self._load_scene_poses(scene_name)886 max_start = len(cam_data) - self.num_frames887 if max_start < 0:888 raise ValueError(f"Scene {scene_name} has fewer than {self.num_frames} frames.")889 start_frame = random.randint(0, max_start)890 end_frame = start_frame + self.num_frames - 1891 892 frames = []893 for i in range(start_frame, end_frame + 1):894 frame_path = os.path.join(self.frames_dir, scene_name, f"{i:04d}.png")895 img = Image.open(frame_path).convert("RGB")896 img = self.crop_and_resize(img, *self.get_height_width(img))897 frames.append(img)898 899 prompt = self._find_nearest_caption(scene_name, start_frame)900 901 rt_list_abs = []902 for i in range(start_frame, end_frame + 1):903 key = str(i)904 if key not in cam_data:905 raise ValueError(f"Scene {scene_name} missing pose for frame {i}.")906 frame_data = cam_data[key]907 raw_pos = frame_data["position"]908 pos = [float(p) * self.cam_position_scale for p in raw_pos]909 rt = self._compute_rt(pos, frame_data["rotation"])910 rt_list_abs.append(rt)911 912 rt_list = self._to_relative_rt(rt_list_abs, rt_list_abs[0])913 pose_indices = list(range(0, self.num_frames, 4))914 actions = [rt_list[i] for i in pose_indices]915 916 return {917 "video": frames,918 "prompt": prompt,919 "actions": actions,920 **self._build_condition_context_payload(921 frames=frames,922 scene_name=scene_name,923 start_frame=start_frame,924 ref_rt=rt_list_abs[0],925 actions=actions,926 ),927 }928 929 def __getitem__(self, data_id):930 n = len(self.scene_names)931 if n == 0:932 raise RuntimeError("CamVideoDataset has no scenes.")933 max_attempts = min(64, n)934 last_error = None935 for attempt in range(max_attempts):936 idx = (data_id + attempt) % n937 scene_name = self.scene_names[idx]938 if scene_name in self.invalid_scenes:939 continue940 try:941 return self._try_get_sample(scene_name)942 except (ValueError, FileNotFoundError, KeyError, OSError) as e:943 self.invalid_scenes.add(scene_name)944 last_error = e945 if attempt < 3 or attempt % 8 == 0:946 print(947 f"[CamVideoDataset] Skipping invalid scene '{scene_name}' "948 f"({type(e).__name__}: {e}); attempt {attempt + 1}/{max_attempts}"949 )950 continue951 raise RuntimeError(952 f"CamVideoDataset: exhausted {max_attempts} attempts starting from index {data_id}; "953 f"last error: {type(last_error).__name__}: {last_error}"954 )955 956 def _build_condition_context_payload(self, frames, scene_name, start_frame, ref_rt, actions):957 if not self.use_condition_context_frames:958 return {}959 payload = {960 "use_condition_context_frames": False,961 "condition_frames": [],962 "condition_frame_indices": [],963 "condition_source": None,964 "condition_actions": [],965 }966 condition_mode = self._sample_condition_mode()967 payload["condition_source"] = condition_mode968 if condition_mode == "text_only":969 return payload970 payload["use_condition_context_frames"] = True971 if self.condition_first_frame:972 payload["condition_frames"].append(frames[0])973 payload["condition_frame_indices"].append(start_frame)974 payload["condition_source"] = "first_frame_only"975 if self.condition_use_camera_pose and actions:976 payload["condition_actions"].append(list(actions[0]))977 if (978 condition_mode == "first_plus_overlap"979 and self.condition_history_keyframes980 and self.num_condition_frames > len(payload["condition_frames"])981 ):982 num_extra_conditions = self.num_condition_frames - len(payload["condition_frames"])983 overlap_payload = self._sample_overlap_conditions(984 scene_name=scene_name,985 start_frame=start_frame,986 ref_rt=ref_rt,987 target_height=frames[0].size[1],988 target_width=frames[0].size[0],989 num_extra_conditions=num_extra_conditions,990 )991 if overlap_payload is None:992 return payload993 overlap_frames, overlap_indices, overlap_actions = overlap_payload994 payload["condition_frames"].extend(overlap_frames)995 payload["condition_frame_indices"].extend(overlap_indices)996 if self.condition_use_camera_pose:997 payload["condition_actions"].extend(overlap_actions)998 payload["condition_source"] = "first_plus_overlap"999 return payload1000 1001 def __len__(self):1002 return len(self.scene_names) * self.repeat1003 1004 1005class DiffusionTrainingModule(torch.nn.Module):1006 def __init__(self):1007 super().__init__()1008 1009 1010 def to(self, *args, **kwargs):1011 for name, model in self.named_children():1012 model.to(*args, **kwargs)1013 return self1014 1015 1016 def trainable_modules(self):1017 trainable_modules = filter(lambda p: p.requires_grad, self.parameters())1018 return trainable_modules1019 1020 1021 def trainable_param_names(self):1022 trainable_param_names = list(filter(lambda named_param: named_param[1].requires_grad, self.named_parameters()))1023 trainable_param_names = set([named_param[0] for named_param in trainable_param_names])1024 return trainable_param_names1025 1026 1027 def add_lora_to_model(self, model, target_modules, lora_rank, lora_alpha=None):1028 if lora_alpha is None:1029 lora_alpha = lora_rank1030 lora_config = LoraConfig(r=lora_rank, lora_alpha=lora_alpha, target_modules=target_modules)1031 model = inject_adapter_in_model(lora_config, model)1032 return model1033 1034 1035 def export_trainable_state_dict(self, state_dict, remove_prefix=None):1036 trainable_param_names = self.trainable_param_names()1037 state_dict = {name: param for name, param in state_dict.items() if name in trainable_param_names}1038 if remove_prefix is not None:1039 state_dict_ = {}1040 for name, param in state_dict.items():1041 if name.startswith(remove_prefix):1042 name = name[len(remove_prefix):]1043 state_dict_[name] = param1044 state_dict = state_dict_1045 return state_dict1046 1047 1048 1049class ModelLogger:1050 def __init__(self, output_path, remove_prefix_in_ckpt=None, state_dict_converter=lambda x:x):1051 self.output_path = output_path1052 self.remove_prefix_in_ckpt = remove_prefix_in_ckpt1053 self.state_dict_converter = state_dict_converter1054 1055 1056 def on_step_end(self, loss):1057 pass1058 1059 1060 def on_epoch_end(self, accelerator, model, epoch_id):1061 accelerator.wait_for_everyone()1062 if accelerator.is_main_process:1063 state_dict = accelerator.get_state_dict(model)1064 state_dict = accelerator.unwrap_model(model).export_trainable_state_dict(state_dict, remove_prefix=self.remove_prefix_in_ckpt)1065 state_dict = self.state_dict_converter(state_dict)1066 os.makedirs(self.output_path, exist_ok=True)1067 path = os.path.join(self.output_path, f"epoch-{epoch_id}.safetensors")1068 accelerator.save(state_dict, path, safe_serialization=True)1069 1070 1071 1072def launch_training_task(1073 dataset: torch.utils.data.Dataset,1074 model: DiffusionTrainingModule,1075 model_logger: ModelLogger,1076 optimizer: torch.optim.Optimizer,1077 scheduler: torch.optim.lr_scheduler.LRScheduler,1078 num_epochs: int = 1,1079 gradient_accumulation_steps: int = 1,1080):1081 dataloader = torch.utils.data.DataLoader(dataset, shuffle=True, collate_fn=lambda x: x[0], drop_last=True)1082 accelerator = Accelerator(gradient_accumulation_steps=gradient_accumulation_steps)1083 model, optimizer, dataloader, scheduler = accelerator.prepare(model, optimizer, dataloader, scheduler)1084 1085 for epoch_id in range(num_epochs):1086 for data in tqdm(dataloader):1087 with accelerator.accumulate(model):1088 optimizer.zero_grad()1089 loss = model(data)1090 accelerator.backward(loss)1091 optimizer.step()1092 model_logger.on_step_end(loss)1093 scheduler.step()1094 model_logger.on_epoch_end(accelerator, model, epoch_id)1095 1096def launch_data_process_task(model: DiffusionTrainingModule, dataset, output_path="./models"):1097 dataloader = torch.utils.data.DataLoader(dataset, shuffle=False, collate_fn=lambda x: x[0], drop_last=True)1098 accelerator = Accelerator()1099 model, dataloader = accelerator.prepare(model, dataloader)1100 os.makedirs(os.path.join(output_path, "data_cache"), exist_ok=True)1101 for data_id, data in enumerate(tqdm(dataloader)):1102 with torch.no_grad():1103 inputs = model.forward_preprocess(data)1104 inputs = {key: inputs[key] for key in model.model_input_keys if key in inputs}1105 torch.save(inputs, os.path.join(output_path, "data_cache", f"{data_id}.pth"))1106 1107 1108 1109def wan_parser():1110 parser = argparse.ArgumentParser(description="Simple example of a training script.")1111 parser.add_argument("--dataset_base_path", type=str, default="", required=True, help="Base path of the dataset.")1112 parser.add_argument("--dataset_metadata_path", type=str, default=None, help="Path to the metadata file of the dataset.")1113 parser.add_argument("--max_pixels", type=int, default=1280*720, help="Maximum number of pixels per frame, used for dynamic resolution..")1114 parser.add_argument("--height", type=int, default=None, help="Height of images or videos. Leave `height` and `width` empty to enable dynamic resolution.")1115 parser.add_argument("--width", type=int, default=None, help="Width of images or videos. Leave `height` and `width` empty to enable dynamic resolution.")1116 parser.add_argument("--num_frames", type=int, default=81, help="Number of frames per video. Frames are sampled from the video prefix.")1117 parser.add_argument("--data_file_keys", type=str, default="image,video", help="Data file keys in the metadata. Comma-separated.")1118 parser.add_argument("--dataset_repeat", type=int, default=1, help="Number of times to repeat the dataset per epoch.")1119 parser.add_argument("--model_paths", type=str, default=None, help="Paths to load models. In JSON format.")1120 parser.add_argument("--model_id_with_origin_paths", type=str, default=None, help="Model ID with origin paths, e.g., Wan-AI/Wan2.1-T2V-1.3B:diffusion_pytorch_model*.safetensors. Comma-separated.")1121 parser.add_argument("--learning_rate", type=float, default=1e-4, help="Learning rate.")1122 parser.add_argument("--num_epochs", type=int, default=1, help="Number of epochs.")1123 parser.add_argument("--output_path", type=str, default="./models", help="Output save path.")1124 parser.add_argument("--remove_prefix_in_ckpt", type=str, default="pipe.dit.", help="Remove prefix in ckpt.")1125 parser.add_argument("--trainable_models", type=str, default=None, help="Models to train, e.g., dit, vae, text_encoder.")1126 parser.add_argument("--lora_base_model", type=str, default=None, help="Which model LoRA is added to.")1127 parser.add_argument("--lora_target_modules", type=str, default="q,k,v,o,ffn.0,ffn.2", help="Which layers LoRA is added to.")1128 parser.add_argument("--lora_rank", type=int, default=32, help="Rank of LoRA.")1129 parser.add_argument("--extra_inputs", default=None, help="Additional model inputs, comma-separated.")1130 parser.add_argument("--use_gradient_checkpointing_offload", default=False, action="store_true", help="Whether to offload gradient checkpointing to CPU memory.")1131 parser.add_argument("--gradient_accumulation_steps", type=int, default=1, help="Gradient accumulation steps.")1132 parser.add_argument("--use_condition_context_frames", default=False, action="store_true", help="Enable appended clean condition latents.")1133 parser.add_argument("--condition_first_frame", default=False, action="store_true", help="Use the current clip first frame as a clean condition frame.")1134 parser.add_argument("--condition_history_keyframes", default=False, action="store_true", help="Use overlap-based keyframes as conditions.")1135 parser.add_argument("--condition_use_camera_pose", default=True, action="store_true", help="Inject camera pose for condition frames.")1136 parser.add_argument("--num_condition_frames", type=int, default=1, help="Number of condition frames.")1137 parser.add_argument("--condition_frame_mode", type=str, default="first_frame_only", help="Condition frame selection mode.")1138 parser.add_argument("--overlap_labels_root", type=str, default=None, help="Root dir for overlap label JSONs.")1139 parser.add_argument("--condition_t2v_ratio", type=float, default=0.10, help="Ratio of text-only condition samples.")1140 parser.add_argument("--condition_i2v_ratio", type=float, default=0.10, help="Ratio of first-frame-only condition samples.")1141 return parser1142 1143 1144 1145def flux_parser():1146 parser = argparse.ArgumentParser(description="Simple example of a training script.")1147 parser.add_argument("--dataset_base_path", type=str, default="", required=True, help="Base path of the dataset.")1148 parser.add_argument("--dataset_metadata_path", type=str, default=None, help="Path to the metadata file of the dataset.")1149 parser.add_argument("--max_pixels", type=int, default=1024*1024, help="Maximum number of pixels per frame, used for dynamic resolution..")1150 parser.add_argument("--height", type=int, default=None, help="Height of images. Leave `height` and `width` empty to enable dynamic resolution.")1151 parser.add_argument("--width", type=int, default=None, help="Width of images. Leave `height` and `width` empty to enable dynamic resolution.")1152 parser.add_argument("--data_file_keys", type=str, default="image", help="Data file keys in the metadata. Comma-separated.")1153 parser.add_argument("--dataset_repeat", type=int, default=1, help="Number of times to repeat the dataset per epoch.")1154 parser.add_argument("--model_paths", type=str, default=None, help="Paths to load models. In JSON format.")1155 parser.add_argument("--model_id_with_origin_paths", type=str, default=None, help="Model ID with origin paths, e.g., Wan-AI/Wan2.1-T2V-1.3B:diffusion_pytorch_model*.safetensors. Comma-separated.")1156 parser.add_argument("--learning_rate", type=float, default=1e-4, help="Learning rate.")1157 parser.add_argument("--num_epochs", type=int, default=1, help="Number of epochs.")1158 parser.add_argument("--output_path", type=str, default="./models", help="Output save path.")1159 parser.add_argument("--remove_prefix_in_ckpt", type=str, default="pipe.dit.", help="Remove prefix in ckpt.")1160 parser.add_argument("--trainable_models", type=str, default=None, help="Models to train, e.g., dit, vae, text_encoder.")1161 parser.add_argument("--lora_base_model", type=str, default=None, help="Which model LoRA is added to.")1162 parser.add_argument("--lora_target_modules", type=str, default="q,k,v,o,ffn.0,ffn.2", help="Which layers LoRA is added to.")1163 parser.add_argument("--lora_rank", type=int, default=32, help="Rank of LoRA.")1164 parser.add_argument("--extra_inputs", default=None, help="Additional model inputs, comma-separated.")1165 parser.add_argument("--align_to_opensource_format", default=False, action="store_true", help="Whether to align the lora format to opensource format. Only for DiT's LoRA.")1166 parser.add_argument("--use_gradient_checkpointing", default=False, action="store_true", help="Whether to use gradient checkpointing.")1167 parser.add_argument("--use_gradient_checkpointing_offload", default=False, action="store_true", help="Whether to offload gradient checkpointing to CPU memory.")1168 parser.add_argument("--gradient_accumulation_steps", type=int, default=1, help="Gradient accumulation steps.")1169 return parser1170 