diff --git a/empty_text_embed.pt b/empty_text_embed.pt deleted file mode 100644 index dcedcde..0000000 Binary files a/empty_text_embed.pt and /dev/null differ diff --git a/ldm/modules/attention.py b/ldm/modules/attention.py index f72de57..779ae2d 100644 --- a/ldm/modules/attention.py +++ b/ldm/modules/attention.py @@ -9,6 +9,9 @@ from typing import Optional, Any from ...ldm.modules.diffusionmodules.util import checkpoint from ...ldm import xformers_state +import comfy.ops +ops = comfy.ops.manual_cast + # try: # import xformers # import xformers.ops @@ -49,7 +52,7 @@ def init_(tensor): class GEGLU(nn.Module): def __init__(self, dim_in, dim_out): super().__init__() - self.proj = nn.Linear(dim_in, dim_out * 2) + self.proj = ops.Linear(dim_in, dim_out * 2) def forward(self, x): x, gate = self.proj(x).chunk(2, dim=-1) @@ -62,14 +65,14 @@ class FeedForward(nn.Module): inner_dim = int(dim * mult) dim_out = default(dim_out, dim) project_in = nn.Sequential( - nn.Linear(dim, inner_dim), + ops.Linear(dim, inner_dim), nn.GELU() ) if not glu else GEGLU(dim, inner_dim) self.net = nn.Sequential( project_in, nn.Dropout(dropout), - nn.Linear(inner_dim, dim_out) + ops.Linear(inner_dim, dim_out) ) def forward(self, x): @@ -95,22 +98,22 @@ class SpatialSelfAttention(nn.Module): self.in_channels = in_channels self.norm = Normalize(in_channels) - self.q = torch.nn.Conv2d(in_channels, + self.q = torch.ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.k = torch.nn.Conv2d(in_channels, + self.k = torch.ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.v = torch.nn.Conv2d(in_channels, + self.v = torch.ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.proj_out = torch.nn.Conv2d(in_channels, + self.proj_out = torch.ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, @@ -151,12 +154,12 @@ class CrossAttention(nn.Module): self.scale = dim_head ** -0.5 self.heads = heads - self.to_q = nn.Linear(query_dim, inner_dim, bias=False) - self.to_k = nn.Linear(context_dim, inner_dim, bias=False) - self.to_v = nn.Linear(context_dim, inner_dim, bias=False) + self.to_q = ops.Linear(query_dim, inner_dim, bias=False) + self.to_k = ops.Linear(context_dim, inner_dim, bias=False) + self.to_v = ops.Linear(context_dim, inner_dim, bias=False) self.to_out = nn.Sequential( - nn.Linear(inner_dim, query_dim), + ops.Linear(inner_dim, query_dim), nn.Dropout(dropout) ) @@ -199,19 +202,19 @@ class MemoryEfficientCrossAttention(nn.Module): # https://github.com/MatthieuTPHR/diffusers/blob/d80b531ff8060ec1ea982b65a1b8df70f73aa67c/src/diffusers/models/attention.py#L223 def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0.0): super().__init__() - print(f"Setting up {self.__class__.__name__}. Query dim is {query_dim}, context_dim is {context_dim} and using " - f"{heads} heads.") + #print(f"Setting up {self.__class__.__name__}. Query dim is {query_dim}, context_dim is {context_dim} and using " + # f"{heads} heads.") inner_dim = dim_head * heads context_dim = default(context_dim, query_dim) self.heads = heads self.dim_head = dim_head - self.to_q = nn.Linear(query_dim, inner_dim, bias=False) - self.to_k = nn.Linear(context_dim, inner_dim, bias=False) - self.to_v = nn.Linear(context_dim, inner_dim, bias=False) + self.to_q = ops.Linear(query_dim, inner_dim, bias=False) + self.to_k = ops.Linear(context_dim, inner_dim, bias=False) + self.to_v = ops.Linear(context_dim, inner_dim, bias=False) - self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim), nn.Dropout(dropout)) + self.to_out = nn.Sequential(ops.Linear(inner_dim, query_dim), nn.Dropout(dropout)) self.attention_op: Optional[Any] = None def forward(self, x, context=None, mask=None): @@ -262,9 +265,9 @@ class BasicTransformerBlock(nn.Module): self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff) self.attn2 = attn_cls(query_dim=dim, context_dim=context_dim, heads=n_heads, dim_head=d_head, dropout=dropout) # is self-attn if context is none - self.norm1 = nn.LayerNorm(dim) - self.norm2 = nn.LayerNorm(dim) - self.norm3 = nn.LayerNorm(dim) + self.norm1 = ops.LayerNorm(dim) + self.norm2 = ops.LayerNorm(dim) + self.norm3 = ops.LayerNorm(dim) self.checkpoint = checkpoint def forward(self, x, context=None): @@ -297,13 +300,13 @@ class SpatialTransformer(nn.Module): inner_dim = n_heads * d_head self.norm = Normalize(in_channels) if not use_linear: - self.proj_in = nn.Conv2d(in_channels, + self.proj_in = ops.Conv2d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0) else: - self.proj_in = nn.Linear(in_channels, inner_dim) + self.proj_in = ops.Linear(in_channels, inner_dim) self.transformer_blocks = nn.ModuleList( [BasicTransformerBlock(inner_dim, n_heads, d_head, dropout=dropout, context_dim=context_dim[d], @@ -311,13 +314,13 @@ class SpatialTransformer(nn.Module): for d in range(depth)] ) if not use_linear: - self.proj_out = zero_module(nn.Conv2d(inner_dim, + self.proj_out = zero_module(ops.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)) else: - self.proj_out = zero_module(nn.Linear(in_channels, inner_dim)) + self.proj_out = zero_module(ops.Linear(in_channels, inner_dim)) self.use_linear = use_linear def forward(self, x, context=None): diff --git a/ldm/modules/diffusionmodules/model.py b/ldm/modules/diffusionmodules/model.py index fb61fed..0ce44f7 100644 --- a/ldm/modules/diffusionmodules/model.py +++ b/ldm/modules/diffusionmodules/model.py @@ -8,7 +8,8 @@ from typing import Optional, Any from ldm.modules.attention import MemoryEfficientCrossAttention from ldm import xformers_state - +import comfy.ops +ops = comfy.ops.manual_cast # try: # import xformers @@ -54,7 +55,7 @@ class Upsample(nn.Module): super().__init__() self.with_conv = with_conv if self.with_conv: - self.conv = torch.nn.Conv2d(in_channels, + self.conv = ops.Conv2d(in_channels, in_channels, kernel_size=3, stride=1, @@ -73,7 +74,7 @@ class Downsample(nn.Module): self.with_conv = with_conv if self.with_conv: # no asymmetric padding in torch conv, must do it ourselves - self.conv = torch.nn.Conv2d(in_channels, + self.conv = ops.Conv2d(in_channels, in_channels, kernel_size=3, stride=2, @@ -99,30 +100,30 @@ class ResnetBlock(nn.Module): self.use_conv_shortcut = conv_shortcut self.norm1 = Normalize(in_channels) - self.conv1 = torch.nn.Conv2d(in_channels, + self.conv1 = ops.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) if temb_channels > 0: - self.temb_proj = torch.nn.Linear(temb_channels, + self.temb_proj = ops.Linear(temb_channels, out_channels) self.norm2 = Normalize(out_channels) self.dropout = torch.nn.Dropout(dropout) - self.conv2 = torch.nn.Conv2d(out_channels, + self.conv2 = ops.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1) if self.in_channels != self.out_channels: if self.use_conv_shortcut: - self.conv_shortcut = torch.nn.Conv2d(in_channels, + self.conv_shortcut = ops.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) else: - self.nin_shortcut = torch.nn.Conv2d(in_channels, + self.nin_shortcut = ops.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, @@ -157,22 +158,22 @@ class AttnBlock(nn.Module): self.in_channels = in_channels self.norm = Normalize(in_channels) - self.q = torch.nn.Conv2d(in_channels, + self.q = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.k = torch.nn.Conv2d(in_channels, + self.k = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.v = torch.nn.Conv2d(in_channels, + self.v = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.proj_out = torch.nn.Conv2d(in_channels, + self.proj_out = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, @@ -216,22 +217,22 @@ class MemoryEfficientAttnBlock(nn.Module): self.in_channels = in_channels self.norm = Normalize(in_channels) - self.q = torch.nn.Conv2d(in_channels, + self.q = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.k = torch.nn.Conv2d(in_channels, + self.k = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.v = torch.nn.Conv2d(in_channels, + self.v = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.proj_out = torch.nn.Conv2d(in_channels, + self.proj_out = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, @@ -318,14 +319,14 @@ class Model(nn.Module): # timestep embedding self.temb = nn.Module() self.temb.dense = nn.ModuleList([ - torch.nn.Linear(self.ch, + ops.Linear(self.ch, self.temb_ch), - torch.nn.Linear(self.temb_ch, + ops.Linear(self.temb_ch, self.temb_ch), ]) # downsampling - self.conv_in = torch.nn.Conv2d(in_channels, + self.conv_in = ops.Conv2d(in_channels, self.ch, kernel_size=3, stride=1, @@ -394,7 +395,7 @@ class Model(nn.Module): # end self.norm_out = Normalize(block_in) - self.conv_out = torch.nn.Conv2d(block_in, + self.conv_out = ops.Conv2d(block_in, out_ch, kernel_size=3, stride=1, @@ -467,7 +468,7 @@ class Encoder(nn.Module): self.in_channels = in_channels # downsampling - self.conv_in = torch.nn.Conv2d(in_channels, + self.conv_in = ops.Conv2d(in_channels, self.ch, kernel_size=3, stride=1, @@ -512,7 +513,7 @@ class Encoder(nn.Module): # end self.norm_out = Normalize(block_in) - self.conv_out = torch.nn.Conv2d(block_in, + self.conv_out = ops.Conv2d(block_in, 2*z_channels if double_z else z_channels, kernel_size=3, stride=1, @@ -571,7 +572,7 @@ class Decoder(nn.Module): self.z_shape, np.prod(self.z_shape))) # z to block_in - self.conv_in = torch.nn.Conv2d(z_channels, + self.conv_in = ops.Conv2d(z_channels, block_in, kernel_size=3, stride=1, @@ -613,7 +614,7 @@ class Decoder(nn.Module): # end self.norm_out = Normalize(block_in) - self.conv_out = torch.nn.Conv2d(block_in, + self.conv_out = ops.Conv2d(block_in, out_ch, kernel_size=3, stride=1, @@ -658,7 +659,7 @@ class Decoder(nn.Module): class SimpleDecoder(nn.Module): def __init__(self, in_channels, out_channels, *args, **kwargs): super().__init__() - self.model = nn.ModuleList([nn.Conv2d(in_channels, in_channels, 1), + self.model = nn.ModuleList([ops.Conv2d(in_channels, in_channels, 1), ResnetBlock(in_channels=in_channels, out_channels=2 * in_channels, temb_channels=0, dropout=0.0), @@ -668,11 +669,11 @@ class SimpleDecoder(nn.Module): ResnetBlock(in_channels=4 * in_channels, out_channels=2 * in_channels, temb_channels=0, dropout=0.0), - nn.Conv2d(2*in_channels, in_channels, 1), + ops.Conv2d(2*in_channels, in_channels, 1), Upsample(in_channels, with_conv=True)]) # end self.norm_out = Normalize(in_channels) - self.conv_out = torch.nn.Conv2d(in_channels, + self.conv_out = ops.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, @@ -719,7 +720,7 @@ class UpsampleDecoder(nn.Module): # end self.norm_out = Normalize(block_in) - self.conv_out = torch.nn.Conv2d(block_in, + self.conv_out = ops.Conv2d(block_in, out_channels, kernel_size=3, stride=1, @@ -744,7 +745,7 @@ class LatentRescaler(nn.Module): super().__init__() # residual block, interpolate, residual block self.factor = factor - self.conv_in = nn.Conv2d(in_channels, + self.conv_in = ops.Conv2d(in_channels, mid_channels, kernel_size=3, stride=1, @@ -759,7 +760,7 @@ class LatentRescaler(nn.Module): temb_channels=0, dropout=0.0) for _ in range(depth)]) - self.conv_out = nn.Conv2d(mid_channels, + self.conv_out = ops.Conv2d(mid_channels, out_channels, kernel_size=1, ) @@ -841,7 +842,7 @@ class Resize(nn.Module): raise NotImplementedError() assert in_channels is not None # no asymmetric padding in torch conv, must do it ourselves - self.conv = torch.nn.Conv2d(in_channels, + self.conv = ops.Conv2d(in_channels, in_channels, kernel_size=4, stride=2, diff --git a/ldm/modules/diffusionmodules/util.py b/ldm/modules/diffusionmodules/util.py index 637363d..76d221e 100644 --- a/ldm/modules/diffusionmodules/util.py +++ b/ldm/modules/diffusionmodules/util.py @@ -17,6 +17,9 @@ from einops import repeat from ldm.util import instantiate_from_config +import comfy.ops +ops = comfy.ops.manual_cast + def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3): if schedule == "linear": @@ -223,11 +226,11 @@ def conv_nd(dims, *args, **kwargs): Create a 1D, 2D, or 3D convolution module. """ if dims == 1: - return nn.Conv1d(*args, **kwargs) + return ops.Conv1d(*args, **kwargs) elif dims == 2: - return nn.Conv2d(*args, **kwargs) + return ops.Conv2d(*args, **kwargs) elif dims == 3: - return nn.Conv3d(*args, **kwargs) + return ops.Conv3d(*args, **kwargs) raise ValueError(f"unsupported dimensions: {dims}") @@ -235,7 +238,7 @@ def linear(*args, **kwargs): """ Create a linear module. """ - return nn.Linear(*args, **kwargs) + return ops.Linear(*args, **kwargs) def avg_pool_nd(dims, *args, **kwargs): diff --git a/ldm/modules/encoders/modules.py b/ldm/modules/encoders/modules.py index ab8b31d..9358ff3 100644 --- a/ldm/modules/encoders/modules.py +++ b/ldm/modules/encoders/modules.py @@ -143,8 +143,10 @@ class FrozenOpenCLIPEmbedder(AbstractEncoder): # def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", max_length=77, def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", max_length=77, freeze=True, layer="last"): + super().__init__() assert layer in self.LAYERS + return model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), pretrained=version) del model.visual self.model = model diff --git a/ldm/modules/midas/midas/blocks.py b/ldm/modules/midas/midas/blocks.py index 2145d18..8837f60 100644 --- a/ldm/modules/midas/midas/blocks.py +++ b/ldm/modules/midas/midas/blocks.py @@ -1,6 +1,9 @@ import torch import torch.nn as nn +import comfy.ops +ops = comfy.ops.manual_cast + from .vit import ( _make_pretrained_vitb_rn50_384, _make_pretrained_vitl16_384, @@ -59,16 +62,16 @@ def _make_scratch(in_shape, out_shape, groups=1, expand=False): out_shape3 = out_shape*4 out_shape4 = out_shape*8 - scratch.layer1_rn = nn.Conv2d( + scratch.layer1_rn = ops.Conv2d( in_shape[0], out_shape1, kernel_size=3, stride=1, padding=1, bias=False, groups=groups ) - scratch.layer2_rn = nn.Conv2d( + scratch.layer2_rn = ops.Conv2d( in_shape[1], out_shape2, kernel_size=3, stride=1, padding=1, bias=False, groups=groups ) - scratch.layer3_rn = nn.Conv2d( + scratch.layer3_rn = ops.Conv2d( in_shape[2], out_shape3, kernel_size=3, stride=1, padding=1, bias=False, groups=groups ) - scratch.layer4_rn = nn.Conv2d( + scratch.layer4_rn = ops.Conv2d( in_shape[3], out_shape4, kernel_size=3, stride=1, padding=1, bias=False, groups=groups ) @@ -164,11 +167,11 @@ class ResidualConvUnit(nn.Module): """ super().__init__() - self.conv1 = nn.Conv2d( + self.conv1 = ops.Conv2d( features, features, kernel_size=3, stride=1, padding=1, bias=True ) - self.conv2 = nn.Conv2d( + self.conv2 = ops.Conv2d( features, features, kernel_size=3, stride=1, padding=1, bias=True ) @@ -244,11 +247,11 @@ class ResidualConvUnit_custom(nn.Module): self.groups=1 - self.conv1 = nn.Conv2d( + self.conv1 = ops.Conv2d( features, features, kernel_size=3, stride=1, padding=1, bias=True, groups=self.groups ) - self.conv2 = nn.Conv2d( + self.conv2 = ops.Conv2d( features, features, kernel_size=3, stride=1, padding=1, bias=True, groups=self.groups ) @@ -310,7 +313,7 @@ class FeatureFusionBlock_custom(nn.Module): if self.expand==True: out_features = features//2 - self.out_conv = nn.Conv2d(features, out_features, kernel_size=1, stride=1, padding=0, bias=True, groups=1) + self.out_conv = ops.Conv2d(features, out_features, kernel_size=1, stride=1, padding=0, bias=True, groups=1) self.resConfUnit1 = ResidualConvUnit_custom(features, activation, bn) self.resConfUnit2 = ResidualConvUnit_custom(features, activation, bn) diff --git a/model/q_sampler.py b/model/q_sampler.py index 8c6450c..883e753 100644 --- a/model/q_sampler.py +++ b/model/q_sampler.py @@ -577,7 +577,8 @@ class SpacedSampler: # fuse by tile_weights on noise (score) noise_buffer /= count pred_x0 = self._predict_xstart_from_eps(x_t=img, t=index, eps=noise_buffer) - tao_index = torch.tensor(torch.round(index * t_max), dtype=torch.int64) + tao_index = torch.round(index * t_max).clone().detach().to(torch.int64) + img = self.q_sample(pred_x0, tao_index) @@ -606,11 +607,11 @@ class SpacedSampler: tile_cond_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8] tile_cond = { "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], - "c_crossattn": [self.model.get_learned_conditioning([positive_prompt] * b)] + "c_crossattn": [empty_text_embed * b] } tile_uncond = { "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], - "c_crossattn": [self.model.get_learned_conditioning([negative_prompt] * b)] + "c_crossattn": [empty_text_embed * b] } # predict noise for this tile tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond) @@ -728,11 +729,12 @@ class SpacedSampler: # accumulate noise noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise count[:, :, hi:hi_end, wi:wi_end] += 1 - pbar.update(1) + # average on noise (score) noise_buffer.div_(count) pred_x0 = self._predict_xstart_from_eps(x_t=img, t=index, eps=noise_buffer) - tao_index = torch.tensor(torch.round(index * t_max), dtype=torch.int64) + tao_index = torch.round(index * t_max).clone().detach().to(torch.int64) + img = self.q_sample(pred_x0, tao_index) noise_buffer.zero_() @@ -795,6 +797,7 @@ class SpacedSampler: noise_buffer.zero_() count.zero_() + pbar.update(1) img = pred_x0 # decode samples of each diffusion process diff --git a/nodes.py b/nodes.py index c16ba43..83bee8e 100644 --- a/nodes.py +++ b/nodes.py @@ -9,13 +9,12 @@ from .model.ccsr_stage1 import ControlLDM from .utils.common import instantiate_from_config, load_state_dict -import comfy.model_management +import comfy.model_management as mm import comfy.utils import folder_paths from nodes import ImageScaleBy from nodes import ImageScale - script_directory = os.path.dirname(os.path.abspath(__file__)) class CCSR_Upscale: @@ -64,29 +63,24 @@ class CCSR_Upscale: CATEGORY = "CCSR" @torch.no_grad() - def process(self, ccsr_model, image, resize_method, scale_by, steps, t_max, t_min, tile_size, tile_stride, color_fix_type, keep_model_loaded, vae_tile_size_encode, vae_tile_size_decode, sampling_method, seed): + def process(self, ccsr_model, image, resize_method, scale_by, steps, t_max, t_min, tile_size, tile_stride, + color_fix_type, keep_model_loaded, vae_tile_size_encode, vae_tile_size_decode, sampling_method, seed): torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) - comfy.model_management.unload_all_models() - device = comfy.model_management.get_torch_device() - config_path = os.path.join(script_directory, "configs/model/ccsr_stage2.yaml") - empty_text_embed = torch.load(os.path.join(script_directory, "empty_text_embed.pt"), map_location=device) - dtype = torch.float16 if comfy.model_management.should_use_fp16() and not comfy.model_management.is_device_mps(device) else torch.float32 - if not hasattr(self, "model") or self.model is None: - config = OmegaConf.load(config_path) - self.model = instantiate_from_config(config) - - load_state_dict(self.model, comfy.utils.load_torch_file(ccsr_model), strict=True) - # reload preprocess model if specified + mm.unload_all_models() + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + dtype = ccsr_model['dtype'] + model = ccsr_model['model'] + + #empty_text_embed = torch.load(os.path.join(script_directory, "empty_text_embed.pt"), map_location=device) + empty_text_embed_sd = comfy.utils.load_torch_file(os.path.join(script_directory, "empty_text_embed.safetensors")) + empty_text_embed = empty_text_embed_sd['empty_text_embed'].to(dtype).to(device) - self.model.freeze() - self.model.to(device, dtype=dtype) - sampler = SpacedSampler(self.model, var_type="fixed_small") + sampler = SpacedSampler(model, var_type="fixed_small") - batch_size = image.shape[0] image, = ImageScaleBy.upscale(self, image, resize_method, scale_by) - # Assuming 'image' is a PyTorch tensor with shape [B, H, W, C] and you want to resize it. B, H, W, C = image.shape # Calculate the new height and width, rounding down to the nearest multiple of 64. @@ -95,28 +89,26 @@ class CCSR_Upscale: # Reorder to [B, C, H, W] before using interpolate. image = image.permute(0, 3, 1, 2).contiguous() - - # Resize the image tensor. - resized_image = F.interpolate(image, size=(new_height, new_width), mode='bicubic', align_corners=False) - - # Move the tensor to the GPU. - #resized_image = resized_image.to(device) + resized_image = F.interpolate(image, size=(new_height, new_width), mode='bilinear', align_corners=False) + strength = 1.0 - self.model.control_scales = [strength] * 13 + model.control_scales = [strength] * 13 + model.to(device, dtype=dtype).eval() + height, width = resized_image.size(-2), resized_image.size(-1) shape = (1, 4, height // 8, width // 8) - x_T = torch.randn(shape, device=self.model.device, dtype=torch.float32) - autocast_condition = dtype == torch.float16 and not comfy.model_management.is_device_mps(device) - out = [] + x_T = torch.randn(shape, device=model.device, dtype=torch.float32) - pbar = comfy.utils.ProgressBar(batch_size) - - with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext(): - for i in range(batch_size): + out = [] + if B > 1: + pbar = comfy.utils.ProgressBar(B) + autocast_condition = dtype == torch.float16 and not mm.is_device_mps(device) + with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext(): + for i in range(B): img = resized_image[i].unsqueeze(0).to(device) if sampling_method == 'ccsr_tiled_mixdiff': - self.model.reset_encoder_decoder() + model.reset_encoder_decoder() print("Using tiled mixdiff") samples = sampler.sample_with_mixdiff_ccsr( empty_text_embed, tile_size=tile_size, tile_stride=tile_stride, @@ -126,7 +118,7 @@ class CCSR_Upscale: color_fix_type=color_fix_type ) elif sampling_method == 'ccsr_tiled_vae_gaussian_weights': - self.model._init_tiled_vae(encoder_tile_size=vae_tile_size_encode // 8, decoder_tile_size=vae_tile_size_decode // 8) + model._init_tiled_vae(encoder_tile_size=vae_tile_size_encode // 8, decoder_tile_size=vae_tile_size_decode // 8) print("Using gaussian weights") samples = sampler.sample_with_tile_ccsr( empty_text_embed, tile_size=tile_size, tile_stride=tile_stride, @@ -136,7 +128,7 @@ class CCSR_Upscale: color_fix_type=color_fix_type ) else: - self.model.reset_encoder_decoder() + model.reset_encoder_decoder() print("no tiling") samples = sampler.sample_ccsr( empty_text_embed, steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=img, @@ -145,9 +137,10 @@ class CCSR_Upscale: color_fix_type=color_fix_type ) out.append(samples.squeeze(0).cpu()) - comfy.model_management.throw_exception_if_processing_interrupted() - pbar.update(1) - print("Sampled image ", i, " out of ", batch_size) + mm.throw_exception_if_processing_interrupted() + if B > 1: + pbar.update(1) + print("Sampled image ", i, " out of ", B) original_height, original_width = H, W processed_height = samples.size(2) @@ -156,8 +149,8 @@ class CCSR_Upscale: resized_back_image, = ImageScale.upscale(self, out_stacked, "lanczos", target_width, processed_height, crop="disabled") if not keep_model_loaded: - self.model = None - comfy.model_management.soft_empty_cache() + model.to(offload_device) + mm.soft_empty_cache() return(resized_back_image,) class CCSR_Model_Select: @@ -173,15 +166,85 @@ class CCSR_Model_Select: CATEGORY = "CCSR" def load_ccsr_checkpoint(self, ckpt_name): + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) + config_path = os.path.join(script_directory, "configs/model/ccsr_stage2.yaml") + dtype = torch.float16 if mm.should_use_fp16() and not mm.is_device_mps(device) else torch.float32 - return (ckpt_path,) + if not hasattr(self, "model") or self.model is None: + config = OmegaConf.load(config_path) + self.model = instantiate_from_config(config) + + load_state_dict(self.model, comfy.utils.load_torch_file(ckpt_path), strict=True) + # reload preprocess model if specified + + ccsr_model = { + 'model': self.model, + 'dtype': dtype + } + return (ccsr_model,) + +class DownloadAndLoadCCSRModel: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "model": ( + [ + 'real-world_ccsr-fp16.safetensors', + 'real-world_ccsr-fp32.safetensors' + ], + ), + }, + } + + RETURN_TYPES = ("CCSRMODEL",) + RETURN_NAMES = ("ccsr_model",) + FUNCTION = "loadmodel" + CATEGORY = "CCSR" + + def loadmodel(self, model): + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + dtype = torch.float16 if 'fp16' in model else torch.float32 + + model_path = os.path.join(folder_paths.models_dir, "CCSR") + safetensors_path = os.path.join(model_path, model) + + if not os.path.exists(safetensors_path): + print(f"Downloading CCSR model to: {model_path}") + from huggingface_hub import snapshot_download + snapshot_download(repo_id="Kijai/ccsr-safetensors", + allow_patterns=[f'*{model}*'], + local_dir=model_path, + local_dir_use_symlinks=False) + + config_path = os.path.join(script_directory, "configs/model/ccsr_stage2.yaml") + config = OmegaConf.load(config_path) + + model = instantiate_from_config(config) + + sd = comfy.utils.load_torch_file(safetensors_path) + + model.load_state_dict(sd, strict=False) + del sd + mm.soft_empty_cache() + + ccsr_model = { + 'model': model, + 'dtype': dtype, + } + + return (ccsr_model,) + NODE_CLASS_MAPPINGS = { "CCSR_Upscale": CCSR_Upscale, - "CCSR_Model_Select": CCSR_Model_Select + "CCSR_Model_Select": CCSR_Model_Select, + "DownloadAndLoadCCSRModel": DownloadAndLoadCCSRModel } NODE_DISPLAY_NAME_MAPPINGS = { "CCSR_Upscale": "CCSR_Upscale", - "CCSR_Model_Select": "CCSR_Model_Select" + "CCSR_Model_Select": "CCSR_Model_Select", + "DownloadAndLoadCCSRModel": "DownloadAndLoad CCSRModel" } \ No newline at end of file