From 84c17561eb813d361eea9edfebd61e07f6c19557 Mon Sep 17 00:00:00 2001 From: sayakpaul Date: Fri, 27 Dec 2024 15:02:21 +0530 Subject: [PATCH] fixes --- finetrainers/trainer.py | 3 ++- finetrainers/utils/torch_utils.py | 6 ++++++ 2 files changed, 8 insertions(+), 1 deletion(-) diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index db98e5d..bbacc49 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -47,7 +47,7 @@ from .utils.data_utils import should_perform_precomputation from .utils.file_utils import string_to_filename from .utils.optimizer_utils import get_optimizer from .utils.memory_utils import get_memory_statistics, free_memory, make_contiguous -from .utils.torch_utils import unwrap_model, align_device_and_dtype +from .utils.torch_utils import unwrap_model, align_device_and_dtype, expand_tensor_to_dims from .utils.checkpointing import get_latest_ckpt_path_to_resume_from, get_intermediate_ckpt_path @@ -696,6 +696,7 @@ class Trainer: device=accelerator.device, dtype=weight_dtype, ) + sigmas = expand_tensor_to_dims(sigmas, ndim=latent_conditions["latents"].ndim) noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise latent_conditions.update({"noisy_latents": noisy_latents}) diff --git a/finetrainers/utils/torch_utils.py b/finetrainers/utils/torch_utils.py index 2aad077..17c5f9c 100644 --- a/finetrainers/utils/torch_utils.py +++ b/finetrainers/utils/torch_utils.py @@ -27,3 +27,9 @@ def align_device_and_dtype( if dtype is not None: x = {k: align_device_and_dtype(v, device, dtype) for k, v in x.items()} return x + + +def expand_tensor_to_dims(tensor, ndim): + while len(tensor.shape) < ndim: + tensor = tensor.unsqueeze(-1) + return tensor \ No newline at end of file