PixArt controlnet test
This commit is contained in:
@@ -38,6 +38,19 @@ pixart_conf = {
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pixart_conf.update({ # controlnet models
|
||||||
|
"ControlPixArtHalf": {
|
||||||
|
"target": "ControlPixArtHalf",
|
||||||
|
"unet_config": pixart_conf["PixArt_XL_2"]["unet_config"],
|
||||||
|
"sampling_settings": pixart_conf["PixArt_XL_2"]["sampling_settings"],
|
||||||
|
},
|
||||||
|
"ControlPixArtMSHalf": {
|
||||||
|
"target": "ControlPixArtMSHalf",
|
||||||
|
"unet_config": pixart_conf["PixArtMS_XL_2"]["unet_config"],
|
||||||
|
"sampling_settings": pixart_conf["PixArtMS_XL_2"]["sampling_settings"],
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
pixart_res = {
|
pixart_res = {
|
||||||
"PixArtMS_XL_2": { # models/PixArtMS 1024x1024
|
"PixArtMS_XL_2": { # models/PixArtMS 1024x1024
|
||||||
'0.25': [512, 2048], '0.26': [512, 1984], '0.27': [512, 1920], '0.28': [512, 1856],
|
'0.25': [512, 2048], '0.26': [512, 1984], '0.27': [512, 1920], '0.28': [512, 1856],
|
||||||
|
|||||||
+33
-1
@@ -3,6 +3,7 @@ import comfy.latent_formats
|
|||||||
import comfy.model_patcher
|
import comfy.model_patcher
|
||||||
import comfy.model_base
|
import comfy.model_base
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
|
import comfy.conds
|
||||||
import torch
|
import torch
|
||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
from .diffusers_convert import convert_state_dict
|
from .diffusers_convert import convert_state_dict
|
||||||
@@ -23,6 +24,27 @@ class EXM_PixArt(comfy.supported_models_base.BASE):
|
|||||||
def model_type(self, state_dict, prefix=""):
|
def model_type(self, state_dict, prefix=""):
|
||||||
return comfy.model_base.ModelType.EPS
|
return comfy.model_base.ModelType.EPS
|
||||||
|
|
||||||
|
class EXM_PixArt_Model(comfy.model_base.BaseModel):
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
|
||||||
|
def extra_conds(self, **kwargs):
|
||||||
|
out = super().extra_conds(**kwargs)
|
||||||
|
|
||||||
|
img_hw = kwargs.get("img_hw", None)
|
||||||
|
if img_hw is not None:
|
||||||
|
out["img_hw"] = comfy.conds.CONDRegular(torch.tensor(img_hw))
|
||||||
|
|
||||||
|
aspect_ratio = kwargs.get("aspect_ratio", None)
|
||||||
|
if aspect_ratio is not None:
|
||||||
|
out["aspect_ratio"] = comfy.conds.CONDRegular(torch.tensor(aspect_ratio))
|
||||||
|
|
||||||
|
cn_hint = kwargs.get("cn_hint", None)
|
||||||
|
if cn_hint is not None:
|
||||||
|
out["cn_hint"] = comfy.conds.CONDRegular(cn_hint)
|
||||||
|
|
||||||
|
return out
|
||||||
|
|
||||||
def load_pixart(model_path, model_conf):
|
def load_pixart(model_path, model_conf):
|
||||||
state_dict = comfy.utils.load_torch_file(model_path)
|
state_dict = comfy.utils.load_torch_file(model_path)
|
||||||
state_dict = state_dict.get("model", state_dict)
|
state_dict = state_dict.get("model", state_dict)
|
||||||
@@ -48,7 +70,7 @@ def load_pixart(model_path, model_conf):
|
|||||||
unet_dtype = manual_cast_dtype
|
unet_dtype = manual_cast_dtype
|
||||||
|
|
||||||
model_conf = EXM_PixArt(model_conf) # convert to object
|
model_conf = EXM_PixArt(model_conf) # convert to object
|
||||||
model = comfy.model_base.BaseModel(
|
model = EXM_PixArt_Model( # same as comfy.model_base.BaseModel
|
||||||
model_conf,
|
model_conf,
|
||||||
model_type=comfy.model_base.ModelType.EPS,
|
model_type=comfy.model_base.ModelType.EPS,
|
||||||
device=model_management.get_torch_device()
|
device=model_management.get_torch_device()
|
||||||
@@ -60,6 +82,16 @@ def load_pixart(model_path, model_conf):
|
|||||||
elif model_conf.model_target == "PixArt":
|
elif model_conf.model_target == "PixArt":
|
||||||
from .models.PixArt import PixArt
|
from .models.PixArt import PixArt
|
||||||
model.diffusion_model = PixArt(**model_conf.unet_config)
|
model.diffusion_model = PixArt(**model_conf.unet_config)
|
||||||
|
elif model_conf.model_target == "ControlPixArtMSHalf":
|
||||||
|
from .models.PixArtMS import PixArtMS
|
||||||
|
from .models.pixart_controlnet import ControlPixArtMSHalf
|
||||||
|
model.diffusion_model = PixArtMS(**model_conf.unet_config)
|
||||||
|
model.diffusion_model = ControlPixArtMSHalf(model.diffusion_model)
|
||||||
|
elif model_conf.model_target == "ControlPixArtHalf":
|
||||||
|
from .models.PixArt import PixArt
|
||||||
|
from .models.pixart_controlnet import ControlPixArtHalf
|
||||||
|
model.diffusion_model = PixArt(**model_conf.unet_config)
|
||||||
|
model.diffusion_model = ControlPixArtHalf(model.diffusion_model)
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError(f"Unknown model target '{model_conf.model_target}'")
|
raise NotImplementedError(f"Unknown model target '{model_conf.model_target}'")
|
||||||
|
|
||||||
|
|||||||
@@ -28,6 +28,7 @@ class PixArtBlock(nn.Module):
|
|||||||
|
|
||||||
def __init__(self, hidden_size, num_heads, mlp_ratio=4.0, drop_path=0., window_size=0, input_size=None, use_rel_pos=False, **block_kwargs):
|
def __init__(self, hidden_size, num_heads, mlp_ratio=4.0, drop_path=0., window_size=0, input_size=None, use_rel_pos=False, **block_kwargs):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
self.hidden_size = hidden_size
|
||||||
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||||
self.attn = WindowAttention(hidden_size, num_heads=num_heads, qkv_bias=True,
|
self.attn = WindowAttention(hidden_size, num_heads=num_heads, qkv_bias=True,
|
||||||
input_size=input_size if window_size == 0 else (window_size, window_size),
|
input_size=input_size if window_size == 0 else (window_size, window_size),
|
||||||
|
|||||||
+15
-11
@@ -182,30 +182,34 @@ class PixArtMS(PixArt):
|
|||||||
x = self.unpatchify(x) # (N, out_channels, H, W)
|
x = self.unpatchify(x) # (N, out_channels, H, W)
|
||||||
return x
|
return x
|
||||||
|
|
||||||
def forward(self, x, timesteps, context, y=None, **kwargs):
|
def forward(self, x, timesteps, context, img_hw=None, aspect_ratio=None, **kwargs):
|
||||||
"""
|
"""
|
||||||
Forward pass that adapts comfy input to original forward function
|
Forward pass that adapts comfy input to original forward function
|
||||||
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
||||||
timesteps: (N,) tensor of diffusion timesteps
|
timesteps: (N,) tensor of diffusion timesteps
|
||||||
context: (N, 1, 120, C) conditioning
|
context: (N, 1, 120, C) conditioning
|
||||||
y: extra conditioning.
|
img_hw: height|width conditioning
|
||||||
|
aspect_ratio: aspect ratio conditioning
|
||||||
"""
|
"""
|
||||||
## aspect ratio based on the latent image shape.
|
## size/ar from cond with fallback based on the latent image shape.
|
||||||
# Ideally, these should only be used as a fallback with the real ones
|
|
||||||
# passed in `y` to allow different values to be used for cont/uncond.
|
|
||||||
bs = x.shape[0]
|
bs = x.shape[0]
|
||||||
data_info = {
|
data_info = {}
|
||||||
"img_hw" : torch.tensor(
|
if img_hw is None:
|
||||||
|
data_info["img_hw"] = torch.tensor(
|
||||||
[[x.shape[2]*8, x.shape[3]*8]],
|
[[x.shape[2]*8, x.shape[3]*8]],
|
||||||
dtype=self.dtype,
|
dtype=self.dtype,
|
||||||
device=x.device
|
device=x.device
|
||||||
).repeat(bs, 1),
|
).repeat(bs, 1)
|
||||||
"aspect_ratio" : torch.tensor(
|
else:
|
||||||
|
data_info["img_hw"] = img_hw.to(x.dtype).to(x.device)
|
||||||
|
if aspect_ratio is None or True:
|
||||||
|
data_info["aspect_ratio"] = torch.tensor(
|
||||||
[[x.shape[2]/x.shape[3]]],
|
[[x.shape[2]/x.shape[3]]],
|
||||||
dtype=self.dtype,
|
dtype=self.dtype,
|
||||||
device=x.device
|
device=x.device
|
||||||
).repeat(bs, 1),
|
).repeat(bs, 1)
|
||||||
}
|
else:
|
||||||
|
data_info["aspect_ratio"] = aspect_ratio.to(x.dtype).to(x.device)
|
||||||
|
|
||||||
## Still accepts the input w/o that dim but returns garbage
|
## Still accepts the input w/o that dim but returns garbage
|
||||||
if len(context.shape) == 3:
|
if len(context.shape) == 3:
|
||||||
|
|||||||
@@ -0,0 +1,312 @@
|
|||||||
|
import re
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from copy import deepcopy
|
||||||
|
from torch import Tensor
|
||||||
|
from torch.nn import Module, Linear, init
|
||||||
|
from typing import Any, Mapping
|
||||||
|
|
||||||
|
from .PixArt import PixArt, get_2d_sincos_pos_embed
|
||||||
|
from .PixArtMS import PixArtMSBlock, PixArtMS
|
||||||
|
from .utils import auto_grad_checkpoint
|
||||||
|
|
||||||
|
# The implementation of ControlNet-Half architrecture
|
||||||
|
# https://github.com/lllyasviel/ControlNet/discussions/188
|
||||||
|
class ControlT2IDitBlockHalf(Module):
|
||||||
|
def __init__(self, base_block: PixArtMSBlock, block_index: 0) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.copied_block = deepcopy(base_block)
|
||||||
|
self.block_index = block_index
|
||||||
|
|
||||||
|
for p in self.copied_block.parameters():
|
||||||
|
p.requires_grad_(True)
|
||||||
|
|
||||||
|
self.copied_block.load_state_dict(base_block.state_dict())
|
||||||
|
self.copied_block.train()
|
||||||
|
|
||||||
|
self.hidden_size = hidden_size = base_block.hidden_size
|
||||||
|
if self.block_index == 0:
|
||||||
|
self.before_proj = Linear(hidden_size, hidden_size)
|
||||||
|
init.zeros_(self.before_proj.weight)
|
||||||
|
init.zeros_(self.before_proj.bias)
|
||||||
|
self.after_proj = Linear(hidden_size, hidden_size)
|
||||||
|
init.zeros_(self.after_proj.weight)
|
||||||
|
init.zeros_(self.after_proj.bias)
|
||||||
|
|
||||||
|
def forward(self, x, y, t, mask=None, c=None):
|
||||||
|
|
||||||
|
if self.block_index == 0:
|
||||||
|
# the first block
|
||||||
|
c = self.before_proj(c)
|
||||||
|
c = self.copied_block(x + c, y, t, mask)
|
||||||
|
c_skip = self.after_proj(c)
|
||||||
|
else:
|
||||||
|
# load from previous c and produce the c for skip connection
|
||||||
|
c = self.copied_block(c, y, t, mask)
|
||||||
|
c_skip = self.after_proj(c)
|
||||||
|
|
||||||
|
return c, c_skip
|
||||||
|
|
||||||
|
|
||||||
|
# The implementation of ControlPixArtHalf net
|
||||||
|
class ControlPixArtHalf(Module):
|
||||||
|
# only support single res model
|
||||||
|
def __init__(self, base_model: PixArt, copy_blocks_num: int = 13) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.dtype = torch.get_default_dtype()
|
||||||
|
self.base_model = base_model.eval()
|
||||||
|
self.controlnet = []
|
||||||
|
self.copy_blocks_num = copy_blocks_num
|
||||||
|
self.total_blocks_num = len(base_model.blocks)
|
||||||
|
for p in self.base_model.parameters():
|
||||||
|
p.requires_grad_(False)
|
||||||
|
|
||||||
|
# Copy first copy_blocks_num block
|
||||||
|
for i in range(copy_blocks_num):
|
||||||
|
self.controlnet.append(ControlT2IDitBlockHalf(base_model.blocks[i], i))
|
||||||
|
self.controlnet = nn.ModuleList(self.controlnet)
|
||||||
|
|
||||||
|
def __getattr__(self, name: str) -> Tensor or Module:
|
||||||
|
if name in ['forward', 'forward_with_dpmsolver', 'forward_with_cfg', 'forward_c', 'load_state_dict']:
|
||||||
|
return self.__dict__[name]
|
||||||
|
elif name in ['base_model', 'controlnet']:
|
||||||
|
return super().__getattr__(name)
|
||||||
|
else:
|
||||||
|
return getattr(self.base_model, name)
|
||||||
|
|
||||||
|
def forward_c(self, c):
|
||||||
|
self.h, self.w = c.shape[-2]//self.patch_size, c.shape[-1]//self.patch_size
|
||||||
|
pos_embed = torch.from_numpy(get_2d_sincos_pos_embed(self.pos_embed.shape[-1], (self.h, self.w), lewei_scale=self.lewei_scale, base_size=self.base_size)).unsqueeze(0).to(c.device).to(self.dtype)
|
||||||
|
return self.x_embedder(c) + pos_embed if c is not None else c
|
||||||
|
|
||||||
|
# def forward(self, x, t, c, **kwargs):
|
||||||
|
# return self.base_model(x, t, c=self.forward_c(c), **kwargs)
|
||||||
|
def forward_raw(self, x, timestep, y, mask=None, data_info=None, c=None, **kwargs):
|
||||||
|
# modify the original PixArtMS forward function
|
||||||
|
if c is not None:
|
||||||
|
c = c.to(self.dtype)
|
||||||
|
c = self.forward_c(c)
|
||||||
|
"""
|
||||||
|
Forward pass of PixArt.
|
||||||
|
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
||||||
|
t: (N,) tensor of diffusion timesteps
|
||||||
|
y: (N, 1, 120, C) tensor of class labels
|
||||||
|
"""
|
||||||
|
x = x.to(self.dtype)
|
||||||
|
timestep = timestep.to(self.dtype)
|
||||||
|
y = y.to(self.dtype)
|
||||||
|
pos_embed = self.pos_embed.to(self.dtype)
|
||||||
|
self.h, self.w = x.shape[-2]//self.patch_size, x.shape[-1]//self.patch_size
|
||||||
|
x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2
|
||||||
|
t = self.t_embedder(timestep.to(x.dtype)) # (N, D)
|
||||||
|
t0 = self.t_block(t)
|
||||||
|
y = self.y_embedder(y, self.training) # (N, 1, L, D)
|
||||||
|
if mask is not None:
|
||||||
|
if mask.shape[0] != y.shape[0]:
|
||||||
|
mask = mask.repeat(y.shape[0] // mask.shape[0], 1)
|
||||||
|
mask = mask.squeeze(1).squeeze(1)
|
||||||
|
y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1])
|
||||||
|
y_lens = mask.sum(dim=1).tolist()
|
||||||
|
else:
|
||||||
|
y_lens = [y.shape[2]] * y.shape[0]
|
||||||
|
y = y.squeeze(1).view(1, -1, x.shape[-1])
|
||||||
|
|
||||||
|
# define the first layer
|
||||||
|
x = auto_grad_checkpoint(self.base_model.blocks[0], x, y, t0, y_lens, **kwargs) # (N, T, D) #support grad checkpoint
|
||||||
|
|
||||||
|
if c is not None:
|
||||||
|
# update c
|
||||||
|
for index in range(1, self.copy_blocks_num + 1):
|
||||||
|
c, c_skip = auto_grad_checkpoint(self.controlnet[index - 1], x, y, t0, y_lens, c, **kwargs)
|
||||||
|
x = auto_grad_checkpoint(self.base_model.blocks[index], x + c_skip, y, t0, y_lens, **kwargs)
|
||||||
|
|
||||||
|
# update x
|
||||||
|
for index in range(self.copy_blocks_num + 1, self.total_blocks_num):
|
||||||
|
x = auto_grad_checkpoint(self.base_model.blocks[index], x, y, t0, y_lens, **kwargs)
|
||||||
|
else:
|
||||||
|
for index in range(1, self.total_blocks_num):
|
||||||
|
x = auto_grad_checkpoint(self.base_model.blocks[index], x, y, t0, y_lens, **kwargs)
|
||||||
|
|
||||||
|
x = self.final_layer(x, t) # (N, T, patch_size ** 2 * out_channels)
|
||||||
|
x = self.unpatchify(x) # (N, out_channels, H, W)
|
||||||
|
return x
|
||||||
|
|
||||||
|
def forward(self, x, timesteps, context, cn_hint=None, **kwargs):
|
||||||
|
"""
|
||||||
|
Forward pass that adapts comfy input to original forward function
|
||||||
|
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
||||||
|
timesteps: (N,) tensor of diffusion timesteps
|
||||||
|
context: (N, 1, 120, C) conditioning
|
||||||
|
cn_hint: controlnet hint
|
||||||
|
"""
|
||||||
|
## Still accepts the input w/o that dim but returns garbage
|
||||||
|
if len(context.shape) == 3:
|
||||||
|
context = context.unsqueeze(1)
|
||||||
|
|
||||||
|
## run original forward pass
|
||||||
|
out = self.forward_raw(
|
||||||
|
x = x.to(self.dtype),
|
||||||
|
timestep = timesteps.to(self.dtype),
|
||||||
|
y = context.to(self.dtype),
|
||||||
|
c = cn_hint,
|
||||||
|
)
|
||||||
|
|
||||||
|
## only return EPS
|
||||||
|
out = out.to(torch.float)
|
||||||
|
eps, rest = out[:, :self.in_channels], out[:, self.in_channels:]
|
||||||
|
return eps
|
||||||
|
|
||||||
|
def forward_with_dpmsolver(self, x, t, y, data_info, c, **kwargs):
|
||||||
|
model_out = self.forward_raw(x, t, y, data_info=data_info, c=c, **kwargs)
|
||||||
|
return model_out.chunk(2, dim=1)[0]
|
||||||
|
|
||||||
|
# def forward_with_dpmsolver(self, x, t, y, data_info, c, **kwargs):
|
||||||
|
# return self.base_model.forward_with_dpmsolver(x, t, y, data_info=data_info, c=self.forward_c(c), **kwargs)
|
||||||
|
|
||||||
|
def forward_with_cfg(self, x, t, y, cfg_scale, data_info, c, **kwargs):
|
||||||
|
return self.base_model.forward_with_cfg(x, t, y, cfg_scale, data_info, c=self.forward_c(c), **kwargs)
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict: Mapping[str, Any], strict: bool = True):
|
||||||
|
if all((k.startswith('base_model') or k.startswith('controlnet')) for k in state_dict.keys()):
|
||||||
|
return super().load_state_dict(state_dict, strict)
|
||||||
|
else:
|
||||||
|
new_key = {}
|
||||||
|
for k in state_dict.keys():
|
||||||
|
new_key[k] = re.sub(r"(blocks\.\d+)(.*)", r"\1.base_block\2", k)
|
||||||
|
for k, v in new_key.items():
|
||||||
|
if k != v:
|
||||||
|
print(f"replace {k} to {v}")
|
||||||
|
state_dict[v] = state_dict.pop(k)
|
||||||
|
|
||||||
|
return self.base_model.load_state_dict(state_dict, strict)
|
||||||
|
|
||||||
|
def unpatchify(self, x):
|
||||||
|
"""
|
||||||
|
x: (N, T, patch_size**2 * C)
|
||||||
|
imgs: (N, H, W, C)
|
||||||
|
"""
|
||||||
|
c = self.out_channels
|
||||||
|
p = self.x_embedder.patch_size[0]
|
||||||
|
assert self.h * self.w == x.shape[1]
|
||||||
|
|
||||||
|
x = x.reshape(shape=(x.shape[0], self.h, self.w, p, p, c))
|
||||||
|
x = torch.einsum('nhwpqc->nchpwq', x)
|
||||||
|
imgs = x.reshape(shape=(x.shape[0], c, self.h * p, self.w * p))
|
||||||
|
return imgs
|
||||||
|
|
||||||
|
# @property
|
||||||
|
# def dtype(self):
|
||||||
|
## 返回模型参数的数据类型
|
||||||
|
# return next(self.parameters()).dtype
|
||||||
|
|
||||||
|
|
||||||
|
# The implementation for PixArtMS_Half + 1024 resolution
|
||||||
|
class ControlPixArtMSHalf(ControlPixArtHalf):
|
||||||
|
# support multi-scale res model (multi-scale model can also be applied to single reso training & inference)
|
||||||
|
def __init__(self, base_model: PixArtMS, copy_blocks_num: int = 13) -> None:
|
||||||
|
super().__init__(base_model=base_model, copy_blocks_num=copy_blocks_num)
|
||||||
|
|
||||||
|
def forward_raw(self, x, timestep, y, mask=None, data_info=None, c=None, **kwargs):
|
||||||
|
# modify the original PixArtMS forward function
|
||||||
|
"""
|
||||||
|
Forward pass of PixArt.
|
||||||
|
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
||||||
|
t: (N,) tensor of diffusion timesteps
|
||||||
|
y: (N, 1, 120, C) tensor of class labels
|
||||||
|
"""
|
||||||
|
if c is not None:
|
||||||
|
c = c.to(self.dtype)
|
||||||
|
c = self.forward_c(c)
|
||||||
|
bs = x.shape[0]
|
||||||
|
x = x.to(self.dtype)
|
||||||
|
timestep = timestep.to(self.dtype)
|
||||||
|
y = y.to(self.dtype)
|
||||||
|
c_size, ar = data_info['img_hw'].to(self.dtype), data_info['aspect_ratio'].to(self.dtype)
|
||||||
|
self.h, self.w = x.shape[-2]//self.patch_size, x.shape[-1]//self.patch_size
|
||||||
|
|
||||||
|
pos_embed = torch.from_numpy(get_2d_sincos_pos_embed(self.pos_embed.shape[-1], (self.h, self.w), lewei_scale=self.lewei_scale, base_size=self.base_size)).unsqueeze(0).to(x.device).to(self.dtype)
|
||||||
|
x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2
|
||||||
|
t = self.t_embedder(timestep) # (N, D)
|
||||||
|
csize = self.csize_embedder(c_size, bs) # (N, D)
|
||||||
|
ar = self.ar_embedder(ar, bs) # (N, D)
|
||||||
|
t = t + torch.cat([csize, ar], dim=1)
|
||||||
|
t0 = self.t_block(t)
|
||||||
|
y = self.y_embedder(y, self.training) # (N, D)
|
||||||
|
if mask is not None:
|
||||||
|
if mask.shape[0] != y.shape[0]:
|
||||||
|
mask = mask.repeat(y.shape[0] // mask.shape[0], 1)
|
||||||
|
mask = mask.squeeze(1).squeeze(1)
|
||||||
|
y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1])
|
||||||
|
y_lens = mask.sum(dim=1).tolist()
|
||||||
|
else:
|
||||||
|
y_lens = [y.shape[2]] * y.shape[0]
|
||||||
|
y = y.squeeze(1).view(1, -1, x.shape[-1])
|
||||||
|
|
||||||
|
# define the first layer
|
||||||
|
x = auto_grad_checkpoint(self.base_model.blocks[0], x, y, t0, y_lens, **kwargs) # (N, T, D) #support grad checkpoint
|
||||||
|
|
||||||
|
if c is not None:
|
||||||
|
# update c
|
||||||
|
for index in range(1, self.copy_blocks_num + 1):
|
||||||
|
c, c_skip = auto_grad_checkpoint(self.controlnet[index - 1], x, y, t0, y_lens, c, **kwargs)
|
||||||
|
x = auto_grad_checkpoint(self.base_model.blocks[index], x + c_skip, y, t0, y_lens, **kwargs)
|
||||||
|
|
||||||
|
# update x
|
||||||
|
for index in range(self.copy_blocks_num + 1, self.total_blocks_num):
|
||||||
|
x = auto_grad_checkpoint(self.base_model.blocks[index], x, y, t0, y_lens, **kwargs)
|
||||||
|
else:
|
||||||
|
for index in range(1, self.total_blocks_num):
|
||||||
|
x = auto_grad_checkpoint(self.base_model.blocks[index], x, y, t0, y_lens, **kwargs)
|
||||||
|
|
||||||
|
x = self.final_layer(x, t) # (N, T, patch_size ** 2 * out_channels)
|
||||||
|
x = self.unpatchify(x) # (N, out_channels, H, W)
|
||||||
|
return x
|
||||||
|
|
||||||
|
def forward(self, x, timesteps, context, img_hw=None, aspect_ratio=None, cn_hint=None, **kwargs):
|
||||||
|
"""
|
||||||
|
Forward pass that adapts comfy input to original forward function
|
||||||
|
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
||||||
|
timesteps: (N,) tensor of diffusion timesteps
|
||||||
|
context: (N, 1, 120, C) conditioning
|
||||||
|
img_hw: height|width conditioning
|
||||||
|
aspect_ratio: aspect ratio conditioning
|
||||||
|
cn_hint: controlnet hint
|
||||||
|
"""
|
||||||
|
## size/ar from cond with fallback based on the latent image shape.
|
||||||
|
bs = x.shape[0]
|
||||||
|
data_info = {}
|
||||||
|
if img_hw is None:
|
||||||
|
data_info["img_hw"] = torch.tensor(
|
||||||
|
[[x.shape[2]*8, x.shape[3]*8]],
|
||||||
|
dtype=self.dtype,
|
||||||
|
device=x.device
|
||||||
|
).repeat(bs, 1)
|
||||||
|
else:
|
||||||
|
data_info["img_hw"] = img_hw.to(x.dtype)
|
||||||
|
if aspect_ratio is None or True:
|
||||||
|
data_info["aspect_ratio"] = torch.tensor(
|
||||||
|
[[x.shape[2]/x.shape[3]]],
|
||||||
|
dtype=self.dtype,
|
||||||
|
device=x.device
|
||||||
|
).repeat(bs, 1)
|
||||||
|
else:
|
||||||
|
data_info["aspect_ratio"] = aspect_ratio.to(x.dtype)
|
||||||
|
|
||||||
|
## Still accepts the input w/o that dim but returns garbage
|
||||||
|
if len(context.shape) == 3:
|
||||||
|
context = context.unsqueeze(1)
|
||||||
|
|
||||||
|
## run original forward pass
|
||||||
|
out = self.forward_raw(
|
||||||
|
x = x.to(self.dtype),
|
||||||
|
timestep = timesteps.to(self.dtype),
|
||||||
|
y = context.to(self.dtype),
|
||||||
|
c = cn_hint,
|
||||||
|
data_info=data_info,
|
||||||
|
)
|
||||||
|
|
||||||
|
## only return EPS
|
||||||
|
out = out.to(torch.float)
|
||||||
|
eps, rest = out[:, :self.in_channels], out[:, self.in_channels:]
|
||||||
|
return eps
|
||||||
+52
-1
@@ -38,7 +38,7 @@ class PixArtResolutionSelect():
|
|||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"model": (list(pixart_conf.keys()),),
|
"model": (list(pixart_res.keys()),),
|
||||||
# keys are the same for both
|
# keys are the same for both
|
||||||
"ratio": (list(pixart_res["PixArtMS_XL_2"].keys()),{"default":"1.00"}),
|
"ratio": (list(pixart_res["PixArtMS_XL_2"].keys()),{"default":"1.00"}),
|
||||||
}
|
}
|
||||||
@@ -92,6 +92,55 @@ class PixArtLoraLoader:
|
|||||||
model_lora = load_pixart_lora(model, lora, lora_path, strength,)
|
model_lora = load_pixart_lora(model, lora, lora_path, strength,)
|
||||||
return (model_lora,)
|
return (model_lora,)
|
||||||
|
|
||||||
|
class PixArtResolutionCond:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"cond": ("CONDITIONING", ),
|
||||||
|
"width": ("INT", {"default": 1024.0, "min": 0, "max": 8192}),
|
||||||
|
"height": ("INT", {"default": 1024.0, "min": 0, "max": 8192}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("CONDITIONING",)
|
||||||
|
RETURN_NAMES = ("cond",)
|
||||||
|
FUNCTION = "add_cond"
|
||||||
|
CATEGORY = "ExtraModels/PixArt"
|
||||||
|
TITLE = "PixArt Resolution Conditioning"
|
||||||
|
|
||||||
|
def add_cond(self, cond, width, height):
|
||||||
|
for c in range(len(cond)):
|
||||||
|
cond[c][1].update({
|
||||||
|
"img_hw": [[height, width]],
|
||||||
|
"aspect_ratio": [[height/width]],
|
||||||
|
})
|
||||||
|
return (cond,)
|
||||||
|
|
||||||
|
class PixArtControlNetCond:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"cond": ("CONDITIONING",),
|
||||||
|
"latent": ("LATENT",),
|
||||||
|
# "image": ("IMAGE",),
|
||||||
|
# "vae": ("VAE",),
|
||||||
|
# "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("CONDITIONING",)
|
||||||
|
RETURN_NAMES = ("cond",)
|
||||||
|
FUNCTION = "add_cond"
|
||||||
|
CATEGORY = "ExtraModels/PixArt"
|
||||||
|
TITLE = "PixArt ControlNet Conditioning"
|
||||||
|
|
||||||
|
def add_cond(self, cond, latent):
|
||||||
|
for c in range(len(cond)):
|
||||||
|
cond[c][1]["cn_hint"] = latent["samples"] * 0.18215
|
||||||
|
return (cond,)
|
||||||
|
|
||||||
class PixArtDPMSampler:
|
class PixArtDPMSampler:
|
||||||
"""
|
"""
|
||||||
The sampler from the reference code.
|
The sampler from the reference code.
|
||||||
@@ -189,4 +238,6 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"PixArtLoraLoader" : PixArtLoraLoader,
|
"PixArtLoraLoader" : PixArtLoraLoader,
|
||||||
"PixArtDPMSampler" : PixArtDPMSampler,
|
"PixArtDPMSampler" : PixArtDPMSampler,
|
||||||
"PixArtT5TextEncode" : PixArtT5TextEncode,
|
"PixArtT5TextEncode" : PixArtT5TextEncode,
|
||||||
|
"PixArtResolutionCond" : PixArtResolutionCond,
|
||||||
|
"PixArtControlNetCond" : PixArtControlNetCond,
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user