mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Recommit the conditioning code.
This commit is contained in:
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
@@ -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"]
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user