Precomputation folder name based on model name (#196)

* update

* update

* fix

* address review comments
This commit is contained in:
Aryan
2025-01-08 08:45:44 +05:30
committed by GitHub
parent 905c22edd8
commit dbffc80e7c
2 changed files with 13 additions and 9 deletions
+7 -3
View File
@@ -217,7 +217,11 @@ class Trainer:
batched_text_conditions[key] = [x[key] for x in text_conditions][0]
return {"latent_conditions": batched_latent_conditions, "text_conditions": batched_text_conditions}
should_precompute = should_perform_precomputation(self.args.data_root)
cleaned_model_id = string_to_filename(self.args.pretrained_model_name_or_path)
precomputation_dir = (
Path(self.args.data_root) / f"{self.args.model_name}_{cleaned_model_id}_{PRECOMPUTED_DIR_NAME}"
)
should_precompute = should_perform_precomputation(precomputation_dir)
if not should_precompute:
logger.info("Precomputed conditions and latents found. Loading precomputed data.")
self.dataloader = torch.utils.data.DataLoader(
@@ -252,8 +256,8 @@ class Trainer:
"Caption dropout is not supported with precomputation yet. This will be supported in the future."
)
conditions_dir = Path(self.args.data_root) / PRECOMPUTED_DIR_NAME / PRECOMPUTED_CONDITIONS_DIR_NAME
latents_dir = Path(self.args.data_root) / PRECOMPUTED_DIR_NAME / PRECOMPUTED_LATENTS_DIR_NAME
conditions_dir = precomputation_dir / PRECOMPUTED_CONDITIONS_DIR_NAME
latents_dir = precomputation_dir / PRECOMPUTED_LATENTS_DIR_NAME
conditions_dir.mkdir(parents=True, exist_ok=True)
latents_dir.mkdir(parents=True, exist_ok=True)
+6 -6
View File
@@ -3,17 +3,17 @@ from typing import Union
from accelerate.logging import get_logger
from ..constants import PRECOMPUTED_CONDITIONS_DIR_NAME, PRECOMPUTED_DIR_NAME, PRECOMPUTED_LATENTS_DIR_NAME
from ..constants import PRECOMPUTED_CONDITIONS_DIR_NAME, PRECOMPUTED_LATENTS_DIR_NAME
logger = get_logger("finetrainers")
def should_perform_precomputation(data_root: Union[str, Path]) -> bool:
if isinstance(data_root, str):
data_root = Path(data_root)
conditions_dir = data_root / PRECOMPUTED_DIR_NAME / PRECOMPUTED_CONDITIONS_DIR_NAME
latents_dir = data_root / PRECOMPUTED_DIR_NAME / PRECOMPUTED_LATENTS_DIR_NAME
def should_perform_precomputation(precomputation_dir: Union[str, Path]) -> bool:
if isinstance(precomputation_dir, str):
precomputation_dir = Path(precomputation_dir)
conditions_dir = precomputation_dir / PRECOMPUTED_CONDITIONS_DIR_NAME
latents_dir = precomputation_dir / PRECOMPUTED_LATENTS_DIR_NAME
if conditions_dir.exists() and latents_dir.exists():
num_files_conditions = len(list(conditions_dir.glob("*.pt")))
num_files_latents = len(list(latents_dir.glob("*.pt")))