address some reviews; change hunyuan checkpoint repo-id

This commit is contained in:
Aryan
2024-12-23 12:54:52 +01:00
parent b045037850
commit 363841ae3f
3 changed files with 12 additions and 12 deletions
+2 -3
View File
@@ -273,8 +273,7 @@ OUTPUT_DIR="/path/to/models/hunyuan-video/hunyuan-video-loras/hunyuan-video_caki
# Model arguments
model_cmd="--model_name hunyuan_video \
--pretrained_model_name_or_path tencent/HunyuanVideo
--revision refs/pr/18"
--pretrained_model_name_or_path hunyuanvideo-community/HunyuanVideo"
# Dataset arguments
dataset_cmd="--data_root $DATA_ROOT \
@@ -356,7 +355,7 @@ import torch
from diffusers import HunyuanVideoPipeline, HunyuanVideoTransformer3DModel
from diffusers.utils import export_to_video
model_id = "tencent/HunyuanVideo"
model_id = "hunyuanvideo-community/HunyuanVideo"
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer", torch_dtype=torch.bfloat16
)
@@ -17,7 +17,7 @@ logger = get_logger("finetrainers") # pylint: disable=invalid-name
def load_condition_models(
model_id: str = "tencent/HunyuanVideo",
model_id: str = "hunyuanvideo-community/HunyuanVideo",
text_encoder_dtype: torch.dtype = torch.float16,
text_encoder_2_dtype: torch.dtype = torch.float16,
revision: Optional[str] = None,
@@ -43,7 +43,7 @@ def load_condition_models(
def load_latent_models(
model_id: str = "tencent/HunyuanVideo",
model_id: str = "hunyuanvideo-community/HunyuanVideo",
vae_dtype: torch.dtype = torch.float16,
revision: Optional[str] = None,
cache_dir: Optional[str] = None,
@@ -56,12 +56,12 @@ def load_latent_models(
def load_diffusion_models(
model_id: str = "tencent/HunyuanVideo",
model_id: str = "hunyuanvideo-community/HunyuanVideo",
transformer_dtype: torch.dtype = torch.bfloat16,
revision: Optional[str] = None,
cache_dir: Optional[str] = None,
**kwargs,
) -> Dict[str, nn.Module]:
) -> Dict[str, Union[nn.Module, FlowMatchEulerDiscreteScheduler]]:
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer", torch_dtype=transformer_dtype, revision=revision, cache_dir=cache_dir
)
@@ -70,7 +70,7 @@ def load_diffusion_models(
def initialize_pipeline(
model_id: str = "tencent/HunyuanVideo",
model_id: str = "hunyuanvideo-community/HunyuanVideo",
text_encoder_dtype: torch.dtype = torch.float16,
text_encoder_2_dtype: torch.dtype = torch.float16,
transformer_dtype: torch.dtype = torch.bfloat16,
+5 -4
View File
@@ -138,7 +138,7 @@ class Trainer:
self.vae = components.get("vae", self.vae)
self.scheduler = components.get("scheduler", self.scheduler)
def _nuke_components(self) -> None:
def _delete_components(self) -> None:
self.tokenizer = None
self.tokenizer_2 = None
self.tokenizer_3 = None
@@ -274,7 +274,7 @@ class Trainer:
torch.save(other_conditions, filename.as_posix())
index += 1
progress_bar.update(1)
self._nuke_components()
self._delete_components()
memory_statistics = get_memory_statistics()
logger.info(f"Memory after precomputing conditions: {json.dumps(memory_statistics, indent=4)}")
@@ -323,7 +323,7 @@ class Trainer:
torch.save(latent_conditions, filename.as_posix())
index += 1
progress_bar.update(1)
self._nuke_components()
self._delete_components()
self.state.accelerator.wait_for_everyone()
logger.info("Precomputation complete")
@@ -669,7 +669,8 @@ class Trainer:
accelerator.backward(loss)
if accelerator.sync_gradients and accelerator.distributed_type != DistributedType.DEEPSPEED:
accelerator.clip_grad_norm_(self.transformer.parameters(), self.args.max_grad_norm)
grad_norm = accelerator.clip_grad_norm_(self.transformer.parameters(), self.args.max_grad_norm)
logs["grad_norm"] = grad_norm
self.optimizer.step()
self.lr_scheduler.step()