Compare commits

...
Author SHA1 Message Date
SolitaryThinker e9d95b1c10 wip 2025-12-12 06:30:08 +00:00
SolitaryThinker 8b55e9706c wip 2025-12-10 01:59:07 +00:00
6 changed files with 272 additions and 30 deletions
+20
View File
@@ -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",
+41
View File
@@ -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
+66 -8
View File
@@ -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()