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