From b10557c31ec6113e7eb6eb4e74abd6916d82c0b5 Mon Sep 17 00:00:00 2001 From: bubbliiiing <3323290568@qq.com> Date: Wed, 29 May 2024 15:02:07 +0800 Subject: [PATCH] fix bug in validation while training --- README.md | 6 +++--- easyanimate/models/autoencoder_magvit.py | 5 +---- predict_t2v.py | 5 ++--- scripts/train_t2iv_lora.py | 2 +- 4 files changed, 7 insertions(+), 11 deletions(-) diff --git a/README.md b/README.md index e9fd17a..6cdad27 100644 --- a/README.md +++ b/README.md @@ -189,7 +189,7 @@ EasyAnimateV2: #### b. Using webui - Step 1: Download the corresponding weights and place them in the models folder. -- Step 2: Run the app. py file to enter the graph page. +- Step 2: Run the app.py file to enter the graph page. - Step 3: Select the generated model based on the page, fill in prompt, neg_prompt, guidance_scale, and seed, click on generate, wait for the generated result, and save the result in the samples folder. ### 2. Model Training @@ -250,9 +250,9 @@ You can also set the path as absolute path as follow: #### b. Video VAE training (optional) Video VAE training is an optional option as we have already provided pre trained Video VAEs. -If you want to train video vae, you can refer to [README] (easyanimate/vae/README. md) in the video vae section. +If you want to train video vae, you can refer to [README](easyanimate/vae/README.md) in the video vae section. -#### c. Video VAE training +#### c. Video DiT training If the data format is relative path during data preprocessing, please set ```scripts/train_t2iv.sh``` as follow. ``` export DATASET_NAME="datasets/internal_datasets/" diff --git a/easyanimate/models/autoencoder_magvit.py b/easyanimate/models/autoencoder_magvit.py index ea24fda..d693767 100644 --- a/easyanimate/models/autoencoder_magvit.py +++ b/easyanimate/models/autoencoder_magvit.py @@ -267,10 +267,7 @@ class AutoencoderKLMagvit(ModelMixin, ConfigMixin, FromOriginalVAEMixin): def _decode(self, z: torch.FloatTensor, return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]: if self.use_tiling and (z.shape[-1] > self.tile_latent_min_size or z.shape[-2] > self.tile_latent_min_size): return self.tiled_decode(z, return_dict=return_dict) - self.post_quant_conv = self.post_quant_conv.float() - self.decoder = self.decoder.float() - - z = self.post_quant_conv(z.float()) + z = self.post_quant_conv(z) dec = self.decoder(z) if not return_dict: diff --git a/predict_t2v.py b/predict_t2v.py index 4c979c2..7841029 100644 --- a/predict_t2v.py +++ b/predict_t2v.py @@ -83,9 +83,8 @@ else: Choosen_AutoencoderKL = AutoencoderKL vae = Choosen_AutoencoderKL.from_pretrained( model_name, - subfolder="vae", - torch_dtype=weight_dtype -) + subfolder="vae" +).to(weight_dtype) if vae_path is not None: print(f"From checkpoint: {vae_path}") diff --git a/scripts/train_t2iv_lora.py b/scripts/train_t2iv_lora.py index abbed56..b509b0a 100644 --- a/scripts/train_t2iv_lora.py +++ b/scripts/train_t2iv_lora.py @@ -133,7 +133,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config, for i in range(len(args.validation_prompts)): with torch.no_grad(): - with torch.autocast("cuda"): + with torch.autocast("cuda", dtype=weight_dtype): sample = pipeline( args.validation_prompts[i], video_length = args.video_sample_n_frames,