Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e9d95b1c10 | ||
|
|
8b55e9706c |
@@ -83,6 +83,9 @@ class PreprocessConfig:
|
||||
speed_factor: float = 1.0
|
||||
drop_short_ratio: float = 1.0
|
||||
do_temporal_sample: bool = False
|
||||
enable_smart_resize: bool = False
|
||||
smart_resize_max_area: int | None = None
|
||||
hw_aspect_threshold: float = 1.5
|
||||
|
||||
# Model configuration
|
||||
training_cfg_rate: float = 0.0
|
||||
@@ -184,6 +187,23 @@ class PreprocessConfig:
|
||||
action=StoreBoolean,
|
||||
default=PreprocessConfig.do_temporal_sample,
|
||||
help="Whether to do temporal sampling")
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}enable-smart-resize",
|
||||
action=StoreBoolean,
|
||||
default=PreprocessConfig.enable_smart_resize,
|
||||
help="Whether to enable smart resizing")
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}smart-resize-max-area",
|
||||
type=int,
|
||||
default=PreprocessConfig.smart_resize_max_area,
|
||||
help="Maximum area for smart resizing")
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}hw-aspect-threshold",
|
||||
type=float,
|
||||
default=PreprocessConfig.hw_aspect_threshold,
|
||||
help=
|
||||
"Height/Width aspect ratio threshold. Allowed range is [1/threshold * target_aspect, threshold * target_aspect]."
|
||||
)
|
||||
|
||||
# Model Training configuration
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}training-cfg-rate",
|
||||
|
||||
@@ -152,3 +152,44 @@ class TemporalRandomCrop:
|
||||
begin_index = random.randint(0, rand_end)
|
||||
end_index = min(begin_index + self.size, total_frames)
|
||||
return begin_index, end_index
|
||||
|
||||
|
||||
def best_output_size(
|
||||
width: int,
|
||||
height: int,
|
||||
width_stride: int,
|
||||
height_stride: int,
|
||||
max_area: int,
|
||||
) -> tuple[int, int]:
|
||||
"""
|
||||
Calculate the best output size (width, height) given the original dimensions, strides and max area.
|
||||
The aspect ratio is preserved as much as possible.
|
||||
|
||||
Args:
|
||||
width (int): Original width
|
||||
height (int): Original height
|
||||
width_stride (int): Width stride requirement
|
||||
height_stride (int): Height stride requirement
|
||||
max_area (int): Maximum allowed area (width * height)
|
||||
|
||||
Returns:
|
||||
tuple[int, int]: (new_width, new_height)
|
||||
"""
|
||||
aspect_ratio = width / height
|
||||
|
||||
# Scale dimensions if they exceed max_area
|
||||
current_area = width * height
|
||||
if current_area > max_area:
|
||||
scale = (max_area / current_area)**0.5
|
||||
width = int(width * scale)
|
||||
height = int(height * scale)
|
||||
|
||||
# Round to the nearest multiple of stride
|
||||
width = round(width / width_stride) * width_stride
|
||||
height = round(height / height_stride) * height_stride
|
||||
|
||||
# Ensure dimensions are at least one stride
|
||||
width = max(width, width_stride)
|
||||
height = max(height, height_stride)
|
||||
|
||||
return width, height
|
||||
|
||||
@@ -5,13 +5,19 @@ from typing import cast
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
import torchvision.transforms.functional as TF
|
||||
from PIL import Image
|
||||
from einops import rearrange
|
||||
from torchvision import transforms
|
||||
|
||||
from fastvideo.configs.configs import VideoLoaderType
|
||||
from fastvideo.dataset.transform import (CenterCropResizeVideo,
|
||||
TemporalRandomCrop)
|
||||
TemporalRandomCrop, best_output_size)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
from fastvideo.pipelines.pipeline_batch_info import (ForwardBatch,
|
||||
PreprocessBatch)
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
@@ -41,6 +47,7 @@ class VideoTransformStage(PipelineStage):
|
||||
batch = cast(PreprocessBatch, batch)
|
||||
assert isinstance(batch.fps, list)
|
||||
assert isinstance(batch.num_frames, list)
|
||||
assert fastvideo_args.preprocess_config is not None
|
||||
|
||||
if batch.data_type != "video":
|
||||
return batch
|
||||
@@ -49,8 +56,17 @@ class VideoTransformStage(PipelineStage):
|
||||
raise ValueError("Video loader is not set")
|
||||
|
||||
video_pixel_batch = []
|
||||
pil_image_batch = []
|
||||
|
||||
enable_smart_resize = fastvideo_args.preprocess_config.enable_smart_resize
|
||||
smart_resize_max_area = fastvideo_args.preprocess_config.smart_resize_max_area
|
||||
if smart_resize_max_area is None:
|
||||
smart_resize_max_area = 480 * 832
|
||||
|
||||
calculated_size = None
|
||||
|
||||
for i in range(len(batch.video_loader)):
|
||||
# logger.info(f"Processing video {i+1}/{len(batch.video_loader)}")
|
||||
frame_interval = batch.fps[i] / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, batch.num_frames[i],
|
||||
@@ -63,8 +79,28 @@ class VideoTransformStage(PipelineStage):
|
||||
else:
|
||||
frame_indices = frame_indices[:self.num_frames]
|
||||
|
||||
logger.info(
|
||||
f"Frame indices selected (count={len(frame_indices)}): [{frame_indices[0]}, ..., {frame_indices[-1]}]"
|
||||
)
|
||||
|
||||
if fastvideo_args.preprocess_config.video_loader_type == VideoLoaderType.TORCHCODEC:
|
||||
video = batch.video_loader[i].get_frames_at(frame_indices).data
|
||||
try:
|
||||
video = batch.video_loader[i].get_frames_at(
|
||||
frame_indices).data
|
||||
except Exception as e:
|
||||
# Try to get filename if available in PreprocessBatch
|
||||
video_path = "unknown"
|
||||
print(f"batch: {batch}")
|
||||
if isinstance(batch, PreprocessBatch) and hasattr(
|
||||
batch, 'video_file_name') and i < len(
|
||||
batch.video_file_name):
|
||||
video_path = batch.video_file_name[i]
|
||||
|
||||
logger.error(
|
||||
f"Failed to load frames for video {video_path}: {e}")
|
||||
logger.error(
|
||||
f"Attempting to load frame indices: {frame_indices}")
|
||||
raise e
|
||||
elif fastvideo_args.preprocess_config.video_loader_type == VideoLoaderType.TORCHVISION:
|
||||
video, _, _ = torchvision.io.read_video(batch.video_loader[i],
|
||||
output_format="TCHW")
|
||||
@@ -73,16 +109,75 @@ class VideoTransformStage(PipelineStage):
|
||||
raise ValueError(
|
||||
f"Invalid video loader type: {fastvideo_args.preprocess_config.video_loader_type}"
|
||||
)
|
||||
video = self.video_transform(video)
|
||||
video_pixel_batch.append(video)
|
||||
|
||||
logger.info(f"Video tensor shape after loading: {video.shape}")
|
||||
|
||||
if enable_smart_resize:
|
||||
if calculated_size is None:
|
||||
_, _, h_in, w_in = video.shape
|
||||
# Get config values
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||
dh, dw = patch_size[1] * vae_stride, patch_size[
|
||||
2] * vae_stride
|
||||
|
||||
ow, oh = best_output_size(w_in, h_in, dw, dh,
|
||||
smart_resize_max_area)
|
||||
calculated_size = (oh, ow)
|
||||
logger.info(
|
||||
f"Smart resize: input=({h_in}, {w_in}), output=({oh}, {ow})"
|
||||
)
|
||||
|
||||
# Resize video frames using CenterCropResizeVideo (efficient)
|
||||
processed_video = CenterCropResizeVideo(calculated_size)(video)
|
||||
logger.info(
|
||||
f"Processed video shape after resize: {processed_video.shape}"
|
||||
)
|
||||
video_pixel_batch.append(processed_video)
|
||||
|
||||
# Process pil_image (condition) with high quality Lanczos if I2V
|
||||
if fastvideo_args.workload_type == WorkloadType.I2V:
|
||||
# Extract first frame
|
||||
img_tensor = video[0] # C, H, W
|
||||
img = TF.to_pil_image(img_tensor)
|
||||
iw, ih = img.width, img.height
|
||||
ow, oh = calculated_size[1], calculated_size[0]
|
||||
|
||||
# Smart Resize logic for PIL image
|
||||
scale = max(ow / iw, oh / ih)
|
||||
resampling = Image.Resampling.LANCZOS if hasattr(
|
||||
Image, 'Resampling') else Image.LANCZOS
|
||||
img = img.resize((round(iw * scale), round(ih * scale)),
|
||||
resampling)
|
||||
|
||||
# center-crop
|
||||
x1 = (img.width - ow) // 2
|
||||
y1 = (img.height - oh) // 2
|
||||
img = img.crop((x1, y1, x1 + ow, y1 + oh))
|
||||
|
||||
# to tensor [0, 255] uint8
|
||||
img_t = torch.from_numpy(np.array(img)).permute(
|
||||
2, 0, 1).unsqueeze(0)
|
||||
pil_image_batch.append(img_t)
|
||||
|
||||
else:
|
||||
video = self.video_transform(video)
|
||||
video_pixel_batch.append(video)
|
||||
|
||||
video_pixel_values = torch.stack(video_pixel_batch)
|
||||
logger.info(
|
||||
f"Final stacked video batch shape: {video_pixel_values.shape}")
|
||||
video_pixel_values = rearrange(video_pixel_values,
|
||||
"b t c h w -> b c t h w")
|
||||
video_pixel_values = video_pixel_values.to(torch.uint8)
|
||||
|
||||
if fastvideo_args.workload_type == WorkloadType.I2V:
|
||||
batch.pil_image = video_pixel_values[:, :, 0, :, :]
|
||||
if enable_smart_resize and len(pil_image_batch) > 0:
|
||||
batch.pil_image = torch.cat(
|
||||
pil_image_batch, dim=0).to(self.device if hasattr(
|
||||
self, 'device') else video_pixel_values.device)
|
||||
else:
|
||||
batch.pil_image = video_pixel_values[:, :, 0, :, :]
|
||||
|
||||
video_pixel_values = video_pixel_values.float() / 255.0
|
||||
batch.latents = video_pixel_values
|
||||
|
||||
@@ -71,6 +71,7 @@ class PreprocessingDataValidator:
|
||||
|
||||
for name, validator in self.validators.items():
|
||||
if not validator(batch):
|
||||
logger.info(f"Failed validation for {name}")
|
||||
self.filter_counts[name] += 1
|
||||
return False
|
||||
|
||||
@@ -87,6 +88,8 @@ class PreprocessingDataValidator:
|
||||
"""Validate resolution constraints"""
|
||||
|
||||
aspect = self.max_height / self.max_width
|
||||
height = None
|
||||
width = None
|
||||
if batch["resolution"] is not None:
|
||||
height = batch["resolution"].get("height", None)
|
||||
width = batch["resolution"].get("width", None)
|
||||
@@ -94,12 +97,15 @@ class PreprocessingDataValidator:
|
||||
if height is None or width is None:
|
||||
return False
|
||||
|
||||
return self._filter_resolution(
|
||||
ret = self._filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=self.hw_aspect_threshold * aspect,
|
||||
min_h_div_w_ratio=1 / self.hw_aspect_threshold * aspect,
|
||||
)
|
||||
if not ret:
|
||||
logger.info(f"failed in resolution: {batch['caption']}")
|
||||
return ret
|
||||
|
||||
def _filter_resolution(self, h: int, w: int, max_h_div_w_ratio: float,
|
||||
min_h_div_w_ratio: float) -> bool:
|
||||
@@ -113,14 +119,19 @@ class PreprocessingDataValidator:
|
||||
if (batch["num_frames"] / batch["fps"]
|
||||
> self.video_length_tolerance_range *
|
||||
(self.num_frames / self.train_fps * self.speed_factor)):
|
||||
logger.info("Failed in 1")
|
||||
return False
|
||||
|
||||
frame_interval = batch["fps"] / self.train_fps
|
||||
frame_interval = (batch["fps"] / self.train_fps) * self.speed_factor
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, batch["num_frames"],
|
||||
frame_interval).astype(int)
|
||||
return not (len(frame_indices) < self.num_frames
|
||||
and random.random() < self.drop_short_ratio)
|
||||
# logger.info("Failed in 2")
|
||||
result = not (len(frame_indices) < self.num_frames
|
||||
and random.random() < self.drop_short_ratio)
|
||||
if not result:
|
||||
logger.info(f"failed in frame_sampling: {batch['caption']}")
|
||||
return result
|
||||
|
||||
def log_validation_stats(self):
|
||||
info = ""
|
||||
@@ -280,9 +291,46 @@ def build_dataset(preprocess_config: PreprocessConfig, split: str,
|
||||
dataset = dataset.shard(num_shards=get_world_size(),
|
||||
index=get_world_rank())
|
||||
elif preprocess_config.dataset_type == DatasetType.MERGED:
|
||||
metadata_json_path = os.path.join(preprocess_config.dataset_path,
|
||||
"videos2caption.json")
|
||||
video_folder = os.path.join(preprocess_config.dataset_path, "videos")
|
||||
merge_txt_path = os.path.join(preprocess_config.dataset_path,
|
||||
"merge.txt")
|
||||
if os.path.exists(merge_txt_path):
|
||||
logger.info(f"Found merge.txt at {merge_txt_path}")
|
||||
with open(merge_txt_path) as f:
|
||||
line = f.read().strip()
|
||||
if "," not in line:
|
||||
raise ValueError(
|
||||
f"Invalid format in {merge_txt_path}: expected 'video_folder,metadata_json_path'"
|
||||
)
|
||||
video_folder, metadata_json_path = line.split(",", 1)
|
||||
video_folder = video_folder.strip()
|
||||
metadata_json_path = metadata_json_path.strip()
|
||||
|
||||
if not os.path.isabs(video_folder):
|
||||
video_folder = os.path.join(preprocess_config.dataset_path,
|
||||
video_folder)
|
||||
if not os.path.isabs(metadata_json_path):
|
||||
metadata_json_path = os.path.join(
|
||||
preprocess_config.dataset_path, metadata_json_path)
|
||||
else:
|
||||
logger.info(
|
||||
f"merge.txt not found at {merge_txt_path}, using default paths")
|
||||
metadata_json_path = os.path.join(preprocess_config.dataset_path,
|
||||
"videos2caption.json")
|
||||
video_folder = os.path.join(preprocess_config.dataset_path,
|
||||
"videos")
|
||||
|
||||
if not os.path.exists(metadata_json_path):
|
||||
logger.error(f"Metadata file not found: {metadata_json_path}")
|
||||
raise FileNotFoundError(
|
||||
f"Metadata file not found: {metadata_json_path}")
|
||||
|
||||
if not os.path.exists(video_folder):
|
||||
logger.error(f"Video folder not found: {video_folder}")
|
||||
raise FileNotFoundError(f"Video folder not found: {video_folder}")
|
||||
|
||||
logger.info(f"Using metadata file: {metadata_json_path}")
|
||||
logger.info(f"Using video folder: {video_folder}")
|
||||
|
||||
dataset = load_dataset("json",
|
||||
data_files=metadata_json_path,
|
||||
split=split)
|
||||
@@ -293,9 +341,18 @@ def build_dataset(preprocess_config: PreprocessConfig, split: str,
|
||||
if "path" in column_names:
|
||||
dataset = dataset.rename_column("path", "name")
|
||||
|
||||
dataset = dataset.filter(validator)
|
||||
print(f"Length of dataset before filtering: {len(dataset)}")
|
||||
if len(dataset) > 0:
|
||||
print(f"DEBUG: First item in dataset: {dataset[0]}")
|
||||
|
||||
# Disable caching to ensure our print statements run
|
||||
dataset = dataset.filter(validator, load_from_cache_file=False)
|
||||
|
||||
validator.log_validation_stats()
|
||||
print(f"Length of dataset after filtering: {len(dataset)}")
|
||||
dataset = dataset.shard(num_shards=get_world_size(),
|
||||
index=get_world_rank())
|
||||
print(f"Length of dataset after sharding: {len(dataset)}")
|
||||
|
||||
# add video column
|
||||
def add_video_column(item: dict[str, Any]) -> dict[str, Any]:
|
||||
@@ -303,6 +360,7 @@ def build_dataset(preprocess_config: PreprocessConfig, split: str,
|
||||
return item
|
||||
|
||||
dataset = dataset.map(add_video_column)
|
||||
print(f"Length of dataset after mapping: {len(dataset)}")
|
||||
if preprocess_config.video_loader_type == VideoLoaderType.TORCHCODEC:
|
||||
dataset = dataset.cast_column("video", Video())
|
||||
else:
|
||||
|
||||
@@ -40,6 +40,7 @@ class PreprocessWorkflow(WorkflowBase):
|
||||
video_length_tolerance_range=preprocess_config.
|
||||
video_length_tolerance_range,
|
||||
drop_short_ratio=preprocess_config.drop_short_ratio,
|
||||
hw_aspect_threshold=preprocess_config.hw_aspect_threshold,
|
||||
)
|
||||
self.add_component("raw_data_validator", raw_data_validator)
|
||||
|
||||
|
||||
@@ -4,17 +4,37 @@ import os
|
||||
import random
|
||||
|
||||
|
||||
def generate_merged_validation_json(input_dir, output_file):
|
||||
# read in video2caption.json
|
||||
with open(os.path.join(input_dir, "video2caption_replace.json"), "r") as f:
|
||||
def generate_merged_validation_json(args):
|
||||
input_file = args.input_file
|
||||
output_validation_file = args.output_validation_file
|
||||
|
||||
if args.output_train_file:
|
||||
output_train_file = args.output_train_file
|
||||
else:
|
||||
base, ext = os.path.splitext(input_file)
|
||||
output_train_file = f"{base}_train{ext}"
|
||||
|
||||
# read in input json
|
||||
print(f"Reading from {input_file}")
|
||||
with open(input_file, "r") as f:
|
||||
video2caption = json.load(f)
|
||||
|
||||
# count how many elements are in the list
|
||||
num_elements = len(video2caption)
|
||||
print(f"Number of elements in video2caption.json: {num_elements}")
|
||||
print(f"Number of elements in input file: {num_elements}")
|
||||
|
||||
# randomly sample 64 elements from the list
|
||||
sampled_elements = random.sample(video2caption, 64)
|
||||
# randomly sample elements from the list
|
||||
num_sample = min(args.num_elements, num_elements)
|
||||
indices = set(random.sample(range(num_elements), num_sample))
|
||||
|
||||
sampled_elements = []
|
||||
remaining_elements = []
|
||||
|
||||
for i in range(num_elements):
|
||||
if i in indices:
|
||||
sampled_elements.append(video2caption[i])
|
||||
else:
|
||||
remaining_elements.append(video2caption[i])
|
||||
|
||||
# Transform sampled elements into validation.json format
|
||||
validation_data = []
|
||||
@@ -23,10 +43,10 @@ def generate_merged_validation_json(input_dir, output_file):
|
||||
validation_entry = {
|
||||
"caption": element["cap"],
|
||||
"video_path": element.get("path", ""),
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
"num_inference_steps": args.num_inference_steps,
|
||||
"height": args.height,
|
||||
"width": args.width,
|
||||
"num_frames": args.num_frames
|
||||
}
|
||||
validation_data.append(validation_entry)
|
||||
|
||||
@@ -36,18 +56,25 @@ def generate_merged_validation_json(input_dir, output_file):
|
||||
}
|
||||
|
||||
# Write the validation JSON to the output file
|
||||
with open(output_file, "w") as f:
|
||||
with open(output_validation_file, "w") as f:
|
||||
json.dump(validation_json, f, indent=2)
|
||||
|
||||
print(f"Generated validation JSON with {len(validation_data)} entries and saved to {output_file}")
|
||||
print(f"Generated validation JSON with {len(validation_data)} entries and saved to {output_validation_file}")
|
||||
|
||||
# Write the remaining JSON to the output train file
|
||||
with open(output_train_file, "w") as f:
|
||||
json.dump(remaining_elements, f, indent=2)
|
||||
|
||||
print(f"Saved remaining {len(remaining_elements)} entries to {output_train_file}")
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset_type: "mixkit"
|
||||
# dataset_type: "merged"
|
||||
parser.add_argument("--dataset_type", choices=["merged"], required=True)
|
||||
parser.add_argument("--input_dir", type=str, required=True)
|
||||
parser.add_argument("--output_file", type=str, required=True)
|
||||
parser.add_argument("--input_file", type=str, required=True, help="Path to input json file")
|
||||
parser.add_argument("--output_validation_file", type=str, required=True, help="Path to output validation json file")
|
||||
parser.add_argument("--output_train_file", type=str, help="Path to output train json file (remaining data). Defaults to {input_filename}_train.json")
|
||||
parser.add_argument("--num_elements", type=int, default=64)
|
||||
parser.add_argument("--num_frames", type=int, default=77)
|
||||
parser.add_argument("--height", type=int, default=480)
|
||||
@@ -56,8 +83,8 @@ def main():
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.dataset_type == "merged":
|
||||
generate_merged_validation_json(args.input_dir, args.output_file)
|
||||
generate_merged_validation_json(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
Reference in New Issue
Block a user