This commit is contained in:
City
2024-12-11 17:24:18 +01:00
parent d5a47fa3e5
commit 6a88140002
3 changed files with 6 additions and 22 deletions
+1 -3
View File
@@ -1,12 +1,10 @@
"""
Model config and setting logic
"""
import math import math
import logging import logging
import comfy.supported_models_base import comfy.supported_models_base
import comfy.supported_models import comfy.supported_models
import comfy.latent_formats import comfy.latent_formats
import comfy.model_base
from .model.pixart import PixArt from .model.pixart import PixArt
from .model.pixartms import PixArtMS from .model.pixartms import PixArtMS
+5 -19
View File
@@ -1,28 +1,14 @@
import comfy.supported_models_base
import comfy.latent_formats
import comfy.model_detection
import comfy.model_patcher
import comfy.model_base
import comfy.utils
import comfy.conds
import logging import logging
import torch
import comfy.utils
import comfy.model_base
import comfy.model_patcher
import comfy.model_detection
from comfy import model_management from comfy import model_management
from .diffusers_convert import convert_state_dict from .diffusers_convert import convert_state_dict
from .config import model_config_from_unet from .config import model_config_from_unet
class PixArtModel(comfy.model_base.BaseModel):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
def extra_conds(self, **kwargs):
out = super().extra_conds(**kwargs)
for name in ["width", "height", "aspect_ratio", "img_hw"]: # TODO: remove last one
out[name] = comfy.conds.CONDRegular(torch.tensor(name))
return out
def load_pixart_state_dict(sd, model_options={}): def load_pixart_state_dict(sd, model_options={}):
# prefix / format # prefix / format
sd = sd.get("model", sd) # ref ckpt sd = sd.get("model", sd) # ref ckpt
View File