More device management

This commit is contained in:
kijai
2024-03-04 16:28:52 +02:00
parent 9296177d09
commit 72c3c7eba3
2 changed files with 9 additions and 3 deletions
+4 -2
View File
@@ -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):
+5 -1
View File
@@ -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.