From 6716ccf8c2d8b9cd0d3bca0bb472ee0045605520 Mon Sep 17 00:00:00 2001 From: Aryan Date: Thu, 9 Jan 2025 11:06:34 +0530 Subject: [PATCH 1/4] update (#198) Co-authored-by: Sayak Paul --- README.md | 8 +++++--- docs/training/ltx_video.md | 6 +++--- 2 files changed, 8 insertions(+), 6 deletions(-) diff --git a/README.md b/README.md index 141cbf2..dc615ec 100644 --- a/README.md +++ b/README.md @@ -43,6 +43,8 @@ Then launch LoRA fine-tuning. Below we provide an example for LTX-Video. We refe
Training command +TODO: LTX does not do too well with the disney dataset. We will update this to use a better example soon. + ```bash #!/bin/bash export WANDB_MODE="offline" @@ -75,18 +77,18 @@ dataset_cmd="--data_root $DATA_ROOT \ dataloader_cmd="--dataloader_num_workers 0" # Diffusion arguments -diffusion_cmd="--flow_resolution_shifting" +diffusion_cmd="--flow_weighting_scheme logit_normal" # Training arguments training_cmd="--training_type lora \ --seed 42 \ --mixed_precision bf16 \ --batch_size 1 \ - --train_steps 1200 \ + --train_steps 3000 \ --rank 128 \ --lora_alpha 128 \ --target_modules to_q to_k to_v to_out.0 \ - --gradient_accumulation_steps 1 \ + --gradient_accumulation_steps 4 \ --gradient_checkpointing \ --checkpointing_steps 500 \ --checkpointing_limit 2 \ diff --git a/docs/training/ltx_video.md b/docs/training/ltx_video.md index 6d14acc..a25390c 100644 --- a/docs/training/ltx_video.md +++ b/docs/training/ltx_video.md @@ -36,18 +36,18 @@ dataset_cmd="--data_root $DATA_ROOT \ dataloader_cmd="--dataloader_num_workers 0" # Diffusion arguments -diffusion_cmd="--flow_resolution_shifting" +diffusion_cmd="--flow_weighting_scheme logit_normal" # Training arguments training_cmd="--training_type lora \ --seed 42 \ --mixed_precision bf16 \ --batch_size 1 \ - --train_steps 1200 \ + --train_steps 3000 \ --rank 128 \ --lora_alpha 128 \ --target_modules to_q to_k to_v to_out.0 \ - --gradient_accumulation_steps 1 \ + --gradient_accumulation_steps 4 \ --gradient_checkpointing \ --checkpointing_steps 500 \ --checkpointing_limit 2 \ From f311f16d988d22fecca5532974b652316fa066d6 Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Thu, 9 Jan 2025 12:43:03 +0530 Subject: [PATCH 2/4] [core] Fix loading of precomputed conditions and latents (#199) * support precomputation. * fixes --- finetrainers/dataset.py | 11 ++++++++--- finetrainers/trainer.py | 8 ++++++-- 2 files changed, 14 insertions(+), 5 deletions(-) diff --git a/finetrainers/dataset.py b/finetrainers/dataset.py index 6054e49..19ebb69 100644 --- a/finetrainers/dataset.py +++ b/finetrainers/dataset.py @@ -353,13 +353,18 @@ class ImageOrVideoDatasetWithResizeAndRectangleCrop(ImageOrVideoDataset): class PrecomputedDataset(Dataset): - def __init__(self, data_root: str) -> None: + def __init__(self, data_root: str, model_name: str = None, cleaned_model_id: str = None) -> None: super().__init__() self.data_root = Path(data_root) - self.latents_path = self.data_root / PRECOMPUTED_DIR_NAME / PRECOMPUTED_LATENTS_DIR_NAME - self.conditions_path = self.data_root / PRECOMPUTED_DIR_NAME / PRECOMPUTED_CONDITIONS_DIR_NAME + if model_name and cleaned_model_id: + precomputation_dir = self.data_root / f"{model_name}_{cleaned_model_id}_{PRECOMPUTED_DIR_NAME}" + self.latents_path = precomputation_dir / PRECOMPUTED_LATENTS_DIR_NAME + self.conditions_path = precomputation_dir / PRECOMPUTED_CONDITIONS_DIR_NAME + else: + self.latents_path = self.data_root / PRECOMPUTED_DIR_NAME / PRECOMPUTED_LATENTS_DIR_NAME + self.conditions_path = self.data_root / PRECOMPUTED_DIR_NAME / PRECOMPUTED_CONDITIONS_DIR_NAME self.latent_conditions = sorted(os.listdir(self.latents_path)) self.text_conditions = sorted(os.listdir(self.conditions_path)) diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index a05a8f0..81c8264 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -225,7 +225,9 @@ class Trainer: if not should_precompute: logger.info("Precomputed conditions and latents found. Loading precomputed data.") self.dataloader = torch.utils.data.DataLoader( - PrecomputedDataset(self.args.data_root), + PrecomputedDataset( + data_root=self.args.data_root, model_name=self.args.model_name, cleaned_model_id=cleaned_model_id + ), batch_size=self.args.batch_size, shuffle=True, collate_fn=collate_fn, @@ -353,7 +355,9 @@ class Trainer: # Update dataloader to use precomputed conditions and latents self.dataloader = torch.utils.data.DataLoader( - PrecomputedDataset(self.args.data_root), + PrecomputedDataset( + data_root=self.args.data_root, model_name=self.args.model_name, cleaned_model_id=cleaned_model_id + ), batch_size=self.args.batch_size, shuffle=True, collate_fn=collate_fn, From b9a6492ad1457fe672b042f5898fa0df5f4af7b0 Mon Sep 17 00:00:00 2001 From: Aryan Date: Thu, 9 Jan 2025 14:22:56 +0530 Subject: [PATCH 3/4] Epoch loss (#201) * update * update --- finetrainers/trainer.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index 81c8264..8f2ebed 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -673,6 +673,8 @@ class Trainer: self.transformer.train() models_to_accumulate = [self.transformer] + epoch_loss = 0.0 + num_loss_updates = 0 for step, batch in enumerate(self.dataloader): logger.debug(f"Starting step {step + 1}") @@ -843,7 +845,10 @@ class Trainer: if should_run_validation: self.validate(global_step) - logs["loss"] = loss.detach().item() + loss_item = loss.detach().item() + epoch_loss += loss_item + num_loss_updates += 1 + logs["step_loss"] = loss_item logs["lr"] = self.lr_scheduler.get_last_lr()[0] progress_bar.set_postfix(logs) accelerator.log(logs, step=global_step) @@ -851,6 +856,9 @@ class Trainer: if global_step >= self.state.train_steps: break + if num_loss_updates > 0: + epoch_loss /= num_loss_updates + accelerator.log({"epoch_loss": epoch_loss}, step=global_step) memory_statistics = get_memory_statistics() logger.info(f"Memory after epoch {epoch + 1}: {json.dumps(memory_statistics, indent=4)}") From 9523a05220cf1b92b1c6541deb4709813e8e4d0e Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Fri, 10 Jan 2025 07:22:00 +0530 Subject: [PATCH 4/4] Shell script to minimally test supported models on a real dataset (#204) * start minimal complete tests * updates * fixes * fixes --- tests/scripts/dummy_cogvideox_lora.sh | 81 +++++++++++++++++++++++ tests/scripts/dummy_hunyuanvideo_lora.sh | 79 ++++++++++++++++++++++ tests/scripts/dummy_ltx_video_lora.sh | 84 ++++++++++++++++++++++++ tests/test_dataset.py | 6 +- tests/test_model_runs_minimally_lora.sh | 48 ++++++++++++++ 5 files changed, 295 insertions(+), 3 deletions(-) create mode 100644 tests/scripts/dummy_cogvideox_lora.sh create mode 100644 tests/scripts/dummy_hunyuanvideo_lora.sh create mode 100644 tests/scripts/dummy_ltx_video_lora.sh create mode 100644 tests/test_model_runs_minimally_lora.sh diff --git a/tests/scripts/dummy_cogvideox_lora.sh b/tests/scripts/dummy_cogvideox_lora.sh new file mode 100644 index 0000000..c1f7bbd --- /dev/null +++ b/tests/scripts/dummy_cogvideox_lora.sh @@ -0,0 +1,81 @@ +#!/bin/bash + +GPU_IDS="0,1" +DATA_ROOT="$ROOT_DIR/video-dataset-disney" +CAPTION_COLUMN="prompt.txt" +VIDEO_COLUMN="videos.txt" +OUTPUT_DIR="cogvideox" +ID_TOKEN="BW_STYLE" + +# Model arguments +model_cmd="--model_name cogvideox \ + --pretrained_model_name_or_path THUDM/CogVideoX-5b" + +# Dataset arguments +dataset_cmd="--data_root $DATA_ROOT \ + --video_column $VIDEO_COLUMN \ + --caption_column $CAPTION_COLUMN \ + --id_token $ID_TOKEN \ + --video_resolution_buckets 49x480x720 \ + --caption_dropout_p 0.05" + +# Dataloader arguments +dataloader_cmd="--dataloader_num_workers 0 --precompute_conditions" + +# Training arguments +training_cmd="--training_type lora \ + --seed 42 \ + --mixed_precision bf16 \ + --batch_size 1 \ + --precompute_conditions \ + --train_steps 10 \ + --rank 128 \ + --lora_alpha 128 \ + --target_modules to_q to_k to_v to_out.0 \ + --gradient_accumulation_steps 1 \ + --gradient_checkpointing \ + --checkpointing_steps 5 \ + --checkpointing_limit 2 \ + --resume_from_checkpoint=latest \ + --enable_slicing \ + --enable_tiling" + +# Optimizer arguments +optimizer_cmd="--optimizer adamw \ + --lr 3e-5 \ + --beta1 0.9 \ + --beta2 0.95 \ + --weight_decay 1e-4 \ + --epsilon 1e-8 \ + --max_grad_norm 1.0" + +# Validation arguments +validation_prompts=$(cat <