This commit is contained in:
sayakpaul
2024-12-27 15:02:21 +05:30
parent 4461cdd98f
commit 84c17561eb
2 changed files with 8 additions and 1 deletions
+2 -1
View File
@@ -47,7 +47,7 @@ from .utils.data_utils import should_perform_precomputation
from .utils.file_utils import string_to_filename from .utils.file_utils import string_to_filename
from .utils.optimizer_utils import get_optimizer from .utils.optimizer_utils import get_optimizer
from .utils.memory_utils import get_memory_statistics, free_memory, make_contiguous 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 from .utils.checkpointing import get_latest_ckpt_path_to_resume_from, get_intermediate_ckpt_path
@@ -696,6 +696,7 @@ class Trainer:
device=accelerator.device, device=accelerator.device,
dtype=weight_dtype, dtype=weight_dtype,
) )
sigmas = expand_tensor_to_dims(sigmas, ndim=latent_conditions["latents"].ndim)
noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise
latent_conditions.update({"noisy_latents": noisy_latents}) latent_conditions.update({"noisy_latents": noisy_latents})
+6
View File
@@ -27,3 +27,9 @@ def align_device_and_dtype(
if dtype is not None: if dtype is not None:
x = {k: align_device_and_dtype(v, device, dtype) for k, v in x.items()} x = {k: align_device_and_dtype(v, device, dtype) for k, v in x.items()}
return x return x
def expand_tensor_to_dims(tensor, ndim):
while len(tensor.shape) < ndim:
tensor = tensor.unsqueeze(-1)
return tensor