This commit is contained in:
Aryan
2024-12-21 15:18:46 +01:00
parent 2da8cb5798
commit 0db7bf28aa
3 changed files with 159 additions and 30 deletions
+105
View File
@@ -376,6 +376,111 @@ output = pipe(
export_to_video(output, "output.mp4", fps=15)
```
### Memory Usage
LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x512x768` resolutions, **without precomputation**:
```
Memory before training start: {
"memory_allocated": 38.889,
"memory_reserved": 39.02,
"max_memory_allocated": 38.889,
"max_memory_reserved": 39.02
}
Training configuration: {
"trainable parameters": 163577856,
"total samples": 69,
"train epochs": 1,
"train steps": 10,
"batches per device": 1,
"total batches observed per epoch": 69,
"train batch size": 1,
"gradient accumulation steps": 1
}
Memory before validation start: {
"memory_allocated": 39.747,
"memory_reserved": 56.266,
"max_memory_allocated": 51.867,
"max_memory_reserved": 56.266
}
Memory after validation end: {
"memory_allocated": 39.748,
"memory_reserved": 41.445,
"max_memory_allocated": 51.867,
"max_memory_reserved": 58.385
}
Memory after epoch 1: {
"memory_allocated": 39.748,
"memory_reserved": 40.91,
"max_memory_allocated": 39.748,
"max_memory_reserved": 40.91
}
Memory after training end: {
"memory_allocated": 25.288,
"memory_reserved": 27.783,
"max_memory_allocated": 39.748,
"max_memory_reserved": 40.91
}
```
LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x512x768` resolutions, **with precomputation**:
```
Memory after precomputing conditions: {
"memory_allocated": 14.232,
"memory_reserved": 14.336,
"max_memory_allocated": 14.395,
"max_memory_reserved": 14.461
}
Memory after precomputing latents: {
"memory_allocated": 14.717,
"memory_reserved": 14.762,
"max_memory_allocated": 16.759,
"max_memory_reserved": 17.244
}
Memory before training start: {
"memory_allocated": 24.195,
"memory_reserved": 26.039,
"max_memory_allocated": 24.195,
"max_memory_reserved": 26.039
}
Training configuration: {
"trainable parameters": 163577856,
"total samples": 1,
"train epochs": 10,
"train steps": 10,
"batches per device": 1,
"total batches observed per epoch": 1,
"train batch size": 1,
"gradient accumulation steps": 1
}
Memory after epoch 1: {
"memory_allocated": 24.83,
"memory_reserved": 42.387,
"max_memory_allocated": 36.357,
"max_memory_reserved": 42.387
}
Memory before validation start: {
"memory_allocated": 24.842,
"memory_reserved": 42.387,
"max_memory_allocated": 36.977,
"max_memory_reserved": 42.387
}
Memory after validation end: {
"memory_allocated": 39.558,
"memory_reserved": 41.039,
"max_memory_allocated": 43.226,
"max_memory_reserved": 46.947
}
Memory after training end: {
"memory_allocated": 24.842,
"memory_reserved": 26.82,
"max_memory_allocated": 39.558,
"max_memory_reserved": 41.039
}
```
</details>
If you would like to use a custom dataset, refer to the dataset preparation guide [here](./assets/dataset.md).
@@ -16,41 +16,44 @@ from PIL import Image
logger = get_logger("finetrainers") # pylint: disable=invalid-name
def load_components(
def load_condition_models(
model_id: str = "tencent/HunyuanVideo",
text_encoder_dtype: torch.dtype = torch.float16,
text_encoder_2_dtype: torch.dtype = torch.float16,
transformer_dtype: torch.dtype = torch.bfloat16,
revision: Optional[str] = None,
cache_dir: Optional[str] = None,
**kwargs,
) -> Dict[str, nn.Module]:
tokenizer = AutoTokenizer.from_pretrained(model_id, subfolder="tokenizer", revision=revision, cache_dir=cache_dir)
text_encoder = LlamaModel.from_pretrained(model_id, subfolder="text_encoder", torch_dtype=text_encoder_dtype, revision=revision, cache_dir=cache_dir)
tokenizer_2 = CLIPTokenizer.from_pretrained(model_id, subfolder="tokenizer_2", revision=revision, cache_dir=cache_dir)
text_encoder_2 = CLIPTextModel.from_pretrained(model_id, subfolder="text_encoder_2", torch_dtype=text_encoder_2_dtype, revision=revision, cache_dir=cache_dir)
return {"tokenizer": tokenizer, "text_encoder": text_encoder, "tokenizer_2": tokenizer_2, "text_encoder_2": text_encoder_2}
def load_latent_models(
model_id: str = "tencent/HunyuanVideo",
vae_dtype: torch.dtype = torch.float16,
revision: Optional[str] = None,
cache_dir: Optional[str] = None,
**kwargs,
) -> Dict[str, nn.Module]:
vae = AutoencoderKLHunyuanVideo.from_pretrained(model_id, subfolder="vae", torch_dtype=vae_dtype, revision=revision, cache_dir=cache_dir)
return {"vae": vae}
def load_diffusion_models(
model_id: str = "tencent/HunyuanVideo",
transformer_dtype: torch.dtype = torch.bfloat16,
revision: Optional[str] = None,
cache_dir: Optional[str] = None,
**kwargs,
) -> Dict[str, nn.Module]:
tokenizer = AutoTokenizer.from_pretrained(model_id, subfolder="tokenizer", revision=revision, cache_dir=cache_dir)
text_encoder = LlamaModel.from_pretrained(
model_id, subfolder="text_encoder", torch_dtype=text_encoder_dtype, revision=revision, cache_dir=cache_dir
)
tokenizer_2 = CLIPTokenizer.from_pretrained(
model_id, subfolder="tokenizer_2", revision=revision, cache_dir=cache_dir
)
text_encoder_2 = CLIPTextModel.from_pretrained(
model_id, subfolder="text_encoder_2", torch_dtype=text_encoder_2_dtype, revision=revision, cache_dir=cache_dir
)
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer", torch_dtype=transformer_dtype, revision=revision, cache_dir=cache_dir
)
vae = AutoencoderKLHunyuanVideo.from_pretrained(
model_id, subfolder="vae", torch_dtype=vae_dtype, revision=revision, cache_dir=cache_dir
)
scheduler = FlowMatchEulerDiscreteScheduler()
return {
"tokenizer": tokenizer,
"text_encoder": text_encoder,
"tokenizer_2": tokenizer_2,
"text_encoder_2": text_encoder_2,
"transformer": transformer,
"vae": vae,
"scheduler": scheduler,
}
return {"transformer": transformer, "scheduler": scheduler}
def initialize_pipeline(
@@ -72,6 +75,7 @@ def initialize_pipeline(
enable_slicing: bool = False,
enable_tiling: bool = False,
enable_model_cpu_offload: bool = False,
**kwargs,
) -> HunyuanVideoPipeline:
component_name_pairs = [
("tokenizer", tokenizer),
@@ -115,7 +119,7 @@ def prepare_conditions(
guidance: float = 1.0,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
max_sequence_length: int = 128,
max_sequence_length: int = 256,
# TODO(aryan): make configurable
prompt_template: Dict[str, Any] = {
"template": (
@@ -129,6 +133,7 @@ def prepare_conditions(
),
"crop_start": 95,
},
**kwargs,
) -> torch.Tensor:
device = device or text_encoder.device
dtype = dtype or text_encoder.dtype
@@ -154,6 +159,7 @@ def prepare_latents(
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
generator: Optional[torch.Generator] = None,
precompute: bool = False,
**kwargs,
) -> torch.Tensor:
device = device or vae.device
@@ -165,9 +171,24 @@ def prepare_latents(
image_or_video = image_or_video.to(device=device, dtype=vae.dtype)
image_or_video = image_or_video.permute(0, 2, 1, 3, 4).contiguous() # [B, C, F, H, W] -> [B, F, C, H, W]
latents = vae.encode(image_or_video).latent_dist.sample(generator=generator)
latents = latents * vae.config.scaling_factor
latents = latents.to(dtype=dtype)
if not precompute:
latents = vae.encode(image_or_video).latent_dist.sample(generator=generator)
latents = latents * vae.config.scaling_factor
latents = latents.to(dtype=dtype)
return {"latents": latents}
else:
if vae.use_slicing and image_or_video.shape[0] > 1:
encoded_slices = [vae._encode(x_slice) for x_slice in image_or_video.split(1)]
h = torch.cat(encoded_slices)
else:
h = vae._encode(image_or_video)
return {"latents": h}
def post_latent_preparation(
latents: torch.Tensor,
**kwargs,
) -> torch.Tensor:
return {"latents": latents}
@@ -309,10 +330,13 @@ def _get_clip_prompt_embeds(
HUNYUAN_VIDEO_T2V_LORA_CONFIG = {
"pipeline_cls": HunyuanVideoPipeline,
"load_components": load_components,
"load_condition_models": load_condition_models,
"load_latent_models": load_latent_models,
"load_diffusion_models": load_diffusion_models,
"initialize_pipeline": initialize_pipeline,
"prepare_conditions": prepare_conditions,
"prepare_latents": prepare_latents,
"post_latent_preparation": post_latent_preparation,
"collate_fn": collate_fn_t2v,
"forward_pass": forward_pass,
"validation": validation,
+1 -1
View File
@@ -319,7 +319,7 @@ class Trainer:
generator=self.state.generator,
precompute=True,
)
filename = latents_dir / f"latents-{i}-{index}.pt"
filename = latents_dir / f"latents-{self.state.accelerator.process_index}-{index}.pt"
torch.save(latent_conditions, filename.as_posix())
index += 1
progress_bar.update(1)