Recommit the conditioning code.

This commit is contained in:
CrossProduct
2025-01-16 05:03:44 +00:00
parent b341b79652
commit 315431be0a
4 changed files with 160 additions and 33 deletions
+10
View File
@@ -249,6 +249,8 @@ class Args:
dataset_file: Optional[str] = None
video_column: str = None
caption_column: str = None
pose_column:str = None
id_token: Optional[str] = None
image_resolution_buckets: List[Tuple[int, int]] = None
video_resolution_buckets: List[Tuple[int, int, int]] = None
@@ -353,6 +355,7 @@ class Args:
"dataset_file": self.dataset_file,
"video_column": self.video_column,
"caption_column": self.caption_column,
"pose_column": self.pose_column,
"id_token": self.id_token,
"image_resolution_buckets": self.image_resolution_buckets,
"video_resolution_buckets": self.video_resolution_buckets,
@@ -550,6 +553,12 @@ def _add_dataset_arguments(parser: argparse.ArgumentParser) -> None:
default="text",
help="The column of the dataset containing the instance prompt for each video. Or, the name of the file in `--data_root` folder containing the line-separated instance prompts.",
)
parser.add_argument(
"--pose_column",
type=str,
default="text",
help="The column of the dataset containing the instance prompt for each video. Or, the name of the file in `--data_root` folder containing the line-separated instance prompts.",
)
parser.add_argument(
"--id_token",
type=str,
@@ -1006,6 +1015,7 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args:
result_args.dataset_file = args.dataset_file
result_args.video_column = args.video_column
result_args.caption_column = args.caption_column
result_args.pose_column = args.pose_column
result_args.id_token = args.id_token
result_args.image_resolution_buckets = args.image_resolution_buckets or DEFAULT_IMAGE_RESOLUTION_BUCKETS
result_args.video_resolution_buckets = args.video_resolution_buckets or DEFAULT_VIDEO_RESOLUTION_BUCKETS
+57 -12
View File
@@ -46,6 +46,7 @@ class ImageOrVideoDataset(Dataset):
data_root: str,
caption_column: str,
video_column: str,
pose_column: str,
resolution_buckets: List[Tuple[int, int, int]],
dataset_file: Optional[str] = None,
id_token: Optional[str] = None,
@@ -55,8 +56,11 @@ class ImageOrVideoDataset(Dataset):
self.data_root = Path(data_root)
self.dataset_file = dataset_file
self.caption_column = caption_column
self.video_column = video_column
self.pose_column = pose_column
self.id_token = f"{id_token.strip()} " if id_token else ""
self.resolution_buckets = resolution_buckets
@@ -72,6 +76,7 @@ class ImageOrVideoDataset(Dataset):
(
self.prompts,
self.video_paths,
self.pose_pathes,
) = self._load_dataset_from_local_path()
elif dataset_file.endswith(".csv"):
(
@@ -141,15 +146,28 @@ class ImageOrVideoDataset(Dataset):
else:
video = self._preprocess_video(video_path)
return {
"prompt": prompt,
"video": video,
"video_metadata": {
"num_frames": video.shape[0],
"height": video.shape[2],
"width": video.shape[3],
},
}
if self.pose_condition_column != None:
pose_video = self._preprocess_video(self.pose_paths[index])
return {
"prompt": prompt,
"video": video,
"pose_video": pose_video,
"video_metadata": {
"num_frames": video.shape[0],
"height": video.shape[2],
"width": video.shape[3],
},
}
else:
return {
"prompt": prompt,
"video": video,
"video_metadata": {
"num_frames": video.shape[0],
"height": video.shape[2],
"width": video.shape[3],
},
}
def _load_dataset_from_local_path(self) -> Tuple[List[str], List[str]]:
if not self.data_root.exists():
@@ -157,7 +175,8 @@ class ImageOrVideoDataset(Dataset):
prompt_path = self.data_root.joinpath(self.caption_column)
video_path = self.data_root.joinpath(self.video_column)
pose_path = self.data_root.joinpath(self.pose_column)
if not prompt_path.exists() or not prompt_path.is_file():
raise ValueError(
"Expected `--caption_column` to be path to a file in `--data_root` containing line-separated text prompts."
@@ -166,18 +185,23 @@ class ImageOrVideoDataset(Dataset):
raise ValueError(
"Expected `--video_column` to be path to a file in `--data_root` containing line-separated paths to video data in the same directory."
)
if not pose_path.exists() or not pose_path.is_file():
raise ValueError(
"Expected `--pose_column` to be path to a file in `--data_root` containing line-separated paths to video data in the same directory."
)
with open(prompt_path, "r", encoding="utf-8") as file:
prompts = [line.strip() for line in file.readlines() if len(line.strip()) > 0]
with open(video_path, "r", encoding="utf-8") as file:
video_paths = [self.data_root.joinpath(line.strip()) for line in file.readlines() if len(line.strip()) > 0]
with open(pose_path, "r", encoding="utf-8") as file:
pose_paths = [self.data_root.joinpath(line.strip()) for line in file.readlines() if len(line.strip()) > 0]
if any(not path.is_file() for path in video_paths):
raise ValueError(
f"Expected `{self.video_column=}` to be a path to a file in `{self.data_root=}` containing line-separated paths to video data but found atleast one path that is not a valid file."
)
return prompts, video_paths
return prompts, video_paths, pose_paths
def _load_dataset_from_csv(self) -> Tuple[List[str], List[str]]:
df = pd.read_csv(self.dataset_file)
@@ -243,7 +267,28 @@ class ImageOrVideoDataset(Dataset):
frames = frames.permute(0, 3, 1, 2).contiguous()
frames = torch.stack([self.video_transforms(frame) for frame in frames], dim=0)
return frames
def _preprocess_video_single_frame(self, path: Path) -> torch.Tensor:
video_reader = decord.VideoReader(uri=path.as_posix())
video_num_frames = len(video_reader)
nearest_frame_bucket = min(
[bucket for bucket in self.resolution_buckets if bucket[0] <= video_num_frames],
key=lambda x: abs(x[0] - min(video_num_frames, self.max_num_frames)),
default=1,
)[0]
frame_indices = [0 for _ in range(video_num_frames)]
frames = video_reader.get_batch(frame_indices)
frames = frames[:nearest_frame_bucket].float()
frames = frames.permute(0, 3, 1, 2).contiguous()
nearest_res = self._find_nearest_resolution(frames.shape[2], frames.shape[3])
frames_resized = torch.stack([resize(frame, nearest_res) for frame in frames], dim=0)
frames = torch.stack([self.video_transforms(frame) for frame in frames_resized], dim=0)
return frames
class ImageOrVideoDatasetWithResizing(ImageOrVideoDataset):
def __init__(self, *args, **kwargs) -> None:
+1
View File
@@ -191,6 +191,7 @@ def collate_fn_t2v(batch: List[List[Dict[str, torch.Tensor]]]) -> Dict[str, torc
return {
"prompts": [x["prompt"] for x in batch[0]],
"videos": torch.stack([x["video"] for x in batch[0]]),
"poses": torch.stack([x["poses"] for x in batch[0]])
}
+92 -21
View File
@@ -100,19 +100,34 @@ class Trainer:
self.state.model_name = self.args.model_name
self.model_config = get_config_from_model_name(self.args.model_name, self.args.training_type)
if self.args.pose_column != None:
self.pose_condition = True
else:
self.pose_condition = False
def prepare_dataset(self) -> None:
# TODO(aryan): Make a background process for fetching
logger.info("Initializing dataset and dataloader")
self.dataset = ImageOrVideoDatasetWithResizing(
data_root=self.args.data_root,
caption_column=self.args.caption_column,
video_column=self.args.video_column,
resolution_buckets=self.args.video_resolution_buckets,
dataset_file=self.args.dataset_file,
id_token=self.args.id_token,
remove_llm_prefixes=self.args.remove_common_llm_caption_prefixes,
)
if self.pose_condition == True:
self.dataset = ImageOrVideoDatasetWithResizing(
data_root=self.args.data_root,
caption_column=self.args.caption_column,
video_column=self.args.video_column,
pose_column=self.args.pose_column,
resolution_buckets=self.args.video_resolution_buckets,
dataset_file=self.args.dataset_file,
id_token=self.args.id_token,
)
else:
self.dataset = ImageOrVideoDatasetWithResizing(
data_root=self.args.data_root,
caption_column=self.args.caption_column,
video_column=self.args.video_column,
resolution_buckets=self.args.video_resolution_buckets,
dataset_file=self.args.dataset_file,
id_token=self.args.id_token,
remove_llm_prefixes=self.args.remove_common_llm_caption_prefixes,
)
self.dataloader = torch.utils.data.DataLoader(
self.dataset,
batch_size=1,
@@ -647,6 +662,12 @@ class Trainer:
if not self.args.precompute_conditions:
videos = batch["videos"]
prompts = batch["prompts"]
if self.pose_condition == True:
poses = batch["poses"]
# first frame video ? cond video.
first_frame_videos = batch["first_frame_videos"]
batch_size = len(prompts)
if self.args.caption_dropout_technique == "empty":
@@ -691,6 +712,32 @@ class Trainer:
batch_size = latent_conditions["latents"].shape[0]
latent_conditions = make_contiguous(latent_conditions)
if self.pose_conditioning:
pose_video_latents = self.model_config["prepare_latents"](
vae=self.vae,
image_or_video=poses,
patch_size=self.transformer_config.patch_size,
patch_size_t=self.transformer_config.patch_size_t,
device=accelerator.device,
dtype=weight_dtype,
generator=generator,
)
pose_video_latents = make_contiguous(pose_video_latents)
single_frame_video_latents = self.model_config["prepare_latents"](
vae=self.vae,
image_or_video=first_frame_videos,
patch_size=self.transformer_config.patch_size,
patch_size_t=self.transformer_config.patch_size_t,
device=accelerator.device,
dtype=weight_dtype,
generator=generator,
)
single_frame_video_latents = make_contiguous(single_frame_video_latents)
text_conditions = make_contiguous(text_conditions)
if self.args.caption_dropout_technique == "zero":
@@ -734,10 +781,20 @@ class Trainer:
timesteps=timesteps,
)
else:
# Default to flow-matching noise addition
noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise
noisy_latents = noisy_latents.to(latent_conditions["latents"].dtype)
if self.pose_conditioning:
# Use the training single frame - sample compare it to target
# todo fix it so the first frame isn't noised?
noisy_latents = (1.0 - sigmas) * single_frame_video_latents["latents"] + sigmas * noise
# add to noisy latent with the single frame repeating the conditioning to the pose latent
noisy_latents = noisy_latents + pose_video_latents["latents"]
single_frame_video_latents.update({"noisy_latents": noisy_latents})
else:
# Default to flow-matching noise addition
noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise
noisy_latents = noisy_latents.to(latent_conditions["latents"].dtype)
# this is used for whatever reason to pass into the
latent_conditions.update({"noisy_latents": noisy_latents})
@@ -749,13 +806,27 @@ class Trainer:
)
weights = expand_tensor_dims(weights, noise.ndim)
pred = self.model_config["forward_pass"](
transformer=self.transformer,
scheduler=self.scheduler,
timesteps=timesteps,
**latent_conditions,
**text_conditions,
)
if self.pose_conditioning:
pred = self.model_config["forward_pass"](
transformer=self.transformer,
scheduler=self.scheduler,
timesteps=timesteps,
**single_frame_video_latents,
**text_conditions,
)
else:
pred = self.model_config["forward_pass"](
transformer=self.transformer,
scheduler=self.scheduler,
timesteps=timesteps,
**latent_conditions,
**text_conditions,
)
target = prepare_target(
scheduler=self.scheduler, noise=noise, latents=latent_conditions["latents"]
)