mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
fixed the training pipeline to address new changes from upstream
This commit is contained in:
@@ -76,7 +76,7 @@ class ImageOrVideoDataset(Dataset):
|
||||
(
|
||||
self.prompts,
|
||||
self.video_paths,
|
||||
self.pose_pathes,
|
||||
self.pose_paths,
|
||||
) = self._load_dataset_from_local_path()
|
||||
elif dataset_file.endswith(".csv"):
|
||||
(
|
||||
@@ -146,12 +146,15 @@ class ImageOrVideoDataset(Dataset):
|
||||
else:
|
||||
video = self._preprocess_video(video_path)
|
||||
|
||||
if self.pose_condition_column != None:
|
||||
pose_video = self._preprocess_video(self.pose_paths[index])
|
||||
if self.pose_column != None:
|
||||
pose = self._preprocess_video(self.pose_paths[index])
|
||||
img_ref = self._preprocess_video_image_reference_video(video_path)
|
||||
|
||||
return {
|
||||
"prompt": prompt,
|
||||
"video": video,
|
||||
"pose_video": pose_video,
|
||||
"pose": pose,
|
||||
"img_ref": img_ref,
|
||||
"video_metadata": {
|
||||
"num_frames": video.shape[0],
|
||||
"height": video.shape[2],
|
||||
@@ -268,7 +271,7 @@ class ImageOrVideoDataset(Dataset):
|
||||
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:
|
||||
def _preprocess_video_image_reference_video(self, path: Path) -> torch.Tensor:
|
||||
video_reader = decord.VideoReader(uri=path.as_posix())
|
||||
|
||||
video_num_frames = len(video_reader)
|
||||
|
||||
@@ -139,15 +139,9 @@ def prepare_latents(
|
||||
if not precompute:
|
||||
latents = vae.encode(image_or_video).latent_dist.sample(generator=generator)
|
||||
|
||||
|
||||
|
||||
latents = latents.to(dtype=dtype)
|
||||
_, _, num_frames, height, width = latents.shape
|
||||
latents = _normalize_latents(latents, vae.latents_mean, vae.latents_std)
|
||||
|
||||
# expand the channel and pack
|
||||
latents = torch.cat([latents,latents],dim=1)
|
||||
|
||||
latents = _pack_latents(latents, patch_size, patch_size_t)
|
||||
return {"latents": latents, "num_frames": num_frames, "height": height, "width": width}
|
||||
else:
|
||||
@@ -191,7 +185,8 @@ 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]])
|
||||
"poses": torch.stack([x["pose"] for x in batch[0]]),
|
||||
"img_refs": torch.stack([x["img_ref"] for x in batch[0]])
|
||||
}
|
||||
|
||||
|
||||
|
||||
+10
-14
@@ -665,8 +665,8 @@ class Trainer:
|
||||
|
||||
if self.pose_condition == True:
|
||||
poses = batch["poses"]
|
||||
# first frame video ? cond video.
|
||||
first_frame_videos = batch["first_frame_videos"]
|
||||
# generate video across all frames
|
||||
img_refs = batch["img_refs"]
|
||||
|
||||
batch_size = len(prompts)
|
||||
|
||||
@@ -726,9 +726,9 @@ class Trainer:
|
||||
|
||||
pose_video_latents = make_contiguous(pose_video_latents)
|
||||
|
||||
single_frame_video_latents = self.model_config["prepare_latents"](
|
||||
img_refs_latents = self.model_config["prepare_latents"](
|
||||
vae=self.vae,
|
||||
image_or_video=first_frame_videos,
|
||||
image_or_video=img_refs,
|
||||
patch_size=self.transformer_config.patch_size,
|
||||
patch_size_t=self.transformer_config.patch_size_t,
|
||||
device=accelerator.device,
|
||||
@@ -736,7 +736,7 @@ class Trainer:
|
||||
generator=generator,
|
||||
)
|
||||
|
||||
single_frame_video_latents = make_contiguous(single_frame_video_latents)
|
||||
img_refs_latents = make_contiguous(img_refs_latents)
|
||||
|
||||
text_conditions = make_contiguous(text_conditions)
|
||||
|
||||
@@ -782,18 +782,15 @@ class Trainer:
|
||||
)
|
||||
else:
|
||||
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 = (1.0 - sigmas) * img_refs_latents["latents"] + sigmas * noise
|
||||
|
||||
noisy_latents = noisy_latents + pose_video_latents["latents"]
|
||||
single_frame_video_latents.update({"noisy_latents": noisy_latents})
|
||||
img_refs_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})
|
||||
@@ -811,11 +808,10 @@ class Trainer:
|
||||
transformer=self.transformer,
|
||||
scheduler=self.scheduler,
|
||||
timesteps=timesteps,
|
||||
**single_frame_video_latents,
|
||||
**img_refs_latents,
|
||||
**text_conditions,
|
||||
)
|
||||
|
||||
|
||||
else:
|
||||
pred = self.model_config["forward_pass"](
|
||||
transformer=self.transformer,
|
||||
|
||||
Reference in New Issue
Block a user