fix bug in validation while training

This commit is contained in:
bubbliiiing
2024-05-29 15:02:07 +08:00
parent 5e1041a104
commit b10557c31e
4 changed files with 7 additions and 11 deletions
+3 -3
View File
@@ -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/"
+1 -4
View File
@@ -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:
+2 -3
View File
@@ -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}")
+1 -1
View File
@@ -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,