Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
utils.py1170 linesDownload Raw Back to trainers
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