More device management
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user