From e50844e79f42ec241d6b5b638979c4c4e9bac4e7 Mon Sep 17 00:00:00 2001 From: aiXander Date: Sat, 30 Mar 2024 00:14:22 -0700 Subject: [PATCH] simplify dataloader --- trainer/dataset_and_utils.py | 24 +++++------------------- 1 file changed, 5 insertions(+), 19 deletions(-) diff --git a/trainer/dataset_and_utils.py b/trainer/dataset_and_utils.py index 8319215..9a5d2b9 100755 --- a/trainer/dataset_and_utils.py +++ b/trainer/dataset_and_utils.py @@ -196,9 +196,7 @@ class PreprocessedDataset(Dataset): ) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]: image_path = self.image_path[idx] image_path = os.path.join(os.path.dirname(self.csv_path), image_path) - image = PIL.Image.open(image_path).convert("RGB") - image = prepare_image(image, self.size, self.size).to( dtype=self.vae_encoder.dtype, device=self.vae_encoder.device ) @@ -210,7 +208,7 @@ class PreprocessedDataset(Dataset): ti1 = self.tokenizer_1( caption, padding="max_length", - max_length=77, + max_length=self.tokenizer_1.model_max_length, truncation=True, add_special_tokens=True, return_tensors="pt", @@ -222,17 +220,13 @@ class PreprocessedDataset(Dataset): ti2 = self.tokenizer_2( caption, padding="max_length", - max_length=77, + max_length=self.tokenizer_2.model_max_length, truncation=True, add_special_tokens=True, return_tensors="pt", ).input_ids.squeeze() vae_latent = self.vae_encoder.encode(image).latent_dist#.sample() - - #if self.scale_vae_latents: - # vae_latent = vae_latent * self.vae_encoder.config.scaling_factor - dummy_vae_latent = vae_latent.sample() if self.mask_path is None: @@ -243,9 +237,7 @@ class PreprocessedDataset(Dataset): else: mask_path = self.mask_path[idx] mask_path = os.path.join(os.path.dirname(self.csv_path), mask_path) - mask = PIL.Image.open(mask_path) - mask = prepare_mask(mask, self.size, self.size).to( dtype=self.vae_encoder.dtype, device=self.vae_encoder.device ) @@ -265,26 +257,20 @@ class PreprocessedDataset(Dataset): else: # sdxl return (ti1, ti2), vae_latent, mask.squeeze() - def atidx( + def __getitem__( self, idx: int ) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]: if self.do_cache: vae_latent = self.vae_latents[idx].sample() if self.scale_vae_latents: vae_latent *= self.vae_scaling_factor - return self.tokens_tuple[idx], vae_latent, self.masks[idx] + return self.tokens_tuple[idx], vae_latent.squeeze(), self.masks[idx] else: tokens, vae_latent, mask = self._process(idx) vae_latent = vae_latent.sample() if self.scale_vae_latents: vae_latent *= self.vae_scaling_factor - return tokens, vae_latent, mask - - def __getitem__( - self, idx: int - ) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]: - token, vae_latent, mask = self.atidx(idx) - return token, vae_latent.squeeze(), mask + return tokens, vae_latent.squeeze(), mask def import_model_class_from_model_name_or_path(