diff --git a/comfy/controlnet.py b/comfy/controlnet.py index 1d24afa6f..d2744e427 100644 --- a/comfy/controlnet.py +++ b/comfy/controlnet.py @@ -60,7 +60,7 @@ class StrengthType(Enum): LINEAR_UP = 2 class ControlBase: - def __init__(self, device=None): + def __init__(self): self.cond_hint_original = None self.cond_hint = None self.strength = 1.0 @@ -72,10 +72,6 @@ class ControlBase: self.compression_ratio = 8 self.upscale_algorithm = 'nearest-exact' self.extra_args = {} - - if device is None: - device = comfy.model_management.get_torch_device() - self.device = device self.previous_controlnet = None self.extra_conds = [] self.strength_type = StrengthType.CONSTANT @@ -185,8 +181,8 @@ class ControlBase: class ControlNet(ControlBase): - def __init__(self, control_model=None, global_average_pooling=False, compression_ratio=8, latent_format=None, device=None, load_device=None, manual_cast_dtype=None, extra_conds=["y"], strength_type=StrengthType.CONSTANT, concat_mask=False): - super().__init__(device) + def __init__(self, control_model=None, global_average_pooling=False, compression_ratio=8, latent_format=None, load_device=None, manual_cast_dtype=None, extra_conds=["y"], strength_type=StrengthType.CONSTANT, concat_mask=False): + super().__init__() self.control_model = control_model self.load_device = load_device if control_model is not None: @@ -242,7 +238,7 @@ class ControlNet(ControlBase): to_concat.append(comfy.utils.repeat_to_batch_size(c, self.cond_hint.shape[0])) self.cond_hint = torch.cat([self.cond_hint] + to_concat, dim=1) - self.cond_hint = self.cond_hint.to(device=self.device, dtype=dtype) + self.cond_hint = self.cond_hint.to(device=x_noisy.device, dtype=dtype) if x_noisy.shape[0] != self.cond_hint.shape[0]: self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number) @@ -341,8 +337,8 @@ class ControlLoraOps: class ControlLora(ControlNet): - def __init__(self, control_weights, global_average_pooling=False, device=None, model_options={}): #TODO? model_options - ControlBase.__init__(self, device) + def __init__(self, control_weights, global_average_pooling=False, model_options={}): #TODO? model_options + ControlBase.__init__(self) self.control_weights = control_weights self.global_average_pooling = global_average_pooling self.extra_conds += ["y"] @@ -662,12 +658,15 @@ def load_controlnet(ckpt_path, model=None, model_options={}): class T2IAdapter(ControlBase): def __init__(self, t2i_model, channels_in, compression_ratio, upscale_algorithm, device=None): - super().__init__(device) + super().__init__() self.t2i_model = t2i_model self.channels_in = channels_in self.control_input = None self.compression_ratio = compression_ratio self.upscale_algorithm = upscale_algorithm + if device is None: + device = comfy.model_management.get_torch_device() + self.device = device def scale_image_to(self, width, height): unshuffle_amount = self.t2i_model.unshuffle_amount diff --git a/comfy/ldm/modules/diffusionmodules/mmdit.py b/comfy/ldm/modules/diffusionmodules/mmdit.py index 759788a97..b085bbc0f 100644 --- a/comfy/ldm/modules/diffusionmodules/mmdit.py +++ b/comfy/ldm/modules/diffusionmodules/mmdit.py @@ -5,7 +5,7 @@ from typing import Dict, Optional import numpy as np import torch import torch.nn as nn -from .. import attention +from ..attention import optimized_attention from einops import rearrange, repeat from .util import timestep_embedding import comfy.ops @@ -266,8 +266,6 @@ def split_qkv(qkv, head_dim): qkv = qkv.reshape(qkv.shape[0], qkv.shape[1], 3, -1, head_dim).movedim(2, 0) return qkv[0], qkv[1], qkv[2] -def optimized_attention(qkv, num_heads): - return attention.optimized_attention(qkv[0], qkv[1], qkv[2], num_heads) class SelfAttention(nn.Module): ATTENTION_MODES = ("xformers", "torch", "torch-hb", "math", "debug") @@ -326,9 +324,9 @@ class SelfAttention(nn.Module): return x def forward(self, x: torch.Tensor) -> torch.Tensor: - qkv = self.pre_attention(x) + q, k, v = self.pre_attention(x) x = optimized_attention( - qkv, num_heads=self.num_heads + q, k, v, heads=self.num_heads ) x = self.post_attention(x) return x @@ -531,8 +529,8 @@ class DismantledBlock(nn.Module): assert not self.pre_only qkv, intermediates = self.pre_attention(x, c) attn = optimized_attention( - qkv, - num_heads=self.attn.num_heads, + qkv[0], qkv[1], qkv[2], + heads=self.attn.num_heads, ) return self.post_attention(attn, *intermediates) @@ -557,8 +555,8 @@ def _block_mixing(context, x, context_block, x_block, c): qkv = tuple(o) attn = optimized_attention( - qkv, - num_heads=x_block.attn.num_heads, + qkv[0], qkv[1], qkv[2], + heads=x_block.attn.num_heads, ) context_attn, x_attn = ( attn[:, : context_qkv[0].shape[1]], @@ -642,7 +640,7 @@ class SelfAttentionContext(nn.Module): def forward(self, x): qkv = self.qkv(x) q, k, v = split_qkv(qkv, self.dim_head) - x = optimized_attention((q.reshape(q.shape[0], q.shape[1], -1), k, v), self.heads) + x = optimized_attention(q.reshape(q.shape[0], q.shape[1], -1), k, v, heads=self.heads) return self.proj(x) class ContextProcessorBlock(nn.Module): diff --git a/comfy/lora.py b/comfy/lora.py index 81cd1696e..b745ca4d5 100644 --- a/comfy/lora.py +++ b/comfy/lora.py @@ -317,6 +317,10 @@ def model_lora_keys_unet(model, key_map={}): key_lora = "lora_transformer_{}".format(k[:-len(".weight")].replace(".", "_")) #OneTrainer lora key_map[key_lora] = to + key_lora = "lycoris_{}".format(k[:-len(".weight")].replace(".", "_")) #simpletuner lycoris format + key_map[key_lora] = to + + if isinstance(model, comfy.model_base.AuraFlow): #Diffusers lora AuraFlow diffusers_keys = comfy.utils.auraflow_to_diffusers(model.model_config.unet_config, output_prefix="diffusion_model.") for k in diffusers_keys: diff --git a/comfy/ops.py b/comfy/ops.py index 5e7c668eb..3c5ba0124 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -264,10 +264,14 @@ def fp8_linear(self, input): scale_input = self.scale_input if scale_weight is None: scale_weight = torch.ones((), device=input.device, dtype=torch.float32) + else: + scale_weight = scale_weight.to(input.device) + if scale_input is None: scale_input = torch.ones((), device=input.device, dtype=torch.float32) inn = input.reshape(-1, input.shape[2]).to(dtype) else: + scale_input = scale_input.to(input.device) inn = (input * (1.0 / scale_input).to(input.dtype)).reshape(-1, input.shape[2]).to(dtype) if bias is not None: