Faster model loading, autodownload node, optimizations
This commit is contained in:
Binary file not shown.
+27
-24
@@ -9,6 +9,9 @@ from typing import Optional, Any
|
|||||||
from ...ldm.modules.diffusionmodules.util import checkpoint
|
from ...ldm.modules.diffusionmodules.util import checkpoint
|
||||||
from ...ldm import xformers_state
|
from ...ldm import xformers_state
|
||||||
|
|
||||||
|
import comfy.ops
|
||||||
|
ops = comfy.ops.manual_cast
|
||||||
|
|
||||||
# try:
|
# try:
|
||||||
# import xformers
|
# import xformers
|
||||||
# import xformers.ops
|
# import xformers.ops
|
||||||
@@ -49,7 +52,7 @@ def init_(tensor):
|
|||||||
class GEGLU(nn.Module):
|
class GEGLU(nn.Module):
|
||||||
def __init__(self, dim_in, dim_out):
|
def __init__(self, dim_in, dim_out):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.proj = nn.Linear(dim_in, dim_out * 2)
|
self.proj = ops.Linear(dim_in, dim_out * 2)
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
x, gate = self.proj(x).chunk(2, dim=-1)
|
x, gate = self.proj(x).chunk(2, dim=-1)
|
||||||
@@ -62,14 +65,14 @@ class FeedForward(nn.Module):
|
|||||||
inner_dim = int(dim * mult)
|
inner_dim = int(dim * mult)
|
||||||
dim_out = default(dim_out, dim)
|
dim_out = default(dim_out, dim)
|
||||||
project_in = nn.Sequential(
|
project_in = nn.Sequential(
|
||||||
nn.Linear(dim, inner_dim),
|
ops.Linear(dim, inner_dim),
|
||||||
nn.GELU()
|
nn.GELU()
|
||||||
) if not glu else GEGLU(dim, inner_dim)
|
) if not glu else GEGLU(dim, inner_dim)
|
||||||
|
|
||||||
self.net = nn.Sequential(
|
self.net = nn.Sequential(
|
||||||
project_in,
|
project_in,
|
||||||
nn.Dropout(dropout),
|
nn.Dropout(dropout),
|
||||||
nn.Linear(inner_dim, dim_out)
|
ops.Linear(inner_dim, dim_out)
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
@@ -95,22 +98,22 @@ class SpatialSelfAttention(nn.Module):
|
|||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
|
|
||||||
self.norm = Normalize(in_channels)
|
self.norm = Normalize(in_channels)
|
||||||
self.q = torch.nn.Conv2d(in_channels,
|
self.q = torch.ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=0)
|
padding=0)
|
||||||
self.k = torch.nn.Conv2d(in_channels,
|
self.k = torch.ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=0)
|
padding=0)
|
||||||
self.v = torch.nn.Conv2d(in_channels,
|
self.v = torch.ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=0)
|
padding=0)
|
||||||
self.proj_out = torch.nn.Conv2d(in_channels,
|
self.proj_out = torch.ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -151,12 +154,12 @@ class CrossAttention(nn.Module):
|
|||||||
self.scale = dim_head ** -0.5
|
self.scale = dim_head ** -0.5
|
||||||
self.heads = heads
|
self.heads = heads
|
||||||
|
|
||||||
self.to_q = nn.Linear(query_dim, inner_dim, bias=False)
|
self.to_q = ops.Linear(query_dim, inner_dim, bias=False)
|
||||||
self.to_k = nn.Linear(context_dim, inner_dim, bias=False)
|
self.to_k = ops.Linear(context_dim, inner_dim, bias=False)
|
||||||
self.to_v = nn.Linear(context_dim, inner_dim, bias=False)
|
self.to_v = ops.Linear(context_dim, inner_dim, bias=False)
|
||||||
|
|
||||||
self.to_out = nn.Sequential(
|
self.to_out = nn.Sequential(
|
||||||
nn.Linear(inner_dim, query_dim),
|
ops.Linear(inner_dim, query_dim),
|
||||||
nn.Dropout(dropout)
|
nn.Dropout(dropout)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -199,19 +202,19 @@ class MemoryEfficientCrossAttention(nn.Module):
|
|||||||
# https://github.com/MatthieuTPHR/diffusers/blob/d80b531ff8060ec1ea982b65a1b8df70f73aa67c/src/diffusers/models/attention.py#L223
|
# 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):
|
def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0.0):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
print(f"Setting up {self.__class__.__name__}. Query dim is {query_dim}, context_dim is {context_dim} and using "
|
#print(f"Setting up {self.__class__.__name__}. Query dim is {query_dim}, context_dim is {context_dim} and using "
|
||||||
f"{heads} heads.")
|
# f"{heads} heads.")
|
||||||
inner_dim = dim_head * heads
|
inner_dim = dim_head * heads
|
||||||
context_dim = default(context_dim, query_dim)
|
context_dim = default(context_dim, query_dim)
|
||||||
|
|
||||||
self.heads = heads
|
self.heads = heads
|
||||||
self.dim_head = dim_head
|
self.dim_head = dim_head
|
||||||
|
|
||||||
self.to_q = nn.Linear(query_dim, inner_dim, bias=False)
|
self.to_q = ops.Linear(query_dim, inner_dim, bias=False)
|
||||||
self.to_k = nn.Linear(context_dim, inner_dim, bias=False)
|
self.to_k = ops.Linear(context_dim, inner_dim, bias=False)
|
||||||
self.to_v = nn.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
|
self.attention_op: Optional[Any] = None
|
||||||
|
|
||||||
def forward(self, x, context=None, mask=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.ff = FeedForward(dim, dropout=dropout, glu=gated_ff)
|
||||||
self.attn2 = attn_cls(query_dim=dim, context_dim=context_dim,
|
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
|
heads=n_heads, dim_head=d_head, dropout=dropout) # is self-attn if context is none
|
||||||
self.norm1 = nn.LayerNorm(dim)
|
self.norm1 = ops.LayerNorm(dim)
|
||||||
self.norm2 = nn.LayerNorm(dim)
|
self.norm2 = ops.LayerNorm(dim)
|
||||||
self.norm3 = nn.LayerNorm(dim)
|
self.norm3 = ops.LayerNorm(dim)
|
||||||
self.checkpoint = checkpoint
|
self.checkpoint = checkpoint
|
||||||
|
|
||||||
def forward(self, x, context=None):
|
def forward(self, x, context=None):
|
||||||
@@ -297,13 +300,13 @@ class SpatialTransformer(nn.Module):
|
|||||||
inner_dim = n_heads * d_head
|
inner_dim = n_heads * d_head
|
||||||
self.norm = Normalize(in_channels)
|
self.norm = Normalize(in_channels)
|
||||||
if not use_linear:
|
if not use_linear:
|
||||||
self.proj_in = nn.Conv2d(in_channels,
|
self.proj_in = ops.Conv2d(in_channels,
|
||||||
inner_dim,
|
inner_dim,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=0)
|
padding=0)
|
||||||
else:
|
else:
|
||||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
self.proj_in = ops.Linear(in_channels, inner_dim)
|
||||||
|
|
||||||
self.transformer_blocks = nn.ModuleList(
|
self.transformer_blocks = nn.ModuleList(
|
||||||
[BasicTransformerBlock(inner_dim, n_heads, d_head, dropout=dropout, context_dim=context_dim[d],
|
[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)]
|
for d in range(depth)]
|
||||||
)
|
)
|
||||||
if not use_linear:
|
if not use_linear:
|
||||||
self.proj_out = zero_module(nn.Conv2d(inner_dim,
|
self.proj_out = zero_module(ops.Conv2d(inner_dim,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=0))
|
padding=0))
|
||||||
else:
|
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
|
self.use_linear = use_linear
|
||||||
|
|
||||||
def forward(self, x, context=None):
|
def forward(self, x, context=None):
|
||||||
|
|||||||
@@ -8,7 +8,8 @@ from typing import Optional, Any
|
|||||||
|
|
||||||
from ldm.modules.attention import MemoryEfficientCrossAttention
|
from ldm.modules.attention import MemoryEfficientCrossAttention
|
||||||
from ldm import xformers_state
|
from ldm import xformers_state
|
||||||
|
import comfy.ops
|
||||||
|
ops = comfy.ops.manual_cast
|
||||||
|
|
||||||
# try:
|
# try:
|
||||||
# import xformers
|
# import xformers
|
||||||
@@ -54,7 +55,7 @@ class Upsample(nn.Module):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.with_conv = with_conv
|
self.with_conv = with_conv
|
||||||
if self.with_conv:
|
if self.with_conv:
|
||||||
self.conv = torch.nn.Conv2d(in_channels,
|
self.conv = ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -73,7 +74,7 @@ class Downsample(nn.Module):
|
|||||||
self.with_conv = with_conv
|
self.with_conv = with_conv
|
||||||
if self.with_conv:
|
if self.with_conv:
|
||||||
# no asymmetric padding in torch conv, must do it ourselves
|
# 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,
|
in_channels,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=2,
|
stride=2,
|
||||||
@@ -99,30 +100,30 @@ class ResnetBlock(nn.Module):
|
|||||||
self.use_conv_shortcut = conv_shortcut
|
self.use_conv_shortcut = conv_shortcut
|
||||||
|
|
||||||
self.norm1 = Normalize(in_channels)
|
self.norm1 = Normalize(in_channels)
|
||||||
self.conv1 = torch.nn.Conv2d(in_channels,
|
self.conv1 = ops.Conv2d(in_channels,
|
||||||
out_channels,
|
out_channels,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=1)
|
padding=1)
|
||||||
if temb_channels > 0:
|
if temb_channels > 0:
|
||||||
self.temb_proj = torch.nn.Linear(temb_channels,
|
self.temb_proj = ops.Linear(temb_channels,
|
||||||
out_channels)
|
out_channels)
|
||||||
self.norm2 = Normalize(out_channels)
|
self.norm2 = Normalize(out_channels)
|
||||||
self.dropout = torch.nn.Dropout(dropout)
|
self.dropout = torch.nn.Dropout(dropout)
|
||||||
self.conv2 = torch.nn.Conv2d(out_channels,
|
self.conv2 = ops.Conv2d(out_channels,
|
||||||
out_channels,
|
out_channels,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=1)
|
padding=1)
|
||||||
if self.in_channels != self.out_channels:
|
if self.in_channels != self.out_channels:
|
||||||
if self.use_conv_shortcut:
|
if self.use_conv_shortcut:
|
||||||
self.conv_shortcut = torch.nn.Conv2d(in_channels,
|
self.conv_shortcut = ops.Conv2d(in_channels,
|
||||||
out_channels,
|
out_channels,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=1)
|
padding=1)
|
||||||
else:
|
else:
|
||||||
self.nin_shortcut = torch.nn.Conv2d(in_channels,
|
self.nin_shortcut = ops.Conv2d(in_channels,
|
||||||
out_channels,
|
out_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -157,22 +158,22 @@ class AttnBlock(nn.Module):
|
|||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
|
|
||||||
self.norm = Normalize(in_channels)
|
self.norm = Normalize(in_channels)
|
||||||
self.q = torch.nn.Conv2d(in_channels,
|
self.q = ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=0)
|
padding=0)
|
||||||
self.k = torch.nn.Conv2d(in_channels,
|
self.k = ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=0)
|
padding=0)
|
||||||
self.v = torch.nn.Conv2d(in_channels,
|
self.v = ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=0)
|
padding=0)
|
||||||
self.proj_out = torch.nn.Conv2d(in_channels,
|
self.proj_out = ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -216,22 +217,22 @@ class MemoryEfficientAttnBlock(nn.Module):
|
|||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
|
|
||||||
self.norm = Normalize(in_channels)
|
self.norm = Normalize(in_channels)
|
||||||
self.q = torch.nn.Conv2d(in_channels,
|
self.q = ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=0)
|
padding=0)
|
||||||
self.k = torch.nn.Conv2d(in_channels,
|
self.k = ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=0)
|
padding=0)
|
||||||
self.v = torch.nn.Conv2d(in_channels,
|
self.v = ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
padding=0)
|
padding=0)
|
||||||
self.proj_out = torch.nn.Conv2d(in_channels,
|
self.proj_out = ops.Conv2d(in_channels,
|
||||||
in_channels,
|
in_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -318,14 +319,14 @@ class Model(nn.Module):
|
|||||||
# timestep embedding
|
# timestep embedding
|
||||||
self.temb = nn.Module()
|
self.temb = nn.Module()
|
||||||
self.temb.dense = nn.ModuleList([
|
self.temb.dense = nn.ModuleList([
|
||||||
torch.nn.Linear(self.ch,
|
ops.Linear(self.ch,
|
||||||
self.temb_ch),
|
self.temb_ch),
|
||||||
torch.nn.Linear(self.temb_ch,
|
ops.Linear(self.temb_ch,
|
||||||
self.temb_ch),
|
self.temb_ch),
|
||||||
])
|
])
|
||||||
|
|
||||||
# downsampling
|
# downsampling
|
||||||
self.conv_in = torch.nn.Conv2d(in_channels,
|
self.conv_in = ops.Conv2d(in_channels,
|
||||||
self.ch,
|
self.ch,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -394,7 +395,7 @@ class Model(nn.Module):
|
|||||||
|
|
||||||
# end
|
# end
|
||||||
self.norm_out = Normalize(block_in)
|
self.norm_out = Normalize(block_in)
|
||||||
self.conv_out = torch.nn.Conv2d(block_in,
|
self.conv_out = ops.Conv2d(block_in,
|
||||||
out_ch,
|
out_ch,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -467,7 +468,7 @@ class Encoder(nn.Module):
|
|||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
|
|
||||||
# downsampling
|
# downsampling
|
||||||
self.conv_in = torch.nn.Conv2d(in_channels,
|
self.conv_in = ops.Conv2d(in_channels,
|
||||||
self.ch,
|
self.ch,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -512,7 +513,7 @@ class Encoder(nn.Module):
|
|||||||
|
|
||||||
# end
|
# end
|
||||||
self.norm_out = Normalize(block_in)
|
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,
|
2*z_channels if double_z else z_channels,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -571,7 +572,7 @@ class Decoder(nn.Module):
|
|||||||
self.z_shape, np.prod(self.z_shape)))
|
self.z_shape, np.prod(self.z_shape)))
|
||||||
|
|
||||||
# z to block_in
|
# z to block_in
|
||||||
self.conv_in = torch.nn.Conv2d(z_channels,
|
self.conv_in = ops.Conv2d(z_channels,
|
||||||
block_in,
|
block_in,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -613,7 +614,7 @@ class Decoder(nn.Module):
|
|||||||
|
|
||||||
# end
|
# end
|
||||||
self.norm_out = Normalize(block_in)
|
self.norm_out = Normalize(block_in)
|
||||||
self.conv_out = torch.nn.Conv2d(block_in,
|
self.conv_out = ops.Conv2d(block_in,
|
||||||
out_ch,
|
out_ch,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -658,7 +659,7 @@ class Decoder(nn.Module):
|
|||||||
class SimpleDecoder(nn.Module):
|
class SimpleDecoder(nn.Module):
|
||||||
def __init__(self, in_channels, out_channels, *args, **kwargs):
|
def __init__(self, in_channels, out_channels, *args, **kwargs):
|
||||||
super().__init__()
|
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,
|
ResnetBlock(in_channels=in_channels,
|
||||||
out_channels=2 * in_channels,
|
out_channels=2 * in_channels,
|
||||||
temb_channels=0, dropout=0.0),
|
temb_channels=0, dropout=0.0),
|
||||||
@@ -668,11 +669,11 @@ class SimpleDecoder(nn.Module):
|
|||||||
ResnetBlock(in_channels=4 * in_channels,
|
ResnetBlock(in_channels=4 * in_channels,
|
||||||
out_channels=2 * in_channels,
|
out_channels=2 * in_channels,
|
||||||
temb_channels=0, dropout=0.0),
|
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)])
|
Upsample(in_channels, with_conv=True)])
|
||||||
# end
|
# end
|
||||||
self.norm_out = Normalize(in_channels)
|
self.norm_out = Normalize(in_channels)
|
||||||
self.conv_out = torch.nn.Conv2d(in_channels,
|
self.conv_out = ops.Conv2d(in_channels,
|
||||||
out_channels,
|
out_channels,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -719,7 +720,7 @@ class UpsampleDecoder(nn.Module):
|
|||||||
|
|
||||||
# end
|
# end
|
||||||
self.norm_out = Normalize(block_in)
|
self.norm_out = Normalize(block_in)
|
||||||
self.conv_out = torch.nn.Conv2d(block_in,
|
self.conv_out = ops.Conv2d(block_in,
|
||||||
out_channels,
|
out_channels,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -744,7 +745,7 @@ class LatentRescaler(nn.Module):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
# residual block, interpolate, residual block
|
# residual block, interpolate, residual block
|
||||||
self.factor = factor
|
self.factor = factor
|
||||||
self.conv_in = nn.Conv2d(in_channels,
|
self.conv_in = ops.Conv2d(in_channels,
|
||||||
mid_channels,
|
mid_channels,
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
stride=1,
|
stride=1,
|
||||||
@@ -759,7 +760,7 @@ class LatentRescaler(nn.Module):
|
|||||||
temb_channels=0,
|
temb_channels=0,
|
||||||
dropout=0.0) for _ in range(depth)])
|
dropout=0.0) for _ in range(depth)])
|
||||||
|
|
||||||
self.conv_out = nn.Conv2d(mid_channels,
|
self.conv_out = ops.Conv2d(mid_channels,
|
||||||
out_channels,
|
out_channels,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
)
|
)
|
||||||
@@ -841,7 +842,7 @@ class Resize(nn.Module):
|
|||||||
raise NotImplementedError()
|
raise NotImplementedError()
|
||||||
assert in_channels is not None
|
assert in_channels is not None
|
||||||
# no asymmetric padding in torch conv, must do it ourselves
|
# 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,
|
in_channels,
|
||||||
kernel_size=4,
|
kernel_size=4,
|
||||||
stride=2,
|
stride=2,
|
||||||
|
|||||||
@@ -17,6 +17,9 @@ from einops import repeat
|
|||||||
|
|
||||||
from ldm.util import instantiate_from_config
|
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):
|
def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3):
|
||||||
if schedule == "linear":
|
if schedule == "linear":
|
||||||
@@ -223,11 +226,11 @@ def conv_nd(dims, *args, **kwargs):
|
|||||||
Create a 1D, 2D, or 3D convolution module.
|
Create a 1D, 2D, or 3D convolution module.
|
||||||
"""
|
"""
|
||||||
if dims == 1:
|
if dims == 1:
|
||||||
return nn.Conv1d(*args, **kwargs)
|
return ops.Conv1d(*args, **kwargs)
|
||||||
elif dims == 2:
|
elif dims == 2:
|
||||||
return nn.Conv2d(*args, **kwargs)
|
return ops.Conv2d(*args, **kwargs)
|
||||||
elif dims == 3:
|
elif dims == 3:
|
||||||
return nn.Conv3d(*args, **kwargs)
|
return ops.Conv3d(*args, **kwargs)
|
||||||
raise ValueError(f"unsupported dimensions: {dims}")
|
raise ValueError(f"unsupported dimensions: {dims}")
|
||||||
|
|
||||||
|
|
||||||
@@ -235,7 +238,7 @@ def linear(*args, **kwargs):
|
|||||||
"""
|
"""
|
||||||
Create a linear module.
|
Create a linear module.
|
||||||
"""
|
"""
|
||||||
return nn.Linear(*args, **kwargs)
|
return ops.Linear(*args, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
def avg_pool_nd(dims, *args, **kwargs):
|
def avg_pool_nd(dims, *args, **kwargs):
|
||||||
|
|||||||
@@ -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", device="cuda", max_length=77,
|
||||||
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", max_length=77,
|
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", max_length=77,
|
||||||
freeze=True, layer="last"):
|
freeze=True, layer="last"):
|
||||||
|
|
||||||
super().__init__()
|
super().__init__()
|
||||||
assert layer in self.LAYERS
|
assert layer in self.LAYERS
|
||||||
|
return
|
||||||
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), pretrained=version)
|
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), pretrained=version)
|
||||||
del model.visual
|
del model.visual
|
||||||
self.model = model
|
self.model = model
|
||||||
|
|||||||
@@ -1,6 +1,9 @@
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
|
import comfy.ops
|
||||||
|
ops = comfy.ops.manual_cast
|
||||||
|
|
||||||
from .vit import (
|
from .vit import (
|
||||||
_make_pretrained_vitb_rn50_384,
|
_make_pretrained_vitb_rn50_384,
|
||||||
_make_pretrained_vitl16_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_shape3 = out_shape*4
|
||||||
out_shape4 = out_shape*8
|
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
|
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
|
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
|
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
|
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__()
|
super().__init__()
|
||||||
|
|
||||||
self.conv1 = nn.Conv2d(
|
self.conv1 = ops.Conv2d(
|
||||||
features, features, kernel_size=3, stride=1, padding=1, bias=True
|
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
|
features, features, kernel_size=3, stride=1, padding=1, bias=True
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -244,11 +247,11 @@ class ResidualConvUnit_custom(nn.Module):
|
|||||||
|
|
||||||
self.groups=1
|
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
|
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
|
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:
|
if self.expand==True:
|
||||||
out_features = features//2
|
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.resConfUnit1 = ResidualConvUnit_custom(features, activation, bn)
|
||||||
self.resConfUnit2 = ResidualConvUnit_custom(features, activation, bn)
|
self.resConfUnit2 = ResidualConvUnit_custom(features, activation, bn)
|
||||||
|
|||||||
+8
-5
@@ -577,7 +577,8 @@ class SpacedSampler:
|
|||||||
# fuse by tile_weights on noise (score)
|
# fuse by tile_weights on noise (score)
|
||||||
noise_buffer /= count
|
noise_buffer /= count
|
||||||
pred_x0 = self._predict_xstart_from_eps(x_t=img, t=index, eps=noise_buffer)
|
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)
|
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_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8]
|
||||||
tile_cond = {
|
tile_cond = {
|
||||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
"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 = {
|
tile_uncond = {
|
||||||
"c_latent": [self.model.apply_condition_encoder(tile_cond_img)],
|
"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
|
# predict noise for this tile
|
||||||
tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond)
|
tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond)
|
||||||
@@ -728,11 +729,12 @@ class SpacedSampler:
|
|||||||
# accumulate noise
|
# accumulate noise
|
||||||
noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise
|
noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise
|
||||||
count[:, :, hi:hi_end, wi:wi_end] += 1
|
count[:, :, hi:hi_end, wi:wi_end] += 1
|
||||||
pbar.update(1)
|
|
||||||
# average on noise (score)
|
# average on noise (score)
|
||||||
noise_buffer.div_(count)
|
noise_buffer.div_(count)
|
||||||
pred_x0 = self._predict_xstart_from_eps(x_t=img, t=index, eps=noise_buffer)
|
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)
|
img = self.q_sample(pred_x0, tao_index)
|
||||||
|
|
||||||
noise_buffer.zero_()
|
noise_buffer.zero_()
|
||||||
@@ -795,6 +797,7 @@ class SpacedSampler:
|
|||||||
|
|
||||||
noise_buffer.zero_()
|
noise_buffer.zero_()
|
||||||
count.zero_()
|
count.zero_()
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
img = pred_x0
|
img = pred_x0
|
||||||
# decode samples of each diffusion process
|
# decode samples of each diffusion process
|
||||||
|
|||||||
@@ -9,13 +9,12 @@ from .model.ccsr_stage1 import ControlLDM
|
|||||||
|
|
||||||
from .utils.common import instantiate_from_config, load_state_dict
|
from .utils.common import instantiate_from_config, load_state_dict
|
||||||
|
|
||||||
import comfy.model_management
|
import comfy.model_management as mm
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
import folder_paths
|
import folder_paths
|
||||||
from nodes import ImageScaleBy
|
from nodes import ImageScaleBy
|
||||||
from nodes import ImageScale
|
from nodes import ImageScale
|
||||||
|
|
||||||
|
|
||||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
|
||||||
class CCSR_Upscale:
|
class CCSR_Upscale:
|
||||||
@@ -64,29 +63,24 @@ class CCSR_Upscale:
|
|||||||
CATEGORY = "CCSR"
|
CATEGORY = "CCSR"
|
||||||
|
|
||||||
@torch.no_grad()
|
@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.manual_seed(seed)
|
||||||
torch.cuda.manual_seed_all(seed)
|
torch.cuda.manual_seed_all(seed)
|
||||||
comfy.model_management.unload_all_models()
|
mm.unload_all_models()
|
||||||
device = comfy.model_management.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
config_path = os.path.join(script_directory, "configs/model/ccsr_stage2.yaml")
|
offload_device = mm.unet_offload_device()
|
||||||
empty_text_embed = torch.load(os.path.join(script_directory, "empty_text_embed.pt"), map_location=device)
|
dtype = ccsr_model['dtype']
|
||||||
dtype = torch.float16 if comfy.model_management.should_use_fp16() and not comfy.model_management.is_device_mps(device) else torch.float32
|
model = ccsr_model['model']
|
||||||
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)
|
#empty_text_embed = torch.load(os.path.join(script_directory, "empty_text_embed.pt"), map_location=device)
|
||||||
# reload preprocess model if specified
|
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()
|
sampler = SpacedSampler(model, var_type="fixed_small")
|
||||||
self.model.to(device, dtype=dtype)
|
|
||||||
sampler = SpacedSampler(self.model, var_type="fixed_small")
|
|
||||||
|
|
||||||
batch_size = image.shape[0]
|
|
||||||
image, = ImageScaleBy.upscale(self, image, resize_method, scale_by)
|
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
|
B, H, W, C = image.shape
|
||||||
|
|
||||||
# Calculate the new height and width, rounding down to the nearest multiple of 64.
|
# 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.
|
# Reorder to [B, C, H, W] before using interpolate.
|
||||||
image = image.permute(0, 3, 1, 2).contiguous()
|
image = image.permute(0, 3, 1, 2).contiguous()
|
||||||
|
resized_image = F.interpolate(image, size=(new_height, new_width), mode='bilinear', align_corners=False)
|
||||||
|
|
||||||
# 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)
|
|
||||||
strength = 1.0
|
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)
|
height, width = resized_image.size(-2), resized_image.size(-1)
|
||||||
shape = (1, 4, height // 8, width // 8)
|
shape = (1, 4, height // 8, width // 8)
|
||||||
x_T = torch.randn(shape, device=self.model.device, dtype=torch.float32)
|
x_T = torch.randn(shape, device=model.device, dtype=torch.float32)
|
||||||
autocast_condition = dtype == torch.float16 and not comfy.model_management.is_device_mps(device)
|
|
||||||
out = []
|
out = []
|
||||||
|
if B > 1:
|
||||||
pbar = comfy.utils.ProgressBar(batch_size)
|
pbar = comfy.utils.ProgressBar(B)
|
||||||
|
autocast_condition = dtype == torch.float16 and not mm.is_device_mps(device)
|
||||||
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
|
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
|
||||||
for i in range(batch_size):
|
for i in range(B):
|
||||||
img = resized_image[i].unsqueeze(0).to(device)
|
img = resized_image[i].unsqueeze(0).to(device)
|
||||||
if sampling_method == 'ccsr_tiled_mixdiff':
|
if sampling_method == 'ccsr_tiled_mixdiff':
|
||||||
self.model.reset_encoder_decoder()
|
model.reset_encoder_decoder()
|
||||||
print("Using tiled mixdiff")
|
print("Using tiled mixdiff")
|
||||||
samples = sampler.sample_with_mixdiff_ccsr(
|
samples = sampler.sample_with_mixdiff_ccsr(
|
||||||
empty_text_embed, tile_size=tile_size, tile_stride=tile_stride,
|
empty_text_embed, tile_size=tile_size, tile_stride=tile_stride,
|
||||||
@@ -126,7 +118,7 @@ class CCSR_Upscale:
|
|||||||
color_fix_type=color_fix_type
|
color_fix_type=color_fix_type
|
||||||
)
|
)
|
||||||
elif sampling_method == 'ccsr_tiled_vae_gaussian_weights':
|
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")
|
print("Using gaussian weights")
|
||||||
samples = sampler.sample_with_tile_ccsr(
|
samples = sampler.sample_with_tile_ccsr(
|
||||||
empty_text_embed, tile_size=tile_size, tile_stride=tile_stride,
|
empty_text_embed, tile_size=tile_size, tile_stride=tile_stride,
|
||||||
@@ -136,7 +128,7 @@ class CCSR_Upscale:
|
|||||||
color_fix_type=color_fix_type
|
color_fix_type=color_fix_type
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.model.reset_encoder_decoder()
|
model.reset_encoder_decoder()
|
||||||
print("no tiling")
|
print("no tiling")
|
||||||
samples = sampler.sample_ccsr(
|
samples = sampler.sample_ccsr(
|
||||||
empty_text_embed, steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=img,
|
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
|
color_fix_type=color_fix_type
|
||||||
)
|
)
|
||||||
out.append(samples.squeeze(0).cpu())
|
out.append(samples.squeeze(0).cpu())
|
||||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
mm.throw_exception_if_processing_interrupted()
|
||||||
|
if B > 1:
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
print("Sampled image ", i, " out of ", batch_size)
|
print("Sampled image ", i, " out of ", B)
|
||||||
|
|
||||||
original_height, original_width = H, W
|
original_height, original_width = H, W
|
||||||
processed_height = samples.size(2)
|
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")
|
resized_back_image, = ImageScale.upscale(self, out_stacked, "lanczos", target_width, processed_height, crop="disabled")
|
||||||
|
|
||||||
if not keep_model_loaded:
|
if not keep_model_loaded:
|
||||||
self.model = None
|
model.to(offload_device)
|
||||||
comfy.model_management.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
return(resized_back_image,)
|
return(resized_back_image,)
|
||||||
|
|
||||||
class CCSR_Model_Select:
|
class CCSR_Model_Select:
|
||||||
@@ -173,15 +166,85 @@ class CCSR_Model_Select:
|
|||||||
CATEGORY = "CCSR"
|
CATEGORY = "CCSR"
|
||||||
|
|
||||||
def load_ccsr_checkpoint(self, ckpt_name):
|
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)
|
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
|
||||||
|
|
||||||
|
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,)
|
||||||
|
|
||||||
return (ckpt_path,)
|
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"CCSR_Upscale": CCSR_Upscale,
|
"CCSR_Upscale": CCSR_Upscale,
|
||||||
"CCSR_Model_Select": CCSR_Model_Select
|
"CCSR_Model_Select": CCSR_Model_Select,
|
||||||
|
"DownloadAndLoadCCSRModel": DownloadAndLoadCCSRModel
|
||||||
}
|
}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
"CCSR_Upscale": "CCSR_Upscale",
|
"CCSR_Upscale": "CCSR_Upscale",
|
||||||
"CCSR_Model_Select": "CCSR_Model_Select"
|
"CCSR_Model_Select": "CCSR_Model_Select",
|
||||||
|
"DownloadAndLoadCCSRModel": "DownloadAndLoad CCSRModel"
|
||||||
}
|
}
|
||||||
Reference in New Issue
Block a user