Files
FizzleDorf-ComfyUI-AIT/patch/0001-AIT-Compile-support.patch
T
2023-10-12 00:34:20 -04:00

776 lines
30 KiB
Diff

From 6e4aeb5f2f5261b87fe6d736f2fabd3b14d95d0b Mon Sep 17 00:00:00 2001
From: comfyanonymous <comfyanonymous@protonmail.com>
Date: Wed, 11 Oct 2023 01:45:00 -0400
Subject: [PATCH] AIT Compile support.
---
ait_convert.py | 98 +++++++++
comfy/cli_args.py | 2 +
comfy/ldm/modules/attention.py | 200 +++++++++++-------
.../modules/diffusionmodules/openaimodel.py | 110 ++++++----
comfy/ldm/modules/diffusionmodules/util.py | 85 ++++++--
comfy/ops.py | 44 ++++
6 files changed, 400 insertions(+), 139 deletions(-)
create mode 100644 ait_convert.py
diff --git a/ait_convert.py b/ait_convert.py
new file mode 100644
index 0000000..0e9fe53
--- /dev/null
+++ b/ait_convert.py
@@ -0,0 +1,98 @@
+import torch
+import comfy.cli_args
+comfy.cli_args.enable_ait = True
+
+import comfy.sd
+import comfy.model_management
+import comfy.model_detection
+import comfy.utils
+import sys
+import os
+
+sys.setrecursionlimit(10000) #Needed for SDXL
+
+
+ckpt_path = "models/checkpoints/sd_xl_base_1.0_0.9vae.safetensors"
+# ckpt_path = "models/checkpoints/cardosAnime_v20.safetensors"
+
+unet_name = "SDXL"
+
+#NOTE: to get CUDA working with cmake/ms build tools on windows copy the contents of:
+# C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.2\extras\visual_studio_integration\MSBuildExtensions
+# to:
+# C:\Program Files (x86)\Microsoft Visual Studio\2022\BuildTools\MSBuild\Microsoft\VC\v170\BuildCustomizations
+
+os.environ["AIT_USE_CMAKE_COMPILATION"] = "1"
+os.environ["AIT_ENABLE_CUDA_LTO"] = "1"
+
+fp16 = True
+sd = comfy.utils.load_torch_file(ckpt_path)
+model_config = comfy.model_detection.model_config_from_unet(sd, "model.diffusion_model.", fp16)
+model_config.unet_config['dtype'] = 'float16'
+model = model_config.get_model(sd, "model.diffusion_model.",)
+
+
+unet = model.diffusion_model
+
+unet.name_parameter_tensor()
+x = dict(unet.named_parameters())
+
+for k in x:
+ print(k, x[k])
+
+from aitemplate.compiler import compile_model
+from aitemplate.frontend import Tensor, IntVar
+from aitemplate.testing import detect_target
+
+def mark_output(y):
+ if type(y) is not tuple:
+ y = (y,)
+ for i in range(len(y)):
+ y[i]._attrs["is_output"] = True
+ y[i]._attrs["name"] = "output_%d" % (i)
+ y_shape = [d._attrs["values"] for d in y[i]._attrs["shape"]]
+ print("AIT output_{} shape: {}".format(i, y_shape))
+
+batch_size = IntVar(values=(1, 16), name="batch_size")
+height = IntVar(values=(32, 256))
+width = IntVar(values=(32, 256))
+prompt_size = IntVar(values=(77, 77*4))
+
+hidden_dim = model_config.unet_config['context_dim']
+
+latent_model_input_ait = Tensor(
+ [batch_size, height, width, 4], name="x", is_input=True, dtype='float16'
+)
+
+timesteps_ait = Tensor([batch_size], name="timesteps", is_input=True)
+text_embeddings_pt_ait = Tensor(
+ [batch_size, prompt_size, hidden_dim], name="context", is_input=True
+)
+
+if model.adm_channels == 0:
+ y_ait = None
+else:
+ y_ait = Tensor([batch_size, model.adm_channels], name="y", is_input=True)
+
+# print(unet)
+
+Y = unet(
+ latent_model_input_ait,
+ timesteps_ait,
+ text_embeddings_pt_ait,
+ y_ait
+)
+
+mark_output(Y)
+
+# print(Y)
+
+target = detect_target(
+ use_fp16_acc=True, convert_conv_to_gemm=True
+)
+
+# print(target)
+
+#print(params_ait.keys())
+model_name = "unet_1"
+compile_model(Y, target, "./optimize/{}".format(unet_name), model_name)
diff --git a/comfy/cli_args.py b/comfy/cli_args.py
index 35d4416..0d647ea 100644
--- a/comfy/cli_args.py
+++ b/comfy/cli_args.py
@@ -105,3 +105,5 @@ if args.windows_standalone_build:
if args.disable_auto_launch:
args.auto_launch = False
+
+enable_ait = False
diff --git a/comfy/ldm/modules/attention.py b/comfy/ldm/modules/attention.py
index fcae6b6..553f799 100644
--- a/comfy/ldm/modules/attention.py
+++ b/comfy/ldm/modules/attention.py
@@ -91,63 +91,6 @@ def zero_module(module):
return module
-def Normalize(in_channels, dtype=None, device=None):
- return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True, dtype=dtype, device=device)
-
-
-class SpatialSelfAttention(nn.Module):
- def __init__(self, in_channels):
- super().__init__()
- self.in_channels = in_channels
-
- self.norm = Normalize(in_channels)
- self.q = torch.nn.Conv2d(in_channels,
- in_channels,
- kernel_size=1,
- stride=1,
- padding=0)
- self.k = torch.nn.Conv2d(in_channels,
- in_channels,
- kernel_size=1,
- stride=1,
- padding=0)
- self.v = torch.nn.Conv2d(in_channels,
- in_channels,
- kernel_size=1,
- stride=1,
- padding=0)
- self.proj_out = torch.nn.Conv2d(in_channels,
- in_channels,
- kernel_size=1,
- stride=1,
- padding=0)
-
- def forward(self, x):
- h_ = x
- h_ = self.norm(h_)
- q = self.q(h_)
- k = self.k(h_)
- v = self.v(h_)
-
- # compute attention
- b,c,h,w = q.shape
- q = rearrange(q, 'b c h w -> b (h w) c')
- k = rearrange(k, 'b c h w -> b c (h w)')
- w_ = torch.einsum('bij,bjk->bik', q, k)
-
- w_ = w_ * (int(c)**(-0.5))
- w_ = torch.nn.functional.softmax(w_, dim=2)
-
- # attend to values
- v = rearrange(v, 'b c h w -> b c (h w)')
- w_ = rearrange(w_, 'b i j -> b j i')
- h_ = torch.einsum('bij,bjk->bik', v, w_)
- h_ = rearrange(h_, 'b c (h w) -> b c h w', h=h)
- h_ = self.proj_out(h_)
-
- return x+h_
-
-
class CrossAttentionBirchSan(nn.Module):
def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0., dtype=None, device=None, operations=comfy.ops):
super().__init__()
@@ -489,19 +432,119 @@ class CrossAttentionPytorch(nn.Module):
return self.to_out(out)
-if model_management.xformers_enabled():
- print("Using xformers cross attention")
- CrossAttention = MemoryEfficientCrossAttention
-elif model_management.pytorch_attention_enabled():
- print("Using pytorch cross attention")
- CrossAttention = CrossAttentionPytorch
+if comfy.ops.ENABLE_AIT:
+ from aitemplate.compiler import ops
+ from aitemplate.frontend import nn
+
+ class GEGLU(nn.Module):
+ def __init__(self, dim_in, dim_out, dtype, operations=None):
+ super().__init__()
+ self.proj = nn.Linear(dim_in, dim_out, specialization="mul", dtype=dtype)
+ self.gate = nn.Linear(dim_in, dim_out, specialization="fast_gelu", dtype=dtype)
+
+ def forward(self, x):
+ return self.proj(x, self.gate(x))
+
+
+ class FeedForward(nn.Module):
+ def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0.0, dtype=None, device=None, operations=None):
+ super().__init__()
+ inner_dim = int(dim * mult)
+ dim_out = default(dim_out, dim)
+ project_in = (
+ nn.Sequential(
+ nn.Linear(dim, inner_dim, specialization="fast_gelu", dtype=dtype),
+ )
+ if not glu
+ else GEGLU(dim, inner_dim, dtype=dtype)
+ )
+
+ self.net = nn.Sequential(
+ project_in, nn.Dropout(dropout), nn.Linear(inner_dim, dim_out)
+ )
+
+ def forward(self, x, residual=None):
+ shape = ops.size()(x)
+ x = self.net(x)
+ x = ops.reshape()(x, shape)
+ if residual is not None:
+ return x + residual
+ else:
+ return x
+
+ class CrossAttention(nn.Module):
+ def __init__(
+ self,
+ query_dim,
+ context_dim=None,
+ heads=8,
+ dim_head=64,
+ dropout=0.0,
+ dtype="float16",
+ device=None,
+ operations=None,
+ ):
+ super().__init__()
+ 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, dtype=dtype)
+ self.to_k = nn.Linear(context_dim, inner_dim, bias=False, dtype=dtype)
+ self.to_v = nn.Linear(context_dim, inner_dim, bias=False, dtype=dtype)
+ self.to_out = nn.Sequential(
+ nn.Linear(inner_dim, query_dim, dtype=dtype), nn.Dropout(dropout)
+ )
+ self.attn_op = ops.mem_eff_attention(causal=False)
+
+ def forward(self, x, context=None, value=None, mask=None):
+ nheads = self.heads
+ d = self.dim_head
+
+ q = self.to_q(x)
+ context = default(context, x)
+ k = self.to_k(context)
+ if value is not None:
+ v = self.to_v(value)
+ del value
+ else:
+ v = self.to_v(context)
+
+ bs = q.shape()[0]
+
+ q = ops.reshape()(q, [bs, -1, self.heads, self.dim_head])
+ k = ops.reshape()(k, [bs, -1, self.heads, self.dim_head])
+ v = ops.reshape()(v, [bs, -1, self.heads, self.dim_head])
+ q = ops.permute()(q, [0, 2, 1, 3])
+ k = ops.permute()(k, [0, 2, 1, 3])
+ v = ops.permute()(v, [0, 2, 1, 3])
+
+ out = self.attn_op(
+ (ops.reshape()(q, [bs, nheads, -1, d])),
+ (ops.reshape()(k, [bs, nheads, -1, d])),
+ (ops.reshape()(v, [bs, nheads, -1, d])),
+ )
+ out = ops.reshape()(out, [bs, -1, nheads * d])
+ proj = self.to_out(out)
+ proj = ops.reshape()(proj, [bs, -1, nheads * d])
+ return proj
+
else:
- if args.use_split_cross_attention:
- print("Using split optimization for cross attention")
- CrossAttention = CrossAttentionDoggettx
+ if model_management.xformers_enabled():
+ print("Using xformers cross attention")
+ CrossAttention = MemoryEfficientCrossAttention
+ elif model_management.pytorch_attention_enabled():
+ print("Using pytorch cross attention")
+ CrossAttention = CrossAttentionPytorch
else:
- print("Using sub quadratic optimization for cross attention, if you have memory or speed issues try using: --use-split-cross-attention")
- CrossAttention = CrossAttentionBirchSan
+ if args.use_split_cross_attention:
+ print("Using split optimization for cross attention")
+ CrossAttention = CrossAttentionDoggettx
+ else:
+ print("Using sub quadratic optimization for cross attention, if you have memory or speed issues try using: --use-split-cross-attention")
+ CrossAttention = CrossAttentionBirchSan
class BasicTransformerBlock(nn.Module):
@@ -654,7 +697,7 @@ class SpatialTransformer(nn.Module):
context_dim = [context_dim] * depth
self.in_channels = in_channels
inner_dim = n_heads * d_head
- self.norm = Normalize(in_channels, dtype=dtype, device=device)
+ self.norm = nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True, dtype=dtype, device=device)
if not use_linear:
self.proj_in = operations.Conv2d(in_channels,
inner_dim,
@@ -682,12 +725,20 @@ class SpatialTransformer(nn.Module):
# note: if no context is given, cross-attention defaults to self-attention
if not isinstance(context, list):
context = [context] * len(self.transformer_blocks)
- b, c, h, w = x.shape
+ if comfy.ops.ENABLE_AIT:
+ b, h, w, c = x.shape()
+ else:
+ b, c, h, w = x.shape
x_in = x
x = self.norm(x)
if not self.use_linear:
x = self.proj_in(x)
- x = rearrange(x, 'b c h w -> b (h w) c').contiguous()
+
+ if comfy.ops.ENABLE_AIT:
+ x = ops.reshape()(x, [b, -1, c])
+ else:
+ x = x.movedim(1, -1).reshape((b, -1, c)).contiguous()
+
if self.use_linear:
x = self.proj_in(x)
for i, block in enumerate(self.transformer_blocks):
@@ -695,7 +746,12 @@ class SpatialTransformer(nn.Module):
x = block(x, context=context[i], transformer_options=transformer_options)
if self.use_linear:
x = self.proj_out(x)
- x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous()
+
+ if comfy.ops.ENABLE_AIT:
+ x = ops.reshape()(x, [b, h, w, c])
+ else:
+ x = x.movedim(-1, 1).reshape((b, c, h, w)).contiguous()
+
if not self.use_linear:
x = self.proj_out(x)
return x + x_in
diff --git a/comfy/ldm/modules/diffusionmodules/openaimodel.py b/comfy/ldm/modules/diffusionmodules/openaimodel.py
index b42637c..1b5f0dd 100644
--- a/comfy/ldm/modules/diffusionmodules/openaimodel.py
+++ b/comfy/ldm/modules/diffusionmodules/openaimodel.py
@@ -17,6 +17,13 @@ from ..attention import SpatialTransformer
from comfy.ldm.util import exists
import comfy.ops
+silu_op = F.silu
+
+if comfy.ops.ENABLE_AIT:
+ from aitemplate.frontend import nn, Tensor
+ from aitemplate.compiler import ops
+ silu_op = ops.silu
+
class TimestepBlock(nn.Module):
"""
Any module where forward() takes timestep embeddings as a second argument.
@@ -80,19 +87,25 @@ class Upsample(nn.Module):
self.conv = operations.conv_nd(dims, self.channels, self.out_channels, 3, padding=padding, dtype=dtype, device=device)
def forward(self, x, output_shape=None):
- assert x.shape[1] == self.channels
- if self.dims == 3:
- shape = [x.shape[2], x.shape[3] * 2, x.shape[4] * 2]
- if output_shape is not None:
- shape[1] = output_shape[3]
- shape[2] = output_shape[4]
+ if comfy.ops.ENABLE_AIT:
+ if output_shape is None:
+ x = nn.Upsampling2d(scale_factor=2.0, mode="nearest")(x)
+ else:
+ b, h, w, c = output_shape()
+ b_, h_, w_, c_ = x.shape()
+ x = nn.Upsampling2d(scale_factor=2.0, mode="nearest")(x, Tensor([b_, h, w, c_]))
else:
- shape = [x.shape[2] * 2, x.shape[3] * 2]
- if output_shape is not None:
- shape[0] = output_shape[2]
- shape[1] = output_shape[3]
-
- x = F.interpolate(x, size=shape, mode="nearest")
+ if self.dims == 3:
+ shape = [x.shape[2], x.shape[3] * 2, x.shape[4] * 2]
+ if output_shape is not None:
+ shape[1] = output_shape[3]
+ shape[2] = output_shape[4]
+ else:
+ shape = [x.shape[2] * 2, x.shape[3] * 2]
+ if output_shape is not None:
+ shape[0] = output_shape[2]
+ shape[1] = output_shape[3]
+ x = F.interpolate(x, size=shape, mode="nearest")
if self.use_conv:
x = self.conv(x)
return x
@@ -122,7 +135,6 @@ class Downsample(nn.Module):
self.op = avg_pool_nd(dims, kernel_size=stride, stride=stride)
def forward(self, x):
- assert x.shape[1] == self.channels
return self.op(x)
@@ -168,8 +180,8 @@ class ResBlock(TimestepBlock):
self.use_scale_shift_norm = use_scale_shift_norm
self.in_layers = nn.Sequential(
- nn.GroupNorm(32, channels, dtype=dtype, device=device),
- nn.SiLU(),
+ comfy.ops.GroupNorm(32, channels, dtype=dtype, device=device, use_swish=True),
+ nn.Identity(),
operations.conv_nd(dims, channels, self.out_channels, 3, padding=1, dtype=dtype, device=device),
)
@@ -185,19 +197,18 @@ class ResBlock(TimestepBlock):
self.h_upd = self.x_upd = nn.Identity()
self.emb_layers = nn.Sequential(
- nn.SiLU(),
+ nn.Identity(),
operations.Linear(
emb_channels,
2 * self.out_channels if use_scale_shift_norm else self.out_channels, dtype=dtype, device=device
),
)
+
self.out_layers = nn.Sequential(
- nn.GroupNorm(32, self.out_channels, dtype=dtype, device=device),
- nn.SiLU(),
+ comfy.ops.GroupNorm(32, self.out_channels, dtype=dtype, device=device, use_swish=True),
+ nn.Identity(),
nn.Dropout(p=dropout),
- zero_module(
- operations.conv_nd(dims, self.out_channels, self.out_channels, 3, padding=1, dtype=dtype, device=device)
- ),
+ operations.conv_nd(dims, self.out_channels, self.out_channels, 3, padding=1, dtype=dtype, device=device),
)
if self.out_channels == channels:
@@ -230,9 +241,12 @@ class ResBlock(TimestepBlock):
h = in_conv(h)
else:
h = self.in_layers(x)
- emb_out = self.emb_layers(emb).type(h.dtype)
- while len(emb_out.shape) < len(h.shape):
- emb_out = emb_out[..., None]
+ emb_out = self.emb_layers(silu_op(emb))
+
+ if comfy.ops.ENABLE_AIT:
+ emb_out = ops.reshape()(emb_out, [emb_out.shape()[0], 1, 1, emb_out.shape()[1]])
+ else:
+ emb_out = emb_out.reshape((emb_out.shape[0], emb_out.shape[1], 1, 1))
if self.use_scale_shift_norm:
out_norm, out_rest = self.out_layers[0], self.out_layers[1:]
scale, shift = th.chunk(emb_out, 2, dim=1)
@@ -316,6 +330,7 @@ class UNetModel(nn.Module):
adm_in_channels=None,
transformer_depth_middle=None,
device=None,
+ dtype=None,
operations=comfy.ops,
):
super().__init__()
@@ -370,19 +385,19 @@ class UNetModel(nn.Module):
self.conv_resample = conv_resample
self.num_classes = num_classes
self.use_checkpoint = use_checkpoint
- self.dtype = th.float16 if use_fp16 else th.float32
- self.dtype = th.bfloat16 if use_bf16 else self.dtype
+ if dtype is not None:
+ self.dtype = dtype
+ else:
+ self.dtype = th.float16 if use_fp16 else th.float32
+ self.dtype = th.bfloat16 if use_bf16 else self.dtype
self.num_heads = num_heads
self.num_head_channels = num_head_channels
self.num_heads_upsample = num_heads_upsample
self.predict_codebook_ids = n_embed is not None
time_embed_dim = model_channels * 4
- self.time_embed = nn.Sequential(
- operations.Linear(model_channels, time_embed_dim, dtype=self.dtype, device=device),
- nn.SiLU(),
- operations.Linear(time_embed_dim, time_embed_dim, dtype=self.dtype, device=device),
- )
+ self.time_embed = comfy.ops.time_embed(model_channels, time_embed_dim, dtype=self.dtype, device=device)
+
if self.num_classes is not None:
if isinstance(self.num_classes, int):
@@ -393,11 +408,7 @@ class UNetModel(nn.Module):
elif self.num_classes == "sequential":
assert adm_in_channels is not None
self.label_emb = nn.Sequential(
- nn.Sequential(
- operations.Linear(adm_in_channels, time_embed_dim, dtype=self.dtype, device=device),
- nn.SiLU(),
- operations.Linear(time_embed_dim, time_embed_dim, dtype=self.dtype, device=device),
- )
+ comfy.ops.time_embed(adm_in_channels, time_embed_dim, dtype=self.dtype, device=device)
)
else:
raise ValueError()
@@ -586,13 +597,14 @@ class UNetModel(nn.Module):
self._feature_size += ch
self.out = nn.Sequential(
- nn.GroupNorm(32, ch, dtype=self.dtype, device=device),
- nn.SiLU(),
- zero_module(operations.conv_nd(dims, model_channels, out_channels, 3, padding=1, dtype=self.dtype, device=device)),
+ comfy.ops.GroupNorm(32, ch, dtype=self.dtype, device=device, use_swish=True),
+ nn.Identity(),
+ operations.conv_nd(dims, model_channels, out_channels, 3, padding=1, dtype=self.dtype, device=device),
)
+
if self.predict_codebook_ids:
self.id_predictor = nn.Sequential(
- nn.GroupNorm(32, ch, dtype=self.dtype, device=device),
+ comfy.ops.GroupNorm(32, ch, dtype=self.dtype, device=device),
operations.conv_nd(dims, model_channels, n_embed, 1, dtype=self.dtype, device=device),
#nn.LogSoftmax(dim=1) # change to cross_entropy and produce non-normalized logits
)
@@ -606,7 +618,10 @@ class UNetModel(nn.Module):
:param y: an [N] Tensor of labels, if class-conditional.
:return: an [N x C x ...] Tensor of outputs.
"""
- transformer_options["original_shape"] = list(x.shape)
+ if comfy.ops.ENABLE_AIT:
+ transformer_options["original_shape"] = list(x.shape())
+ else:
+ transformer_options["original_shape"] = list(x.shape)
transformer_options["current_index"] = 0
transformer_patches = transformer_options.get("patches", {})
@@ -614,14 +629,13 @@ class UNetModel(nn.Module):
self.num_classes is not None
), "must specify y if and only if the model is class-conditional"
hs = []
- t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(self.dtype)
+ t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False, dtype=x.dtype)
emb = self.time_embed(t_emb)
if self.num_classes is not None:
- assert y.shape[0] == x.shape[0]
emb = emb + self.label_emb(y)
- h = x.type(self.dtype)
+ h = x
for id, module in enumerate(self.input_blocks):
transformer_options["block"] = ("input", id)
h = forward_timestep_embed(module, h, emb, context, transformer_options)
@@ -650,14 +664,18 @@ class UNetModel(nn.Module):
for p in patch:
h, hsp = p(h, hsp, transformer_options)
- h = th.cat([h, hsp], dim=1)
+ if comfy.ops.ENABLE_AIT:
+ h = ops.concatenate()([h, hsp], dim=-1)
+ else:
+ h = th.cat([h, hsp], dim=1)
+
del hsp
if len(hs) > 0:
output_shape = hs[-1].shape
else:
output_shape = None
h = forward_timestep_embed(module, h, emb, context, transformer_options, output_shape)
- h = h.type(x.dtype)
+
if self.predict_codebook_ids:
return self.id_predictor(h)
else:
diff --git a/comfy/ldm/modules/diffusionmodules/util.py b/comfy/ldm/modules/diffusionmodules/util.py
index d890c80..ae42f61 100644
--- a/comfy/ldm/modules/diffusionmodules/util.py
+++ b/comfy/ldm/modules/diffusionmodules/util.py
@@ -157,28 +157,71 @@ class CheckpointFunction(torch.autograd.Function):
del output_tensors
return (None, None) + input_grads
+if comfy.ops.ENABLE_AIT:
+ from aitemplate.compiler import ops
+ from aitemplate.frontend import nn, Tensor
+ def timestep_embedding(timesteps, embedding_dim, max_period=10000, repeat_only=False, dtype=None):
+ scale = 1.0
+ flip_sin_to_cos = True
+ downscale_freq_shift = 0
+ """
+ This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings.
+
+ :param timesteps: a 1-D Tensor of N indices, one per batch element.
+ These may be fractional.
+ :param embedding_dim: the dimension of the output. :param max_period: controls the minimum frequency of the
+ embeddings. :return: an [N x dim] Tensor of positional embeddings.
+ """
+ assert timesteps._rank() == 1, "Timesteps should be a 1d-array"
+
+ half_dim = embedding_dim // 2
+
+ exponent = (-math.log(max_period)) * Tensor(
+ shape=[half_dim], dtype="float16", name="arange"
+ )
-def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False):
- """
- Create sinusoidal timestep embeddings.
- :param timesteps: a 1-D Tensor of N indices, one per batch element.
- These may be fractional.
- :param dim: the dimension of the output.
- :param max_period: controls the minimum frequency of the embeddings.
- :return: an [N x dim] Tensor of positional embeddings.
- """
- if not repeat_only:
- half = dim // 2
- freqs = torch.exp(
- -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
- ).to(device=timesteps.device)
- args = timesteps[:, None].float() * freqs[None]
- embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
- if dim % 2:
- embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
- else:
- embedding = repeat(timesteps, 'b -> b d', d=dim)
- return embedding
+ exponent = exponent * (1.0 / (half_dim - downscale_freq_shift))
+
+ emb = ops.exp(exponent)
+ emb = ops.reshape()(timesteps, [-1, 1]) * ops.reshape()(emb, [1, -1])
+
+ # scale embeddings
+ emb = scale * emb
+
+ # concat sine and cosine embeddings
+ if flip_sin_to_cos:
+ emb = ops.concatenate()(
+ [ops.cos(emb), ops.sin(emb)],
+ dim=-1,
+ )
+ else:
+ emb = ops.concatenate()(
+ [ops.sin(emb), ops.cos(emb)],
+ dim=-1,
+ )
+ return emb
+else:
+ def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False, dtype=torch.float32):
+ """
+ Create sinusoidal timestep embeddings.
+ :param timesteps: a 1-D Tensor of N indices, one per batch element.
+ These may be fractional.
+ :param dim: the dimension of the output.
+ :param max_period: controls the minimum frequency of the embeddings.
+ :return: an [N x dim] Tensor of positional embeddings.
+ """
+ if not repeat_only:
+ half = dim // 2
+ freqs = torch.exp(
+ -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
+ ).to(device=timesteps.device)
+ args = timesteps[:, None].float() * freqs[None]
+ embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
+ if dim % 2:
+ embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
+ else:
+ embedding = repeat(timesteps, 'b -> b d', d=dim)
+ return embedding.to(dtype)
def zero_module(module):
diff --git a/comfy/ops.py b/comfy/ops.py
index 610d545..6d8b077 100644
--- a/comfy/ops.py
+++ b/comfy/ops.py
@@ -1,6 +1,10 @@
import torch
from contextlib import contextmanager
+import comfy.cli_args
+
+ENABLE_AIT = comfy.cli_args.enable_ait
+
class Linear(torch.nn.Module):
def __init__(self, in_features: int, out_features: int, bias: bool = True,
device=None, dtype=None) -> None:
@@ -21,12 +25,52 @@ class Conv2d(torch.nn.Conv2d):
def reset_parameters(self):
return None
+class GroupNorm(torch.nn.GroupNorm):
+ def __init__(self, num_groups: int, num_channels: int, eps: float = 1e-5, affine: bool = True, device=None, dtype=None, use_swish=False):
+ super().__init__(num_groups, num_channels, eps, affine, device, dtype)
+ if use_swish:
+ self.swish = torch.nn.SiLU()
+ else:
+ self.swish = torch.nn.Identity()
+
+ def forward(self, input):
+ return self.swish(super().forward(input))
+
+def time_embed(model_channels, embed_dim, dtype, device):
+ return torch.nn.Sequential(
+ Linear(model_channels, embed_dim, dtype=dtype, device=device),
+ torch.nn.SiLU(),
+ Linear(embed_dim, embed_dim, dtype=dtype, device=device),
+ )
+
+if ENABLE_AIT:
+ from aitemplate.frontend import nn, Tensor
+ class Linear(nn.Linear):
+ def __init__(self, in_features: int, out_features: int, bias: bool = True,
+ device=None, dtype=None) -> None:
+ return super().__init__(in_features, out_features, bias, dtype=dtype)
+
+ def Conv2d(in_channels, out_channels, kernel_size, stride=1, padding=0, dilation=1, groups=1, bias=True, padding_mode='zeros', device=None, dtype=None):
+ if bias == True and padding_mode == 'zeros':
+ return nn.Conv2dBias(in_channels, out_channels, kernel_size, stride=stride, padding=padding, dilation=dilation, groups=groups, dtype=dtype)
+ else:
+ print("error UNIMPLEMENTED")
+
+ def time_embed(model_channels, embed_dim, dtype, device):
+ return nn.Sequential(
+ nn.Linear(model_channels, embed_dim, specialization="swish", dtype=dtype),
+ nn.Identity(),
+ nn.Linear(embed_dim, embed_dim, dtype=dtype),
+ )
+ GroupNorm = nn.GroupNorm
+
def conv_nd(dims, *args, **kwargs):
if dims == 2:
return Conv2d(*args, **kwargs)
else:
raise ValueError(f"unsupported dimensions: {dims}")
+
@contextmanager
def use_comfy_ops(device=None, dtype=None): # Kind of an ugly hack but I can't think of a better way
old_torch_nn_linear = torch.nn.Linear
--
2.37.1.windows.1