diff --git a/sgm/modules/diffusionmodules/sampling.py b/sgm/modules/diffusionmodules/sampling.py index e914fd8..ff1be88 100644 --- a/sgm/modules/diffusionmodules/sampling.py +++ b/sgm/modules/diffusionmodules/sampling.py @@ -20,6 +20,8 @@ from ...util import append_dims, default, instantiate_from_config DEFAULT_GUIDER = {"target": ".sgm.modules.diffusionmodules.guiders.IdentityGuider"} +import comfy.model_management +device = comfy.model_management.get_torch_device() class BaseDiffusionSampler: def __init__( @@ -39,7 +41,7 @@ class BaseDiffusionSampler: ) ) self.verbose = verbose - self.device = device + self.device = comfy.model_management.get_torch_device() def prepare_sampling_loop(self, x, cond, uc=None, num_steps=None): sigmas = self.discretization( @@ -531,7 +533,7 @@ def gaussian_weights(tile_width, tile_height, nbatches): for y in range(latent_height)] weights = np.outer(y_probs, x_probs) - return torch.tile(torch.tensor(weights, device='cuda'), (nbatches, 4, 1, 1)) + return torch.tile(torch.tensor(weights, device=device), (nbatches, 4, 1, 1)) def _sliding_windows(h: int, w: int, tile_size: int, tile_stride: int): diff --git a/sgm/modules/diffusionmodules/util.py b/sgm/modules/diffusionmodules/util.py index 65f0552..65ccd0a 100644 --- a/sgm/modules/diffusionmodules/util.py +++ b/sgm/modules/diffusionmodules/util.py @@ -15,6 +15,9 @@ import torch import torch.nn as nn from einops import repeat +import comfy.model_management +device = comfy.model_management.get_torch_device() +from contextlib import nullcontext def make_beta_schedule( schedule, @@ -185,7 +188,8 @@ class CheckpointFunction(torch.autograd.Function): @staticmethod def backward(ctx, *output_grads): ctx.input_tensors = [x.detach().requires_grad_(True) for x in ctx.input_tensors] - with torch.enable_grad(), torch.cuda.amp.autocast(**ctx.gpu_autocast_kwargs): + autocast_condition = (ctx.input_tensors.dtype == torch.float16 or ctx.input_tensors.dtype == torch.bfloat16) and not comfy.model_management.is_device_mps(device) + with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=ctx.input_tensors.dtype) if autocast_condition else nullcontext(): # Fixes a bug where the first op in run_function modifies the # Tensor storage in place, which is not allowed for detach()'d # Tensors.