This commit is contained in:
Muhamad Syaifullah 2024-10-16 14:22:18 +07:00
parent 0dbba9f751
commit 82cbb7647b
No known key found for this signature in database
176 changed files with 903 additions and 903 deletions

View File

@ -47,7 +47,7 @@ def pull(repo, remote_name='origin', branch='master'):
pygit2.option(pygit2.GIT_OPT_SET_OWNER_VALIDATION, 0) pygit2.option(pygit2.GIT_OPT_SET_OWNER_VALIDATION, 0)
repo_path = str(sys.argv[1]) repo_path = str(sys.argv[1])
repo = pygit2.Repository(repo_path) repo = pygit2.Repository(repo_path)
ident = pygit2.Signature('comfyui', 'comfy@ui') ident = pygit2.Signature('comfyui', 'seap@ui')
try: try:
print("stashing current changes") print("stashing current changes")
repo.stash(ident) repo.stash(ident)

View File

@ -9,7 +9,7 @@ class AppSettings():
def get_settings(self, request): def get_settings(self, request):
file = self.user_manager.get_request_user_filepath( file = self.user_manager.get_request_user_filepath(
request, "comfy.settings.json") request, "seap.settings.json")
if os.path.isfile(file): if os.path.isfile(file):
with open(file) as f: with open(file) as f:
return json.load(f) return json.load(f)
@ -18,7 +18,7 @@ class AppSettings():
def save_settings(self, request, settings): def save_settings(self, request, settings):
file = self.user_manager.get_request_user_filepath( file = self.user_manager.get_request_user_filepath(
request, "comfy.settings.json") request, "seap.settings.json")
with open(file, "w") as f: with open(file, "w") as f:
f.write(json.dumps(settings, indent=4)) f.write(json.dumps(settings, indent=4))

View File

@ -12,7 +12,7 @@ from typing import TypedDict, Optional
import requests import requests
from typing_extensions import NotRequired from typing_extensions import NotRequired
from comfy.cli_args import DEFAULT_VERSION_STRING from seap.cli_args import DEFAULT_VERSION_STRING
REQUEST_TIMEOUT = 10 # seconds REQUEST_TIMEOUT = 10 # seconds

View File

@ -6,7 +6,7 @@ import glob
import shutil import shutil
from aiohttp import web from aiohttp import web
from urllib import parse from urllib import parse
from comfy.cli_args import args from seap.cli_args import args
import folder_paths import folder_paths
from .app_settings import AppSettings from .app_settings import AppSettings
@ -38,8 +38,8 @@ class UserManager():
def get_request_user_id(self, request): def get_request_user_id(self, request):
user = "default" user = "default"
if args.multi_user and "comfy-user" in request.headers: if args.multi_user and "seap-user" in request.headers:
user = request.headers["comfy-user"] user = request.headers["seap-user"]
if user not in self.users: if user not in self.users:
raise KeyError("Unknown user: " + user) raise KeyError("Unknown user: " + user)

View File

@ -1,6 +1,6 @@
import os import os
import importlib.util import importlib.util
from comfy.cli_args import args from seap.cli_args import args
import subprocess import subprocess
#Can't use pytorch to get the GPU names because the cuda malloc has to be set before the first import. #Can't use pytorch to get the GPU names because the cuda malloc has to be set before the first import.

View File

@ -2,7 +2,7 @@ from PIL import Image, ImageOps
from io import BytesIO from io import BytesIO
import numpy as np import numpy as np
import struct import struct
import comfy.utils import seap.utils
import time import time
#You can use this node to save full size images through the websocket, the #You can use this node to save full size images through the websocket, the
@ -27,7 +27,7 @@ class SaveImageWebsocket:
CATEGORY = "api/image" CATEGORY = "api/image"
def save_images(self, images): def save_images(self, images):
pbar = comfy.utils.ProgressBar(images.shape[0]) pbar = seap.utils.ProgressBar(images.shape[0])
step = 0 step = 0
for image in images: for image in images:
i = 255. * image.cpu().numpy() i = 255. * image.cpu().numpy()

View File

@ -12,11 +12,11 @@ from typing import List, Literal, NamedTuple, Optional
import torch import torch
import nodes import nodes
import comfy.model_management import seap.model_management
from comfy_execution.graph import get_input_info, ExecutionList, DynamicPrompt, ExecutionBlocker from seap_execution.graph import get_input_info, ExecutionList, DynamicPrompt, ExecutionBlocker
from comfy_execution.graph_utils import is_link, GraphBuilder from seap_execution.graph_utils import is_link, GraphBuilder
from comfy_execution.caching import HierarchicalCache, LRUCache, CacheKeySetInputSignature, CacheKeySetID from seap_execution.caching import HierarchicalCache, LRUCache, CacheKeySetInputSignature, CacheKeySetID
from comfy.cli_args import args from seap.cli_args import args
class ExecutionResult(Enum): class ExecutionResult(Enum):
SUCCESS = 0 SUCCESS = 0
@ -371,7 +371,7 @@ def execute(server, dynprompt, caches, current_item, extra_data, executed, promp
pending_subgraph_results[unique_id] = cached_outputs pending_subgraph_results[unique_id] = cached_outputs
return (ExecutionResult.PENDING, None, None) return (ExecutionResult.PENDING, None, None)
caches.outputs.set(unique_id, output_data) caches.outputs.set(unique_id, output_data)
except comfy.model_management.InterruptProcessingException as iex: except seap.model_management.InterruptProcessingException as iex:
logging.info("Processing interrupted") logging.info("Processing interrupted")
# skip formatting inputs/outputs # skip formatting inputs/outputs
@ -399,9 +399,9 @@ def execute(server, dynprompt, caches, current_item, extra_data, executed, promp
"traceback": traceback.format_tb(tb), "traceback": traceback.format_tb(tb),
"current_inputs": input_data_formatted "current_inputs": input_data_formatted
} }
if isinstance(ex, comfy.model_management.OOM_EXCEPTION): if isinstance(ex, seap.model_management.OOM_EXCEPTION):
logging.error("Got an OOM, unloading all loaded models.") logging.error("Got an OOM, unloading all loaded models.")
comfy.model_management.unload_all_models() seap.model_management.unload_all_models()
return (ExecutionResult.FAILURE, error_details, ex) return (ExecutionResult.FAILURE, error_details, ex)
@ -435,7 +435,7 @@ class PromptExecutor:
# First, send back the status to the frontend depending # First, send back the status to the frontend depending
# on the exception type # on the exception type
if isinstance(ex, comfy.model_management.InterruptProcessingException): if isinstance(ex, seap.model_management.InterruptProcessingException):
mes = { mes = {
"prompt_id": prompt_id, "prompt_id": prompt_id,
"node_id": node_id, "node_id": node_id,
@ -480,7 +480,7 @@ class PromptExecutor:
if self.caches.outputs.get(node_id) is not None: if self.caches.outputs.get(node_id) is not None:
cached_nodes.append(node_id) cached_nodes.append(node_id)
comfy.model_management.cleanup_models(keep_clone_weights_loaded=True) seap.model_management.cleanup_models(keep_clone_weights_loaded=True)
self.add_message("execution_cached", self.add_message("execution_cached",
{ "nodes": cached_nodes, "prompt_id": prompt_id}, { "nodes": cached_nodes, "prompt_id": prompt_id},
broadcast=False) broadcast=False)
@ -523,8 +523,8 @@ class PromptExecutor:
"meta": meta_outputs, "meta": meta_outputs,
} }
self.server.last_node_id = None self.server.last_node_id = None
if comfy.model_management.DISABLE_SMART_MEMORY: if seap.model_management.DISABLE_SMART_MEMORY:
comfy.model_management.unload_all_models() seap.model_management.unload_all_models()

View File

@ -2,11 +2,11 @@ import torch
from PIL import Image from PIL import Image
import struct import struct
import numpy as np import numpy as np
from comfy.cli_args import args, LatentPreviewMethod from seap.cli_args import args, LatentPreviewMethod
from comfy.taesd.taesd import TAESD from seap.taesd.taesd import TAESD
import comfy.model_management import seap.model_management
import folder_paths import folder_paths
import comfy.utils import seap.utils
import logging import logging
MAX_PREVIEW_RESOLUTION = args.preview_size MAX_PREVIEW_RESOLUTION = args.preview_size
@ -14,7 +14,7 @@ MAX_PREVIEW_RESOLUTION = args.preview_size
def preview_to_image(latent_image): def preview_to_image(latent_image):
latents_ubyte = (((latent_image + 1.0) / 2.0).clamp(0, 1) # change scale from -1..1 to 0..1 latents_ubyte = (((latent_image + 1.0) / 2.0).clamp(0, 1) # change scale from -1..1 to 0..1
.mul(0xFF) # to 0..255 .mul(0xFF) # to 0..255
).to(device="cpu", dtype=torch.uint8, non_blocking=comfy.model_management.device_supports_non_blocking(latent_image.device)) ).to(device="cpu", dtype=torch.uint8, non_blocking=seap.model_management.device_supports_non_blocking(latent_image.device))
return Image.fromarray(latents_ubyte.numpy()) return Image.fromarray(latents_ubyte.numpy())
@ -89,7 +89,7 @@ def prepare_callback(model, steps, x0_output_dict=None):
previewer = get_previewer(model.load_device, model.model.latent_format) previewer = get_previewer(model.load_device, model.model.latent_format)
pbar = comfy.utils.ProgressBar(steps) pbar = seap.utils.ProgressBar(steps)
def callback(step, x0, x, total_steps): def callback(step, x0, x, total_steps):
if x0_output_dict is not None: if x0_output_dict is not None:
x0_output_dict["x0"] = x0 x0_output_dict["x0"] = x0

24
main.py
View File

@ -1,11 +1,11 @@
import comfy.options import seap.options
comfy.options.enable_args_parsing() seap.options.enable_args_parsing()
import os import os
import importlib.util import importlib.util
import folder_paths import folder_paths
import time import time
from comfy.cli_args import args from seap.cli_args import args
from app.logger import setup_logger from app.logger import setup_logger
@ -85,17 +85,17 @@ if args.windows_standalone_build:
except: except:
pass pass
import comfy.utils import seap.utils
import execution import execution
import server import server
from server import BinaryEventTypes from server import BinaryEventTypes
import nodes import nodes
import comfy.model_management import seap.model_management
def cuda_malloc_warning(): def cuda_malloc_warning():
device = comfy.model_management.get_torch_device() device = seap.model_management.get_torch_device()
device_name = comfy.model_management.get_torch_device_name(device) device_name = seap.model_management.get_torch_device_name(device)
cuda_malloc_warning = False cuda_malloc_warning = False
if "cudaMallocAsync" in device_name: if "cudaMallocAsync" in device_name:
for b in cuda_malloc.blacklist: for b in cuda_malloc.blacklist:
@ -141,7 +141,7 @@ def prompt_worker(q, server):
free_memory = flags.get("free_memory", False) free_memory = flags.get("free_memory", False)
if flags.get("unload_models", free_memory): if flags.get("unload_models", free_memory):
comfy.model_management.unload_all_models() seap.model_management.unload_all_models()
need_gc = True need_gc = True
last_gc_collect = 0 last_gc_collect = 0
@ -153,9 +153,9 @@ def prompt_worker(q, server):
if need_gc: if need_gc:
current_time = time.perf_counter() current_time = time.perf_counter()
if (current_time - last_gc_collect) > gc_collect_interval: if (current_time - last_gc_collect) > gc_collect_interval:
comfy.model_management.cleanup_models() seap.model_management.cleanup_models()
gc.collect() gc.collect()
comfy.model_management.soft_empty_cache() seap.model_management.soft_empty_cache()
last_gc_collect = current_time last_gc_collect = current_time
need_gc = False need_gc = False
@ -168,13 +168,13 @@ async def run(server, address='', port=8188, verbose=True, call_on_start=None):
def hijack_progress(server): def hijack_progress(server):
def hook(value, total, preview_image): def hook(value, total, preview_image):
comfy.model_management.throw_exception_if_processing_interrupted() seap.model_management.throw_exception_if_processing_interrupted()
progress = {"value": value, "max": total, "prompt_id": server.last_prompt_id, "node": server.last_node_id} progress = {"value": value, "max": total, "prompt_id": server.last_prompt_id, "node": server.last_node_id}
server.send_sync("progress", progress, server.client_id) server.send_sync("progress", progress, server.client_id)
if preview_image is not None: if preview_image is not None:
server.send_sync(BinaryEventTypes.UNENCODED_PREVIEW_IMAGE, preview_image, server.client_id) server.send_sync(BinaryEventTypes.UNENCODED_PREVIEW_IMAGE, preview_image, server.client_id)
comfy.utils.set_progress_bar_global_hook(hook) seap.utils.set_progress_bar_global_hook(hook)
def cleanup_temp(): def cleanup_temp():

View File

@ -1,6 +1,6 @@
import hashlib import hashlib
from comfy.cli_args import args from seap.cli_args import args
from PIL import ImageFile, UnidentifiedImageError from PIL import ImageFile, UnidentifiedImageError

118
nodes.py
View File

@ -16,19 +16,19 @@ from PIL.PngImagePlugin import PngInfo
import numpy as np import numpy as np
import safetensors.torch import safetensors.torch
sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy")) sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "seap"))
import comfy.diffusers_load import seap.diffusers_load
import comfy.samplers import seap.samplers
import comfy.sample import seap.sample
import comfy.sd import seap.sd
import comfy.utils import seap.utils
import comfy.controlnet import seap.controlnet
import comfy.clip_vision import seap.clip_vision
import comfy.model_management import seap.model_management
from comfy.cli_args import args from seap.cli_args import args
import importlib import importlib
@ -37,10 +37,10 @@ import latent_preview
import node_helpers import node_helpers
def before_node_execution(): def before_node_execution():
comfy.model_management.throw_exception_if_processing_interrupted() seap.model_management.throw_exception_if_processing_interrupted()
def interrupt_processing(value=True): def interrupt_processing(value=True):
comfy.model_management.interrupt_current_processing(value) seap.model_management.interrupt_current_processing(value)
MAX_RESOLUTION=16384 MAX_RESOLUTION=16384
@ -462,7 +462,7 @@ class SaveLatent:
output["latent_tensor"] = samples["samples"] output["latent_tensor"] = samples["samples"]
output["latent_format_version_0"] = torch.tensor([]) output["latent_format_version_0"] = torch.tensor([])
comfy.utils.save_torch_file(output, file, metadata=metadata) seap.utils.save_torch_file(output, file, metadata=metadata)
return { "ui": { "latents": results } } return { "ui": { "latents": results } }
@ -516,7 +516,7 @@ class CheckpointLoader:
def load_checkpoint(self, config_name, ckpt_name): def load_checkpoint(self, config_name, ckpt_name):
config_path = folder_paths.get_full_path("configs", config_name) config_path = folder_paths.get_full_path("configs", config_name)
ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name) ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name)
return comfy.sd.load_checkpoint(config_path, ckpt_path, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings")) return seap.sd.load_checkpoint(config_path, ckpt_path, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings"))
class CheckpointLoaderSimple: class CheckpointLoaderSimple:
@classmethod @classmethod
@ -537,7 +537,7 @@ class CheckpointLoaderSimple:
def load_checkpoint(self, ckpt_name): def load_checkpoint(self, ckpt_name):
ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name) ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name)
out = comfy.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings")) out = seap.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings"))
return out[:3] return out[:3]
class DiffusersLoader: class DiffusersLoader:
@ -564,7 +564,7 @@ class DiffusersLoader:
model_path = path model_path = path
break break
return comfy.diffusers_load.load_diffusers(model_path, output_vae=output_vae, output_clip=output_clip, embedding_directory=folder_paths.get_folder_paths("embeddings")) return seap.diffusers_load.load_diffusers(model_path, output_vae=output_vae, output_clip=output_clip, embedding_directory=folder_paths.get_folder_paths("embeddings"))
class unCLIPCheckpointLoader: class unCLIPCheckpointLoader:
@ -579,7 +579,7 @@ class unCLIPCheckpointLoader:
def load_checkpoint(self, ckpt_name, output_vae=True, output_clip=True): def load_checkpoint(self, ckpt_name, output_vae=True, output_clip=True):
ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name) ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name)
out = comfy.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, output_clipvision=True, embedding_directory=folder_paths.get_folder_paths("embeddings")) out = seap.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, output_clipvision=True, embedding_directory=folder_paths.get_folder_paths("embeddings"))
return out return out
class CLIPSetLastLayer: class CLIPSetLastLayer:
@ -636,10 +636,10 @@ class LoraLoader:
del temp del temp
if lora is None: if lora is None:
lora = comfy.utils.load_torch_file(lora_path, safe_load=True) lora = seap.utils.load_torch_file(lora_path, safe_load=True)
self.loaded_lora = (lora_path, lora) self.loaded_lora = (lora_path, lora)
model_lora, clip_lora = comfy.sd.load_lora_for_models(model, clip, lora, strength_model, strength_clip) model_lora, clip_lora = seap.sd.load_lora_for_models(model, clip, lora, strength_model, strength_clip)
return (model_lora, clip_lora) return (model_lora, clip_lora)
class LoraLoaderModelOnly(LoraLoader): class LoraLoaderModelOnly(LoraLoader):
@ -704,11 +704,11 @@ class VAELoader:
encoder = next(filter(lambda a: a.startswith("{}_encoder.".format(name)), approx_vaes)) encoder = next(filter(lambda a: a.startswith("{}_encoder.".format(name)), approx_vaes))
decoder = next(filter(lambda a: a.startswith("{}_decoder.".format(name)), approx_vaes)) decoder = next(filter(lambda a: a.startswith("{}_decoder.".format(name)), approx_vaes))
enc = comfy.utils.load_torch_file(folder_paths.get_full_path_or_raise("vae_approx", encoder)) enc = seap.utils.load_torch_file(folder_paths.get_full_path_or_raise("vae_approx", encoder))
for k in enc: for k in enc:
sd["taesd_encoder.{}".format(k)] = enc[k] sd["taesd_encoder.{}".format(k)] = enc[k]
dec = comfy.utils.load_torch_file(folder_paths.get_full_path_or_raise("vae_approx", decoder)) dec = seap.utils.load_torch_file(folder_paths.get_full_path_or_raise("vae_approx", decoder))
for k in dec: for k in dec:
sd["taesd_decoder.{}".format(k)] = dec[k] sd["taesd_decoder.{}".format(k)] = dec[k]
@ -740,8 +740,8 @@ class VAELoader:
sd = self.load_taesd(vae_name) sd = self.load_taesd(vae_name)
else: else:
vae_path = folder_paths.get_full_path_or_raise("vae", vae_name) vae_path = folder_paths.get_full_path_or_raise("vae", vae_name)
sd = comfy.utils.load_torch_file(vae_path) sd = seap.utils.load_torch_file(vae_path)
vae = comfy.sd.VAE(sd=sd) vae = seap.sd.VAE(sd=sd)
return (vae,) return (vae,)
class ControlNetLoader: class ControlNetLoader:
@ -756,7 +756,7 @@ class ControlNetLoader:
def load_controlnet(self, control_net_name): def load_controlnet(self, control_net_name):
controlnet_path = folder_paths.get_full_path_or_raise("controlnet", control_net_name) controlnet_path = folder_paths.get_full_path_or_raise("controlnet", control_net_name)
controlnet = comfy.controlnet.load_controlnet(controlnet_path) controlnet = seap.controlnet.load_controlnet(controlnet_path)
return (controlnet,) return (controlnet,)
class DiffControlNetLoader: class DiffControlNetLoader:
@ -772,7 +772,7 @@ class DiffControlNetLoader:
def load_controlnet(self, model, control_net_name): def load_controlnet(self, model, control_net_name):
controlnet_path = folder_paths.get_full_path_or_raise("controlnet", control_net_name) controlnet_path = folder_paths.get_full_path_or_raise("controlnet", control_net_name)
controlnet = comfy.controlnet.load_controlnet(controlnet_path, model) controlnet = seap.controlnet.load_controlnet(controlnet_path, model)
return (controlnet,) return (controlnet,)
@ -879,7 +879,7 @@ class UNETLoader:
model_options["dtype"] = torch.float8_e5m2 model_options["dtype"] = torch.float8_e5m2
unet_path = folder_paths.get_full_path_or_raise("diffusion_models", unet_name) unet_path = folder_paths.get_full_path_or_raise("diffusion_models", unet_name)
model = comfy.sd.load_diffusion_model(unet_path, model_options=model_options) model = seap.sd.load_diffusion_model(unet_path, model_options=model_options)
return (model,) return (model,)
class CLIPLoader: class CLIPLoader:
@ -895,16 +895,16 @@ class CLIPLoader:
def load_clip(self, clip_name, type="stable_diffusion"): def load_clip(self, clip_name, type="stable_diffusion"):
if type == "stable_cascade": if type == "stable_cascade":
clip_type = comfy.sd.CLIPType.STABLE_CASCADE clip_type = seap.sd.CLIPType.STABLE_CASCADE
elif type == "sd3": elif type == "sd3":
clip_type = comfy.sd.CLIPType.SD3 clip_type = seap.sd.CLIPType.SD3
elif type == "stable_audio": elif type == "stable_audio":
clip_type = comfy.sd.CLIPType.STABLE_AUDIO clip_type = seap.sd.CLIPType.STABLE_AUDIO
else: else:
clip_type = comfy.sd.CLIPType.STABLE_DIFFUSION clip_type = seap.sd.CLIPType.STABLE_DIFFUSION
clip_path = folder_paths.get_full_path_or_raise("clip", clip_name) clip_path = folder_paths.get_full_path_or_raise("clip", clip_name)
clip = comfy.sd.load_clip(ckpt_paths=[clip_path], embedding_directory=folder_paths.get_folder_paths("embeddings"), clip_type=clip_type) clip = seap.sd.load_clip(ckpt_paths=[clip_path], embedding_directory=folder_paths.get_folder_paths("embeddings"), clip_type=clip_type)
return (clip,) return (clip,)
class DualCLIPLoader: class DualCLIPLoader:
@ -923,13 +923,13 @@ class DualCLIPLoader:
clip_path1 = folder_paths.get_full_path_or_raise("clip", clip_name1) clip_path1 = folder_paths.get_full_path_or_raise("clip", clip_name1)
clip_path2 = folder_paths.get_full_path_or_raise("clip", clip_name2) clip_path2 = folder_paths.get_full_path_or_raise("clip", clip_name2)
if type == "sdxl": if type == "sdxl":
clip_type = comfy.sd.CLIPType.STABLE_DIFFUSION clip_type = seap.sd.CLIPType.STABLE_DIFFUSION
elif type == "sd3": elif type == "sd3":
clip_type = comfy.sd.CLIPType.SD3 clip_type = seap.sd.CLIPType.SD3
elif type == "flux": elif type == "flux":
clip_type = comfy.sd.CLIPType.FLUX clip_type = seap.sd.CLIPType.FLUX
clip = comfy.sd.load_clip(ckpt_paths=[clip_path1, clip_path2], embedding_directory=folder_paths.get_folder_paths("embeddings"), clip_type=clip_type) clip = seap.sd.load_clip(ckpt_paths=[clip_path1, clip_path2], embedding_directory=folder_paths.get_folder_paths("embeddings"), clip_type=clip_type)
return (clip,) return (clip,)
class CLIPVisionLoader: class CLIPVisionLoader:
@ -944,7 +944,7 @@ class CLIPVisionLoader:
def load_clip(self, clip_name): def load_clip(self, clip_name):
clip_path = folder_paths.get_full_path_or_raise("clip_vision", clip_name) clip_path = folder_paths.get_full_path_or_raise("clip_vision", clip_name)
clip_vision = comfy.clip_vision.load(clip_path) clip_vision = seap.clip_vision.load(clip_path)
return (clip_vision,) return (clip_vision,)
class CLIPVisionEncode: class CLIPVisionEncode:
@ -974,7 +974,7 @@ class StyleModelLoader:
def load_style_model(self, style_model_name): def load_style_model(self, style_model_name):
style_model_path = folder_paths.get_full_path_or_raise("style_models", style_model_name) style_model_path = folder_paths.get_full_path_or_raise("style_models", style_model_name)
style_model = comfy.sd.load_style_model(style_model_path) style_model = seap.sd.load_style_model(style_model_path)
return (style_model,) return (style_model,)
@ -1039,7 +1039,7 @@ class GLIGENLoader:
def load_gligen(self, gligen_name): def load_gligen(self, gligen_name):
gligen_path = folder_paths.get_full_path_or_raise("gligen", gligen_name) gligen_path = folder_paths.get_full_path_or_raise("gligen", gligen_name)
gligen = comfy.sd.load_gligen(gligen_path) gligen = seap.sd.load_gligen(gligen_path)
return (gligen,) return (gligen,)
class GLIGENTextBoxApply: class GLIGENTextBoxApply:
@ -1075,7 +1075,7 @@ class GLIGENTextBoxApply:
class EmptyLatentImage: class EmptyLatentImage:
def __init__(self): def __init__(self):
self.device = comfy.model_management.intermediate_device() self.device = seap.model_management.intermediate_device()
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@ -1187,7 +1187,7 @@ class LatentUpscale:
width = max(64, width) width = max(64, width)
height = max(64, height) height = max(64, height)
s["samples"] = comfy.utils.common_upscale(samples["samples"], width // 8, height // 8, upscale_method, crop) s["samples"] = seap.utils.common_upscale(samples["samples"], width // 8, height // 8, upscale_method, crop)
return (s,) return (s,)
class LatentUpscaleBy: class LatentUpscaleBy:
@ -1206,7 +1206,7 @@ class LatentUpscaleBy:
s = samples.copy() s = samples.copy()
width = round(samples["samples"].shape[3] * scale_by) width = round(samples["samples"].shape[3] * scale_by)
height = round(samples["samples"].shape[2] * scale_by) height = round(samples["samples"].shape[2] * scale_by)
s["samples"] = comfy.utils.common_upscale(samples["samples"], width, height, upscale_method, "disabled") s["samples"] = seap.utils.common_upscale(samples["samples"], width, height, upscale_method, "disabled")
return (s,) return (s,)
class LatentRotate: class LatentRotate:
@ -1322,7 +1322,7 @@ class LatentBlend:
if samples1.shape != samples2.shape: if samples1.shape != samples2.shape:
samples2.permute(0, 3, 1, 2) samples2.permute(0, 3, 1, 2)
samples2 = comfy.utils.common_upscale(samples2, samples1.shape[3], samples1.shape[2], 'bicubic', crop='center') samples2 = seap.utils.common_upscale(samples2, samples1.shape[3], samples1.shape[2], 'bicubic', crop='center')
samples2.permute(0, 2, 3, 1) samples2.permute(0, 2, 3, 1)
samples_blended = self.blend_mode(samples1, samples2, blend_mode) samples_blended = self.blend_mode(samples1, samples2, blend_mode)
@ -1387,23 +1387,23 @@ class SetLatentNoiseMask:
def common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent, denoise=1.0, disable_noise=False, start_step=None, last_step=None, force_full_denoise=False): def common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent, denoise=1.0, disable_noise=False, start_step=None, last_step=None, force_full_denoise=False):
latent_image = latent["samples"] latent_image = latent["samples"]
latent_image = comfy.sample.fix_empty_latent_channels(model, latent_image) latent_image = seap.sample.fix_empty_latent_channels(model, latent_image)
if disable_noise: if disable_noise:
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu") noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
else: else:
batch_inds = latent["batch_index"] if "batch_index" in latent else None batch_inds = latent["batch_index"] if "batch_index" in latent else None
noise = comfy.sample.prepare_noise(latent_image, seed, batch_inds) noise = seap.sample.prepare_noise(latent_image, seed, batch_inds)
noise_mask = None noise_mask = None
if "noise_mask" in latent: if "noise_mask" in latent:
noise_mask = latent["noise_mask"] noise_mask = latent["noise_mask"]
callback = latent_preview.prepare_callback(model, steps) callback = latent_preview.prepare_callback(model, steps)
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED disable_pbar = not seap.utils.PROGRESS_BAR_ENABLED
samples = comfy.sample.sample(model, noise, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, samples = seap.sample.sample(model, noise, steps, cfg, sampler_name, scheduler, positive, negative, latent_image,
denoise=denoise, disable_noise=disable_noise, start_step=start_step, last_step=last_step, denoise=denoise, disable_noise=disable_noise, start_step=start_step, last_step=last_step,
force_full_denoise=force_full_denoise, noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=seed) force_full_denoise=force_full_denoise, noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=seed)
out = latent.copy() out = latent.copy()
out["samples"] = samples out["samples"] = samples
return (out, ) return (out, )
@ -1417,8 +1417,8 @@ class KSampler:
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "tooltip": "The random seed used for creating the noise."}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "tooltip": "The random seed used for creating the noise."}),
"steps": ("INT", {"default": 20, "min": 1, "max": 10000, "tooltip": "The number of steps used in the denoising process."}), "steps": ("INT", {"default": 20, "min": 1, "max": 10000, "tooltip": "The number of steps used in the denoising process."}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01, "tooltip": "The Classifier-Free Guidance scale balances creativity and adherence to the prompt. Higher values result in images more closely matching the prompt however too high values will negatively impact quality."}), "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01, "tooltip": "The Classifier-Free Guidance scale balances creativity and adherence to the prompt. Higher values result in images more closely matching the prompt however too high values will negatively impact quality."}),
"sampler_name": (comfy.samplers.KSampler.SAMPLERS, {"tooltip": "The algorithm used when sampling, this can affect the quality, speed, and style of the generated output."}), "sampler_name": (seap.samplers.KSampler.SAMPLERS, {"tooltip": "The algorithm used when sampling, this can affect the quality, speed, and style of the generated output."}),
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, {"tooltip": "The scheduler controls how noise is gradually removed to form the image."}), "scheduler": (seap.samplers.KSampler.SCHEDULERS, {"tooltip": "The scheduler controls how noise is gradually removed to form the image."}),
"positive": ("CONDITIONING", {"tooltip": "The conditioning describing the attributes you want to include in the image."}), "positive": ("CONDITIONING", {"tooltip": "The conditioning describing the attributes you want to include in the image."}),
"negative": ("CONDITIONING", {"tooltip": "The conditioning describing the attributes you want to exclude from the image."}), "negative": ("CONDITIONING", {"tooltip": "The conditioning describing the attributes you want to exclude from the image."}),
"latent_image": ("LATENT", {"tooltip": "The latent image to denoise."}), "latent_image": ("LATENT", {"tooltip": "The latent image to denoise."}),
@ -1445,8 +1445,8 @@ class KSamplerAdvanced:
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), "noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}), "steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"sampler_name": (comfy.samplers.KSampler.SAMPLERS, ), "sampler_name": (seap.samplers.KSampler.SAMPLERS,),
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ), "scheduler": (seap.samplers.KSampler.SCHEDULERS,),
"positive": ("CONDITIONING", ), "positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ), "negative": ("CONDITIONING", ),
"latent_image": ("LATENT", ), "latent_image": ("LATENT", ),
@ -1686,7 +1686,7 @@ class ImageScale:
elif height == 0: elif height == 0:
height = max(1, round(samples.shape[2] * width / samples.shape[3])) height = max(1, round(samples.shape[2] * width / samples.shape[3]))
s = comfy.utils.common_upscale(samples, width, height, upscale_method, crop) s = seap.utils.common_upscale(samples, width, height, upscale_method, crop)
s = s.movedim(1,-1) s = s.movedim(1,-1)
return (s,) return (s,)
@ -1706,7 +1706,7 @@ class ImageScaleBy:
samples = image.movedim(-1,1) samples = image.movedim(-1,1)
width = round(samples.shape[3] * scale_by) width = round(samples.shape[3] * scale_by)
height = round(samples.shape[2] * scale_by) height = round(samples.shape[2] * scale_by)
s = comfy.utils.common_upscale(samples, width, height, upscale_method, "disabled") s = seap.utils.common_upscale(samples, width, height, upscale_method, "disabled")
s = s.movedim(1,-1) s = s.movedim(1,-1)
return (s,) return (s,)
@ -1738,7 +1738,7 @@ class ImageBatch:
def batch(self, image1, image2): def batch(self, image1, image2):
if image1.shape[1:] != image2.shape[1:]: if image1.shape[1:] != image2.shape[1:]:
image2 = comfy.utils.common_upscale(image2.movedim(-1,1), image1.shape[2], image1.shape[1], "bilinear", "center").movedim(1,-1) image2 = seap.utils.common_upscale(image2.movedim(-1, 1), image1.shape[2], image1.shape[1], "bilinear", "center").movedim(1, -1)
s = torch.cat((image1, image2), dim=0) s = torch.cat((image1, image2), dim=0)
return (s,) return (s,)
@ -2061,13 +2061,13 @@ def init_builtin_extra_nodes():
""" """
Initializes the built-in extra nodes in ComfyUI. Initializes the built-in extra nodes in ComfyUI.
This function loads the extra node files located in the "comfy_extras" directory and imports them into ComfyUI. This function loads the extra node files located in the "seap_extras" directory and imports them into ComfyUI.
If any of the extra node files fail to import, a warning message is logged. If any of the extra node files fail to import, a warning message is logged.
Returns: Returns:
None None
""" """
extras_dir = os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy_extras") extras_dir = os.path.join(os.path.dirname(os.path.realpath(__file__)), "seap_extras")
extras_files = [ extras_files = [
"nodes_latent.py", "nodes_latent.py",
"nodes_hypernetwork.py", "nodes_hypernetwork.py",
@ -2115,7 +2115,7 @@ def init_builtin_extra_nodes():
import_failed = [] import_failed = []
for node_file in extras_files: for node_file in extras_files:
if not load_custom_node(os.path.join(extras_dir, node_file), module_parent="comfy_extras"): if not load_custom_node(os.path.join(extras_dir, node_file), module_parent="seap_extras"):
import_failed.append(node_file) import_failed.append(node_file)
return import_failed return import_failed
@ -2130,7 +2130,7 @@ def init_extra_nodes(init_custom_nodes=True):
logging.info("Skipping loading of custom nodes") logging.info("Skipping loading of custom nodes")
if len(import_failed) > 0: if len(import_failed) > 0:
logging.warning("WARNING: some comfy_extras/ nodes did not import correctly. This may be because they are missing some dependencies.\n") logging.warning("WARNING: some seap_extras/ nodes did not import correctly. This may be because they are missing some dependencies.\n")
for node in import_failed: for node in import_failed:
logging.warning("IMPORT FAILED: {}".format(node)) logging.warning("IMPORT FAILED: {}".format(node))
logging.warning("\nThis issue might be caused by new missing dependencies added the last time you updated ComfyUI.") logging.warning("\nThis issue might be caused by new missing dependencies added the last time you updated ComfyUI.")

View File

@ -15,8 +15,8 @@ from ..ldm.modules.diffusionmodules.openaimodel import UNetModel, TimestepEmbedS
from ..ldm.util import exists from ..ldm.util import exists
from .control_types import UNION_CONTROLNET_TYPES from .control_types import UNION_CONTROLNET_TYPES
from collections import OrderedDict from collections import OrderedDict
import comfy.ops import seap.ops
from comfy.ldm.modules.attention import optimized_attention from seap.ldm.modules.attention import optimized_attention
class OptimizedAttention(nn.Module): class OptimizedAttention(nn.Module):
def __init__(self, c, nhead, dropout=0.0, dtype=None, device=None, operations=None): def __init__(self, c, nhead, dropout=0.0, dtype=None, device=None, operations=None):
@ -95,7 +95,7 @@ class ControlNet(nn.Module):
attn_precision=None, attn_precision=None,
union_controlnet_num_control_type=None, union_controlnet_num_control_type=None,
device=None, device=None,
operations=comfy.ops.disable_weight_init, operations=seap.ops.disable_weight_init,
**kwargs, **kwargs,
): ):
super().__init__() super().__init__()

View File

@ -1,8 +1,8 @@
import torch import torch
from typing import Dict, Optional from typing import Dict, Optional
import comfy.ldm.modules.diffusionmodules.mmdit import seap.ldm.modules.diffusionmodules.mmdit
class ControlNet(comfy.ldm.modules.diffusionmodules.mmdit.MMDiT): class ControlNet(seap.ldm.modules.diffusionmodules.mmdit.MMDiT):
def __init__( def __init__(
self, self,
num_blocks = None, num_blocks = None,
@ -21,7 +21,7 @@ class ControlNet(comfy.ldm.modules.diffusionmodules.mmdit.MMDiT):
if control_latent_channels is None: if control_latent_channels is None:
control_latent_channels = self.in_channels control_latent_channels = self.in_channels
self.pos_embed_input = comfy.ldm.modules.diffusionmodules.mmdit.PatchEmbed( self.pos_embed_input = seap.ldm.modules.diffusionmodules.mmdit.PatchEmbed(
None, None,
self.patch_size, self.patch_size,
control_latent_channels, control_latent_channels,

View File

@ -2,7 +2,7 @@ import argparse
import enum import enum
import os import os
from typing import Optional from typing import Optional
import comfy.options import seap.options
class EnumAction(argparse.Action): class EnumAction(argparse.Action):
@ -173,7 +173,7 @@ parser.add_argument(
parser.add_argument("--user-directory", type=is_valid_directory, default=None, help="Set the ComfyUI user directory with an absolute path.") parser.add_argument("--user-directory", type=is_valid_directory, default=None, help="Set the ComfyUI user directory with an absolute path.")
if comfy.options.args_parsing: if seap.options.args_parsing:
args = parser.parse_args() args = parser.parse_args()
else: else:
args = parser.parse_args([]) args = parser.parse_args([])

View File

@ -1,6 +1,6 @@
import torch import torch
from comfy.ldm.modules.attention import optimized_attention_for_device from seap.ldm.modules.attention import optimized_attention_for_device
import comfy.ops import seap.ops
class CLIPAttention(torch.nn.Module): class CLIPAttention(torch.nn.Module):
def __init__(self, embed_dim, heads, dtype, device, operations): def __init__(self, embed_dim, heads, dtype, device, operations):
@ -78,7 +78,7 @@ class CLIPEmbeddings(torch.nn.Module):
self.position_embedding = operations.Embedding(num_positions, embed_dim, dtype=dtype, device=device) self.position_embedding = operations.Embedding(num_positions, embed_dim, dtype=dtype, device=device)
def forward(self, input_tokens, dtype=torch.float32): def forward(self, input_tokens, dtype=torch.float32):
return self.token_embedding(input_tokens, out_dtype=dtype) + comfy.ops.cast_to(self.position_embedding.weight, dtype=dtype, device=input_tokens.device) return self.token_embedding(input_tokens, out_dtype=dtype) + seap.ops.cast_to(self.position_embedding.weight, dtype=dtype, device=input_tokens.device)
class CLIPTextModel_(torch.nn.Module): class CLIPTextModel_(torch.nn.Module):
@ -159,7 +159,7 @@ class CLIPVisionEmbeddings(torch.nn.Module):
def forward(self, pixel_values): def forward(self, pixel_values):
embeds = self.patch_embedding(pixel_values).flatten(2).transpose(1, 2) embeds = self.patch_embedding(pixel_values).flatten(2).transpose(1, 2)
return torch.cat([comfy.ops.cast_to_input(self.class_embedding, embeds).expand(pixel_values.shape[0], 1, -1), embeds], dim=1) + comfy.ops.cast_to_input(self.position_embedding.weight, embeds) return torch.cat([seap.ops.cast_to_input(self.class_embedding, embeds).expand(pixel_values.shape[0], 1, -1), embeds], dim=1) + seap.ops.cast_to_input(self.position_embedding.weight, embeds)
class CLIPVision(torch.nn.Module): class CLIPVision(torch.nn.Module):

View File

@ -4,11 +4,11 @@ import torch
import json import json
import logging import logging
import comfy.ops import seap.ops
import comfy.model_patcher import seap.model_patcher
import comfy.model_management import seap.model_management
import comfy.utils import seap.utils
import comfy.clip_model import seap.clip_model
class Output: class Output:
def __getitem__(self, key): def __getitem__(self, key):
@ -35,13 +35,13 @@ class ClipVisionModel():
config = json.load(f) config = json.load(f)
self.image_size = config.get("image_size", 224) self.image_size = config.get("image_size", 224)
self.load_device = comfy.model_management.text_encoder_device() self.load_device = seap.model_management.text_encoder_device()
offload_device = comfy.model_management.text_encoder_offload_device() offload_device = seap.model_management.text_encoder_offload_device()
self.dtype = comfy.model_management.text_encoder_dtype(self.load_device) self.dtype = seap.model_management.text_encoder_dtype(self.load_device)
self.model = comfy.clip_model.CLIPVisionModelProjection(config, self.dtype, offload_device, comfy.ops.manual_cast) self.model = seap.clip_model.CLIPVisionModelProjection(config, self.dtype, offload_device, seap.ops.manual_cast)
self.model.eval() self.model.eval()
self.patcher = comfy.model_patcher.ModelPatcher(self.model, load_device=self.load_device, offload_device=offload_device) self.patcher = seap.model_patcher.ModelPatcher(self.model, load_device=self.load_device, offload_device=offload_device)
def load_sd(self, sd): def load_sd(self, sd):
return self.model.load_state_dict(sd, strict=False) return self.model.load_state_dict(sd, strict=False)
@ -50,14 +50,14 @@ class ClipVisionModel():
return self.model.state_dict() return self.model.state_dict()
def encode_image(self, image): def encode_image(self, image):
comfy.model_management.load_model_gpu(self.patcher) seap.model_management.load_model_gpu(self.patcher)
pixel_values = clip_preprocess(image.to(self.load_device), size=self.image_size).float() pixel_values = clip_preprocess(image.to(self.load_device), size=self.image_size).float()
out = self.model(pixel_values=pixel_values, intermediate_output=-2) out = self.model(pixel_values=pixel_values, intermediate_output=-2)
outputs = Output() outputs = Output()
outputs["last_hidden_state"] = out[0].to(comfy.model_management.intermediate_device()) outputs["last_hidden_state"] = out[0].to(seap.model_management.intermediate_device())
outputs["image_embeds"] = out[2].to(comfy.model_management.intermediate_device()) outputs["image_embeds"] = out[2].to(seap.model_management.intermediate_device())
outputs["penultimate_hidden_states"] = out[1].to(comfy.model_management.intermediate_device()) outputs["penultimate_hidden_states"] = out[1].to(seap.model_management.intermediate_device())
return outputs return outputs
def convert_to_transformers(sd, prefix): def convert_to_transformers(sd, prefix):

View File

@ -3,7 +3,7 @@ from typing import Callable, Protocol, TypedDict, Optional, List
class UnetApplyFunction(Protocol): class UnetApplyFunction(Protocol):
"""Function signature protocol on comfy.model_base.BaseModel.apply_model""" """Function signature protocol on seap.model_base.BaseModel.apply_model"""
def __call__(self, x: torch.Tensor, t: torch.Tensor, **kwargs) -> torch.Tensor: def __call__(self, x: torch.Tensor, t: torch.Tensor, **kwargs) -> torch.Tensor:
pass pass

View File

@ -1,6 +1,6 @@
import torch import torch
import math import math
import comfy.utils import seap.utils
def lcm(a, b): #TODO: eventually replace by math.lcm (added in python3.9) def lcm(a, b): #TODO: eventually replace by math.lcm (added in python3.9)
@ -14,7 +14,7 @@ class CONDRegular:
return self.__class__(cond) return self.__class__(cond)
def process_cond(self, batch_size, device, **kwargs): def process_cond(self, batch_size, device, **kwargs):
return self._copy_with(comfy.utils.repeat_to_batch_size(self.cond, batch_size).to(device)) return self._copy_with(seap.utils.repeat_to_batch_size(self.cond, batch_size).to(device))
def can_concat(self, other): def can_concat(self, other):
if self.cond.shape != other.cond.shape: if self.cond.shape != other.cond.shape:
@ -35,7 +35,7 @@ class CONDNoiseShape(CONDRegular):
for i in range(dims): for i in range(dims):
data = data.narrow(i + 2, area[i + dims], area[i]) data = data.narrow(i + 2, area[i + dims], area[i])
return self._copy_with(comfy.utils.repeat_to_batch_size(data, batch_size).to(device)) return self._copy_with(seap.utils.repeat_to_batch_size(data, batch_size).to(device))
class CONDCrossAttn(CONDRegular): class CONDCrossAttn(CONDRegular):

View File

@ -22,19 +22,19 @@ from enum import Enum
import math import math
import os import os
import logging import logging
import comfy.utils import seap.utils
import comfy.model_management import seap.model_management
import comfy.model_detection import seap.model_detection
import comfy.model_patcher import seap.model_patcher
import comfy.ops import seap.ops
import comfy.latent_formats import seap.latent_formats
import comfy.cldm.cldm import seap.cldm.cldm
import comfy.t2i_adapter.adapter import seap.t2i_adapter.adapter
import comfy.ldm.cascade.controlnet import seap.ldm.cascade.controlnet
import comfy.cldm.mmdit import seap.cldm.mmdit
import comfy.ldm.hydit.controlnet import seap.ldm.hydit.controlnet
import comfy.ldm.flux.controlnet import seap.ldm.flux.controlnet
def broadcast_image_to(tensor, target_batch_size, batched_number): def broadcast_image_to(tensor, target_batch_size, batched_number):
@ -74,7 +74,7 @@ class ControlBase:
self.extra_args = {} self.extra_args = {}
if device is None: if device is None:
device = comfy.model_management.get_torch_device() device = seap.model_management.get_torch_device()
self.device = device self.device = device
self.previous_controlnet = None self.previous_controlnet = None
self.extra_conds = [] self.extra_conds = []
@ -190,7 +190,7 @@ class ControlNet(ControlBase):
self.control_model = control_model self.control_model = control_model
self.load_device = load_device self.load_device = load_device
if control_model is not None: if control_model is not None:
self.control_model_wrapped = comfy.model_patcher.ModelPatcher(self.control_model, load_device=load_device, offload_device=comfy.model_management.unet_offload_device()) self.control_model_wrapped = seap.model_patcher.ModelPatcher(self.control_model, load_device=load_device, offload_device=seap.model_management.unet_offload_device())
self.compression_ratio = compression_ratio self.compression_ratio = compression_ratio
self.global_average_pooling = global_average_pooling self.global_average_pooling = global_average_pooling
@ -227,19 +227,19 @@ class ControlNet(ControlBase):
else: else:
if self.latent_format is not None: if self.latent_format is not None:
raise ValueError("This Controlnet needs a VAE but none was provided, please use a ControlNetApply node with a VAE input and connect it.") raise ValueError("This Controlnet needs a VAE but none was provided, please use a ControlNetApply node with a VAE input and connect it.")
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * compression_ratio, x_noisy.shape[2] * compression_ratio, self.upscale_algorithm, "center") self.cond_hint = seap.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * compression_ratio, x_noisy.shape[2] * compression_ratio, self.upscale_algorithm, "center")
if self.vae is not None: if self.vae is not None:
loaded_models = comfy.model_management.loaded_models(only_currently_used=True) loaded_models = seap.model_management.loaded_models(only_currently_used=True)
self.cond_hint = self.vae.encode(self.cond_hint.movedim(1, -1)) self.cond_hint = self.vae.encode(self.cond_hint.movedim(1, -1))
comfy.model_management.load_models_gpu(loaded_models) seap.model_management.load_models_gpu(loaded_models)
if self.latent_format is not None: if self.latent_format is not None:
self.cond_hint = self.latent_format.process_in(self.cond_hint) self.cond_hint = self.latent_format.process_in(self.cond_hint)
if len(self.extra_concat_orig) > 0: if len(self.extra_concat_orig) > 0:
to_concat = [] to_concat = []
for c in self.extra_concat_orig: for c in self.extra_concat_orig:
c = c.to(self.cond_hint.device) c = c.to(self.cond_hint.device)
c = comfy.utils.common_upscale(c, self.cond_hint.shape[3], self.cond_hint.shape[2], self.upscale_algorithm, "center") c = seap.utils.common_upscale(c, self.cond_hint.shape[3], self.cond_hint.shape[2], self.upscale_algorithm, "center")
to_concat.append(comfy.utils.repeat_to_batch_size(c, self.cond_hint.shape[0])) to_concat.append(seap.utils.repeat_to_batch_size(c, self.cond_hint.shape[0]))
self.cond_hint = torch.cat([self.cond_hint] + to_concat, dim=1) self.cond_hint = torch.cat([self.cond_hint] + to_concat, dim=1)
self.cond_hint = self.cond_hint.to(device=self.device, dtype=dtype) self.cond_hint = self.cond_hint.to(device=self.device, dtype=dtype)
@ -280,7 +280,7 @@ class ControlNet(ControlBase):
super().cleanup() super().cleanup()
class ControlLoraOps: class ControlLoraOps:
class Linear(torch.nn.Module, comfy.ops.CastWeightBiasOp): class Linear(torch.nn.Module, seap.ops.CastWeightBiasOp):
def __init__(self, in_features: int, out_features: int, bias: bool = True, def __init__(self, in_features: int, out_features: int, bias: bool = True,
device=None, dtype=None) -> None: device=None, dtype=None) -> None:
factory_kwargs = {'device': device, 'dtype': dtype} factory_kwargs = {'device': device, 'dtype': dtype}
@ -293,13 +293,13 @@ class ControlLoraOps:
self.bias = None self.bias = None
def forward(self, input): def forward(self, input):
weight, bias = comfy.ops.cast_bias_weight(self, input) weight, bias = seap.ops.cast_bias_weight(self, input)
if self.up is not None: if self.up is not None:
return torch.nn.functional.linear(input, weight + (torch.mm(self.up.flatten(start_dim=1), self.down.flatten(start_dim=1))).reshape(self.weight.shape).type(input.dtype), bias) return torch.nn.functional.linear(input, weight + (torch.mm(self.up.flatten(start_dim=1), self.down.flatten(start_dim=1))).reshape(self.weight.shape).type(input.dtype), bias)
else: else:
return torch.nn.functional.linear(input, weight, bias) return torch.nn.functional.linear(input, weight, bias)
class Conv2d(torch.nn.Module, comfy.ops.CastWeightBiasOp): class Conv2d(torch.nn.Module, seap.ops.CastWeightBiasOp):
def __init__( def __init__(
self, self,
in_channels, in_channels,
@ -333,7 +333,7 @@ class ControlLoraOps:
def forward(self, input): def forward(self, input):
weight, bias = comfy.ops.cast_bias_weight(self, input) weight, bias = seap.ops.cast_bias_weight(self, input)
if self.up is not None: if self.up is not None:
return torch.nn.functional.conv2d(input, weight + (torch.mm(self.up.flatten(start_dim=1), self.down.flatten(start_dim=1))).reshape(self.weight.shape).type(input.dtype), bias, self.stride, self.padding, self.dilation, self.groups) return torch.nn.functional.conv2d(input, weight + (torch.mm(self.up.flatten(start_dim=1), self.down.flatten(start_dim=1))).reshape(self.weight.shape).type(input.dtype), bias, self.stride, self.padding, self.dilation, self.groups)
else: else:
@ -355,17 +355,17 @@ class ControlLora(ControlNet):
self.manual_cast_dtype = model.manual_cast_dtype self.manual_cast_dtype = model.manual_cast_dtype
dtype = model.get_dtype() dtype = model.get_dtype()
if self.manual_cast_dtype is None: if self.manual_cast_dtype is None:
class control_lora_ops(ControlLoraOps, comfy.ops.disable_weight_init): class control_lora_ops(ControlLoraOps, seap.ops.disable_weight_init):
pass pass
else: else:
class control_lora_ops(ControlLoraOps, comfy.ops.manual_cast): class control_lora_ops(ControlLoraOps, seap.ops.manual_cast):
pass pass
dtype = self.manual_cast_dtype dtype = self.manual_cast_dtype
controlnet_config["operations"] = control_lora_ops controlnet_config["operations"] = control_lora_ops
controlnet_config["dtype"] = dtype controlnet_config["dtype"] = dtype
self.control_model = comfy.cldm.cldm.ControlNet(**controlnet_config) self.control_model = seap.cldm.cldm.ControlNet(**controlnet_config)
self.control_model.to(comfy.model_management.get_torch_device()) self.control_model.to(seap.model_management.get_torch_device())
diffusion_model = model.diffusion_model diffusion_model = model.diffusion_model
sd = diffusion_model.state_dict() sd = diffusion_model.state_dict()
cm = self.control_model.state_dict() cm = self.control_model.state_dict()
@ -373,13 +373,13 @@ class ControlLora(ControlNet):
for k in sd: for k in sd:
weight = sd[k] weight = sd[k]
try: try:
comfy.utils.set_attr_param(self.control_model, k, weight) seap.utils.set_attr_param(self.control_model, k, weight)
except: except:
pass pass
for k in self.control_weights: for k in self.control_weights:
if k not in {"lora_controlnet"}: if k not in {"lora_controlnet"}:
comfy.utils.set_attr_param(self.control_model, k, self.control_weights[k].to(dtype).to(comfy.model_management.get_torch_device())) seap.utils.set_attr_param(self.control_model, k, self.control_weights[k].to(dtype).to(seap.model_management.get_torch_device()))
def copy(self): def copy(self):
c = ControlLora(self.control_weights, global_average_pooling=self.global_average_pooling) c = ControlLora(self.control_weights, global_average_pooling=self.global_average_pooling)
@ -396,29 +396,29 @@ class ControlLora(ControlNet):
return out return out
def inference_memory_requirements(self, dtype): def inference_memory_requirements(self, dtype):
return comfy.utils.calculate_parameters(self.control_weights) * comfy.model_management.dtype_size(dtype) + ControlBase.inference_memory_requirements(self, dtype) return seap.utils.calculate_parameters(self.control_weights) * seap.model_management.dtype_size(dtype) + ControlBase.inference_memory_requirements(self, dtype)
def controlnet_config(sd, model_options={}): def controlnet_config(sd, model_options={}):
model_config = comfy.model_detection.model_config_from_unet(sd, "", True) model_config = seap.model_detection.model_config_from_unet(sd, "", True)
unet_dtype = model_options.get("dtype", None) unet_dtype = model_options.get("dtype", None)
if unet_dtype is None: if unet_dtype is None:
weight_dtype = comfy.utils.weight_dtype(sd) weight_dtype = seap.utils.weight_dtype(sd)
supported_inference_dtypes = list(model_config.supported_inference_dtypes) supported_inference_dtypes = list(model_config.supported_inference_dtypes)
if weight_dtype is not None: if weight_dtype is not None:
supported_inference_dtypes.append(weight_dtype) supported_inference_dtypes.append(weight_dtype)
unet_dtype = comfy.model_management.unet_dtype(model_params=-1, supported_dtypes=supported_inference_dtypes) unet_dtype = seap.model_management.unet_dtype(model_params=-1, supported_dtypes=supported_inference_dtypes)
load_device = comfy.model_management.get_torch_device() load_device = seap.model_management.get_torch_device()
manual_cast_dtype = comfy.model_management.unet_manual_cast(unet_dtype, load_device) manual_cast_dtype = seap.model_management.unet_manual_cast(unet_dtype, load_device)
operations = model_options.get("custom_operations", None) operations = model_options.get("custom_operations", None)
if operations is None: if operations is None:
operations = comfy.ops.pick_operations(unet_dtype, manual_cast_dtype, disable_fast_fp8=True) operations = seap.ops.pick_operations(unet_dtype, manual_cast_dtype, disable_fast_fp8=True)
offload_device = comfy.model_management.unet_offload_device() offload_device = seap.model_management.unet_offload_device()
return model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device return model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device
def controlnet_load_state_dict(control_model, sd): def controlnet_load_state_dict(control_model, sd):
@ -432,9 +432,9 @@ def controlnet_load_state_dict(control_model, sd):
return control_model return control_model
def load_controlnet_mmdit(sd, model_options={}): def load_controlnet_mmdit(sd, model_options={}):
new_sd = comfy.model_detection.convert_diffusers_mmdit(sd, "") new_sd = seap.model_detection.convert_diffusers_mmdit(sd, "")
model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device = controlnet_config(new_sd, model_options=model_options) model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device = controlnet_config(new_sd, model_options=model_options)
num_blocks = comfy.model_detection.count_blocks(new_sd, 'joint_blocks.{}.') num_blocks = seap.model_detection.count_blocks(new_sd, 'joint_blocks.{}.')
for k in sd: for k in sd:
new_sd[k] = sd[k] new_sd[k] = sd[k]
@ -443,10 +443,10 @@ def load_controlnet_mmdit(sd, model_options={}):
if control_latent_channels == 17: #inpaint controlnet if control_latent_channels == 17: #inpaint controlnet
concat_mask = True concat_mask = True
control_model = comfy.cldm.mmdit.ControlNet(num_blocks=num_blocks, control_latent_channels=control_latent_channels, operations=operations, device=offload_device, dtype=unet_dtype, **model_config.unet_config) control_model = seap.cldm.mmdit.ControlNet(num_blocks=num_blocks, control_latent_channels=control_latent_channels, operations=operations, device=offload_device, dtype=unet_dtype, **model_config.unet_config)
control_model = controlnet_load_state_dict(control_model, new_sd) control_model = controlnet_load_state_dict(control_model, new_sd)
latent_format = comfy.latent_formats.SD3() latent_format = seap.latent_formats.SD3()
latent_format.shift_factor = 0 #SD3 controlnet weirdness latent_format.shift_factor = 0 #SD3 controlnet weirdness
control = ControlNet(control_model, compression_ratio=1, latent_format=latent_format, concat_mask=concat_mask, load_device=load_device, manual_cast_dtype=manual_cast_dtype) control = ControlNet(control_model, compression_ratio=1, latent_format=latent_format, concat_mask=concat_mask, load_device=load_device, manual_cast_dtype=manual_cast_dtype)
return control return control
@ -455,24 +455,24 @@ def load_controlnet_mmdit(sd, model_options={}):
def load_controlnet_hunyuandit(controlnet_data, model_options={}): def load_controlnet_hunyuandit(controlnet_data, model_options={}):
model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device = controlnet_config(controlnet_data, model_options=model_options) model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device = controlnet_config(controlnet_data, model_options=model_options)
control_model = comfy.ldm.hydit.controlnet.HunYuanControlNet(operations=operations, device=offload_device, dtype=unet_dtype) control_model = seap.ldm.hydit.controlnet.HunYuanControlNet(operations=operations, device=offload_device, dtype=unet_dtype)
control_model = controlnet_load_state_dict(control_model, controlnet_data) control_model = controlnet_load_state_dict(control_model, controlnet_data)
latent_format = comfy.latent_formats.SDXL() latent_format = seap.latent_formats.SDXL()
extra_conds = ['text_embedding_mask', 'encoder_hidden_states_t5', 'text_embedding_mask_t5', 'image_meta_size', 'style', 'cos_cis_img', 'sin_cis_img'] extra_conds = ['text_embedding_mask', 'encoder_hidden_states_t5', 'text_embedding_mask_t5', 'image_meta_size', 'style', 'cos_cis_img', 'sin_cis_img']
control = ControlNet(control_model, compression_ratio=1, latent_format=latent_format, load_device=load_device, manual_cast_dtype=manual_cast_dtype, extra_conds=extra_conds, strength_type=StrengthType.CONSTANT) control = ControlNet(control_model, compression_ratio=1, latent_format=latent_format, load_device=load_device, manual_cast_dtype=manual_cast_dtype, extra_conds=extra_conds, strength_type=StrengthType.CONSTANT)
return control return control
def load_controlnet_flux_xlabs_mistoline(sd, mistoline=False, model_options={}): def load_controlnet_flux_xlabs_mistoline(sd, mistoline=False, model_options={}):
model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device = controlnet_config(sd, model_options=model_options) model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device = controlnet_config(sd, model_options=model_options)
control_model = comfy.ldm.flux.controlnet.ControlNetFlux(mistoline=mistoline, operations=operations, device=offload_device, dtype=unet_dtype, **model_config.unet_config) control_model = seap.ldm.flux.controlnet.ControlNetFlux(mistoline=mistoline, operations=operations, device=offload_device, dtype=unet_dtype, **model_config.unet_config)
control_model = controlnet_load_state_dict(control_model, sd) control_model = controlnet_load_state_dict(control_model, sd)
extra_conds = ['y', 'guidance'] extra_conds = ['y', 'guidance']
control = ControlNet(control_model, load_device=load_device, manual_cast_dtype=manual_cast_dtype, extra_conds=extra_conds) control = ControlNet(control_model, load_device=load_device, manual_cast_dtype=manual_cast_dtype, extra_conds=extra_conds)
return control return control
def load_controlnet_flux_instantx(sd, model_options={}): def load_controlnet_flux_instantx(sd, model_options={}):
new_sd = comfy.model_detection.convert_diffusers_mmdit(sd, "") new_sd = seap.model_detection.convert_diffusers_mmdit(sd, "")
model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device = controlnet_config(new_sd, model_options=model_options) model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device = controlnet_config(new_sd, model_options=model_options)
for k in sd: for k in sd:
new_sd[k] = sd[k] new_sd[k] = sd[k]
@ -487,16 +487,16 @@ def load_controlnet_flux_instantx(sd, model_options={}):
if control_latent_channels == 17: if control_latent_channels == 17:
concat_mask = True concat_mask = True
control_model = comfy.ldm.flux.controlnet.ControlNetFlux(latent_input=True, num_union_modes=num_union_modes, control_latent_channels=control_latent_channels, operations=operations, device=offload_device, dtype=unet_dtype, **model_config.unet_config) control_model = seap.ldm.flux.controlnet.ControlNetFlux(latent_input=True, num_union_modes=num_union_modes, control_latent_channels=control_latent_channels, operations=operations, device=offload_device, dtype=unet_dtype, **model_config.unet_config)
control_model = controlnet_load_state_dict(control_model, new_sd) control_model = controlnet_load_state_dict(control_model, new_sd)
latent_format = comfy.latent_formats.Flux() latent_format = seap.latent_formats.Flux()
extra_conds = ['y', 'guidance'] extra_conds = ['y', 'guidance']
control = ControlNet(control_model, compression_ratio=1, latent_format=latent_format, concat_mask=concat_mask, load_device=load_device, manual_cast_dtype=manual_cast_dtype, extra_conds=extra_conds) control = ControlNet(control_model, compression_ratio=1, latent_format=latent_format, concat_mask=concat_mask, load_device=load_device, manual_cast_dtype=manual_cast_dtype, extra_conds=extra_conds)
return control return control
def convert_mistoline(sd): def convert_mistoline(sd):
return comfy.utils.state_dict_prefix_replace(sd, {"single_controlnet_blocks.": "controlnet_single_blocks."}) return seap.utils.state_dict_prefix_replace(sd, {"single_controlnet_blocks.": "controlnet_single_blocks."})
def load_controlnet_state_dict(state_dict, model=None, model_options={}): def load_controlnet_state_dict(state_dict, model=None, model_options={}):
@ -511,8 +511,8 @@ def load_controlnet_state_dict(state_dict, model=None, model_options={}):
supported_inference_dtypes = None supported_inference_dtypes = None
if "controlnet_cond_embedding.conv_in.weight" in controlnet_data: #diffusers format if "controlnet_cond_embedding.conv_in.weight" in controlnet_data: #diffusers format
controlnet_config = comfy.model_detection.unet_config_from_diffusers_unet(controlnet_data) controlnet_config = seap.model_detection.unet_config_from_diffusers_unet(controlnet_data)
diffusers_keys = comfy.utils.unet_to_diffusers(controlnet_config) diffusers_keys = seap.utils.unet_to_diffusers(controlnet_config)
diffusers_keys["controlnet_mid_block.weight"] = "middle_block_out.0.weight" diffusers_keys["controlnet_mid_block.weight"] = "middle_block_out.0.weight"
diffusers_keys["controlnet_mid_block.bias"] = "middle_block_out.0.bias" diffusers_keys["controlnet_mid_block.bias"] = "middle_block_out.0.bias"
@ -586,40 +586,40 @@ def load_controlnet_state_dict(state_dict, model=None, model_options={}):
return net return net
if controlnet_config is None: if controlnet_config is None:
model_config = comfy.model_detection.model_config_from_unet(controlnet_data, prefix, True) model_config = seap.model_detection.model_config_from_unet(controlnet_data, prefix, True)
supported_inference_dtypes = list(model_config.supported_inference_dtypes) supported_inference_dtypes = list(model_config.supported_inference_dtypes)
controlnet_config = model_config.unet_config controlnet_config = model_config.unet_config
unet_dtype = model_options.get("dtype", None) unet_dtype = model_options.get("dtype", None)
if unet_dtype is None: if unet_dtype is None:
weight_dtype = comfy.utils.weight_dtype(controlnet_data) weight_dtype = seap.utils.weight_dtype(controlnet_data)
if supported_inference_dtypes is None: if supported_inference_dtypes is None:
supported_inference_dtypes = [comfy.model_management.unet_dtype()] supported_inference_dtypes = [seap.model_management.unet_dtype()]
if weight_dtype is not None: if weight_dtype is not None:
supported_inference_dtypes.append(weight_dtype) supported_inference_dtypes.append(weight_dtype)
unet_dtype = comfy.model_management.unet_dtype(model_params=-1, supported_dtypes=supported_inference_dtypes) unet_dtype = seap.model_management.unet_dtype(model_params=-1, supported_dtypes=supported_inference_dtypes)
load_device = comfy.model_management.get_torch_device() load_device = seap.model_management.get_torch_device()
manual_cast_dtype = comfy.model_management.unet_manual_cast(unet_dtype, load_device) manual_cast_dtype = seap.model_management.unet_manual_cast(unet_dtype, load_device)
operations = model_options.get("custom_operations", None) operations = model_options.get("custom_operations", None)
if operations is None: if operations is None:
operations = comfy.ops.pick_operations(unet_dtype, manual_cast_dtype) operations = seap.ops.pick_operations(unet_dtype, manual_cast_dtype)
controlnet_config["operations"] = operations controlnet_config["operations"] = operations
controlnet_config["dtype"] = unet_dtype controlnet_config["dtype"] = unet_dtype
controlnet_config["device"] = comfy.model_management.unet_offload_device() controlnet_config["device"] = seap.model_management.unet_offload_device()
controlnet_config.pop("out_channels") controlnet_config.pop("out_channels")
controlnet_config["hint_channels"] = controlnet_data["{}input_hint_block.0.weight".format(prefix)].shape[1] controlnet_config["hint_channels"] = controlnet_data["{}input_hint_block.0.weight".format(prefix)].shape[1]
control_model = comfy.cldm.cldm.ControlNet(**controlnet_config) control_model = seap.cldm.cldm.ControlNet(**controlnet_config)
if pth: if pth:
if 'difference' in controlnet_data: if 'difference' in controlnet_data:
if model is not None: if model is not None:
comfy.model_management.load_models_gpu([model]) seap.model_management.load_models_gpu([model])
model_sd = model.model_state_dict() model_sd = model.model_state_dict()
for x in controlnet_data: for x in controlnet_data:
c_m = "control_model." c_m = "control_model."
@ -655,7 +655,7 @@ def load_controlnet(ckpt_path, model=None, model_options={}):
if filename.endswith("_shuffle") or filename.endswith("_shuffle_fp16"): #TODO: smarter way of enabling global_average_pooling if filename.endswith("_shuffle") or filename.endswith("_shuffle_fp16"): #TODO: smarter way of enabling global_average_pooling
model_options["global_average_pooling"] = True model_options["global_average_pooling"] = True
cnet = load_controlnet_state_dict(comfy.utils.load_torch_file(ckpt_path, safe_load=True), model=model, model_options=model_options) cnet = load_controlnet_state_dict(seap.utils.load_torch_file(ckpt_path, safe_load=True), model=model, model_options=model_options)
if cnet is None: if cnet is None:
logging.error("error checkpoint does not contain controlnet or t2i adapter data {}".format(ckpt_path)) logging.error("error checkpoint does not contain controlnet or t2i adapter data {}".format(ckpt_path))
return cnet return cnet
@ -693,7 +693,7 @@ class T2IAdapter(ControlBase):
self.control_input = None self.control_input = None
self.cond_hint = None self.cond_hint = None
width, height = self.scale_image_to(x_noisy.shape[3] * self.compression_ratio, x_noisy.shape[2] * self.compression_ratio) width, height = self.scale_image_to(x_noisy.shape[3] * self.compression_ratio, x_noisy.shape[2] * self.compression_ratio)
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, width, height, self.upscale_algorithm, "center").float().to(self.device) self.cond_hint = seap.utils.common_upscale(self.cond_hint_original, width, height, self.upscale_algorithm, "center").float().to(self.device)
if self.channels_in == 1 and self.cond_hint.shape[1] > 1: if self.channels_in == 1 and self.cond_hint.shape[1] > 1:
self.cond_hint = torch.mean(self.cond_hint, 1, keepdim=True) self.cond_hint = torch.mean(self.cond_hint, 1, keepdim=True)
if x_noisy.shape[0] != self.cond_hint.shape[0]: if x_noisy.shape[0] != self.cond_hint.shape[0]:
@ -728,12 +728,12 @@ def load_t2i_adapter(t2i_data, model_options={}): #TODO: model_options
prefix_replace["adapter.body.{}.resnets.{}.".format(i, j)] = "body.{}.".format(i * 2 + j) prefix_replace["adapter.body.{}.resnets.{}.".format(i, j)] = "body.{}.".format(i * 2 + j)
prefix_replace["adapter.body.{}.".format(i, j)] = "body.{}.".format(i * 2) prefix_replace["adapter.body.{}.".format(i, j)] = "body.{}.".format(i * 2)
prefix_replace["adapter."] = "" prefix_replace["adapter."] = ""
t2i_data = comfy.utils.state_dict_prefix_replace(t2i_data, prefix_replace) t2i_data = seap.utils.state_dict_prefix_replace(t2i_data, prefix_replace)
keys = t2i_data.keys() keys = t2i_data.keys()
if "body.0.in_conv.weight" in keys: if "body.0.in_conv.weight" in keys:
cin = t2i_data['body.0.in_conv.weight'].shape[1] cin = t2i_data['body.0.in_conv.weight'].shape[1]
model_ad = comfy.t2i_adapter.adapter.Adapter_light(cin=cin, channels=[320, 640, 1280, 1280], nums_rb=4) model_ad = seap.t2i_adapter.adapter.Adapter_light(cin=cin, channels=[320, 640, 1280, 1280], nums_rb=4)
elif 'conv_in.weight' in keys: elif 'conv_in.weight' in keys:
cin = t2i_data['conv_in.weight'].shape[1] cin = t2i_data['conv_in.weight'].shape[1]
channel = t2i_data['conv_in.weight'].shape[0] channel = t2i_data['conv_in.weight'].shape[0]
@ -745,13 +745,13 @@ def load_t2i_adapter(t2i_data, model_options={}): #TODO: model_options
xl = False xl = False
if cin == 256 or cin == 768: if cin == 256 or cin == 768:
xl = True xl = True
model_ad = comfy.t2i_adapter.adapter.Adapter(cin=cin, channels=[channel, channel*2, channel*4, channel*4][:4], nums_rb=2, ksize=ksize, sk=True, use_conv=use_conv, xl=xl) model_ad = seap.t2i_adapter.adapter.Adapter(cin=cin, channels=[channel, channel * 2, channel * 4, channel * 4][:4], nums_rb=2, ksize=ksize, sk=True, use_conv=use_conv, xl=xl)
elif "backbone.0.0.weight" in keys: elif "backbone.0.0.weight" in keys:
model_ad = comfy.ldm.cascade.controlnet.ControlNet(c_in=t2i_data['backbone.0.0.weight'].shape[1], proj_blocks=[0, 4, 8, 12, 51, 55, 59, 63]) model_ad = seap.ldm.cascade.controlnet.ControlNet(c_in=t2i_data['backbone.0.0.weight'].shape[1], proj_blocks=[0, 4, 8, 12, 51, 55, 59, 63])
compression_ratio = 32 compression_ratio = 32
upscale_algorithm = 'bilinear' upscale_algorithm = 'bilinear'
elif "backbone.10.blocks.0.weight" in keys: elif "backbone.10.blocks.0.weight" in keys:
model_ad = comfy.ldm.cascade.controlnet.ControlNet(c_in=t2i_data['backbone.0.weight'].shape[1], bottleneck_mode="large", proj_blocks=[0, 4, 8, 12, 51, 55, 59, 63]) model_ad = seap.ldm.cascade.controlnet.ControlNet(c_in=t2i_data['backbone.0.weight'].shape[1], bottleneck_mode="large", proj_blocks=[0, 4, 8, 12, 51, 55, 59, 63])
compression_ratio = 1 compression_ratio = 1
upscale_algorithm = 'nearest-exact' upscale_algorithm = 'nearest-exact'
else: else:

View File

@ -1,6 +1,6 @@
import os import os
import comfy.sd import seap.sd
def first_file(path, filenames): def first_file(path, filenames):
for f in filenames: for f in filenames:
@ -22,15 +22,15 @@ def load_diffusers(model_path, output_vae=True, output_clip=True, embedding_dire
if text_encoder2_path is not None: if text_encoder2_path is not None:
text_encoder_paths.append(text_encoder2_path) text_encoder_paths.append(text_encoder2_path)
unet = comfy.sd.load_diffusion_model(unet_path) unet = seap.sd.load_diffusion_model(unet_path)
clip = None clip = None
if output_clip: if output_clip:
clip = comfy.sd.load_clip(text_encoder_paths, embedding_directory=embedding_directory) clip = seap.sd.load_clip(text_encoder_paths, embedding_directory=embedding_directory)
vae = None vae = None
if output_vae: if output_vae:
sd = comfy.utils.load_torch_file(vae_path) sd = seap.utils.load_torch_file(vae_path)
vae = comfy.sd.VAE(sd=sd) vae = seap.sd.VAE(sd=sd)
return (unet, clip, vae) return (unet, clip, vae)

View File

@ -2,8 +2,8 @@ import torch
from torch import nn from torch import nn
from .ldm.modules.attention import CrossAttention from .ldm.modules.attention import CrossAttention
from inspect import isfunction from inspect import isfunction
import comfy.ops import seap.ops
ops = comfy.ops.manual_cast ops = seap.ops.manual_cast
def exists(val): def exists(val):
return val is not None return val is not None

View File

@ -8,8 +8,8 @@ from tqdm.auto import trange, tqdm
from . import utils from . import utils
from . import deis from . import deis
import comfy.model_patcher import seap.model_patcher
import comfy.model_sampling import seap.model_sampling
def append_zero(x): def append_zero(x):
return torch.cat([x, x.new_zeros([1])]) return torch.cat([x, x.new_zeros([1])])
@ -521,7 +521,7 @@ def sample_dpm_adaptive(model, x, sigma_min, sigma_max, extra_args=None, callbac
@torch.no_grad() @torch.no_grad()
def sample_dpmpp_2s_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None): def sample_dpmpp_2s_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST): if isinstance(model.inner_model.inner_model.model_sampling, seap.model_sampling.CONST):
return sample_dpmpp_2s_ancestral_RF(model, x, sigmas, extra_args, callback, disable, eta, s_noise, noise_sampler) return sample_dpmpp_2s_ancestral_RF(model, x, sigmas, extra_args, callback, disable, eta, s_noise, noise_sampler)
"""Ancestral sampling with DPM-Solver++(2S) second-order steps.""" """Ancestral sampling with DPM-Solver++(2S) second-order steps."""
@ -1071,7 +1071,7 @@ def sample_euler_cfg_pp(model, x, sigmas, extra_args=None, callback=None, disabl
return args["denoised"] return args["denoised"]
model_options = extra_args.get("model_options", {}).copy() model_options = extra_args.get("model_options", {}).copy()
extra_args["model_options"] = comfy.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True) extra_args["model_options"] = seap.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True)
s_in = x.new_ones([x.shape[0]]) s_in = x.new_ones([x.shape[0]])
for i in trange(len(sigmas) - 1, disable=disable): for i in trange(len(sigmas) - 1, disable=disable):
@ -1096,7 +1096,7 @@ def sample_euler_ancestral_cfg_pp(model, x, sigmas, extra_args=None, callback=No
return args["denoised"] return args["denoised"]
model_options = extra_args.get("model_options", {}).copy() model_options = extra_args.get("model_options", {}).copy()
extra_args["model_options"] = comfy.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True) extra_args["model_options"] = seap.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True)
s_in = x.new_ones([x.shape[0]]) s_in = x.new_ones([x.shape[0]])
for i in trange(len(sigmas) - 1, disable=disable): for i in trange(len(sigmas) - 1, disable=disable):
@ -1122,7 +1122,7 @@ def sample_dpmpp_2s_ancestral_cfg_pp(model, x, sigmas, extra_args=None, callback
return args["denoised"] return args["denoised"]
model_options = extra_args.get("model_options", {}).copy() model_options = extra_args.get("model_options", {}).copy()
extra_args["model_options"] = comfy.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True) extra_args["model_options"] = seap.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True)
s_in = x.new_ones([x.shape[0]]) s_in = x.new_ones([x.shape[0]])
sigma_fn = lambda t: t.neg().exp() sigma_fn = lambda t: t.neg().exp()
@ -1167,7 +1167,7 @@ def sample_dpmpp_2m_cfg_pp(model, x, sigmas, extra_args=None, callback=None, dis
return args["denoised"] return args["denoised"]
model_options = extra_args.get("model_options", {}).copy() model_options = extra_args.get("model_options", {}).copy()
extra_args["model_options"] = comfy.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True) extra_args["model_options"] = seap.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True)
for i in trange(len(sigmas) - 1, disable=disable): for i in trange(len(sigmas) - 1, disable=disable):
denoised = model(x, sigmas[i] * s_in, **extra_args) denoised = model(x, sigmas[i] * s_in, **extra_args)

View File

@ -4,8 +4,8 @@ import torch
from torch import nn from torch import nn
from typing import Literal, Dict, Any from typing import Literal, Dict, Any
import math import math
import comfy.ops import seap.ops
ops = comfy.ops.disable_weight_init ops = seap.ops.disable_weight_init
def vae_sample(mean, scale): def vae_sample(mean, scale):
stdev = nn.functional.softplus(scale) + 1e-4 stdev = nn.functional.softplus(scale) + 1e-4

View File

@ -1,6 +1,6 @@
# code adapted from: https://github.com/Stability-AI/stable-audio-tools # code adapted from: https://github.com/Stability-AI/stable-audio-tools
from comfy.ldm.modules.attention import optimized_attention from seap.ldm.modules.attention import optimized_attention
import typing as tp import typing as tp
import torch import torch
@ -9,7 +9,7 @@ from einops import rearrange
from torch import nn from torch import nn
from torch.nn import functional as F from torch.nn import functional as F
import math import math
import comfy.ops import seap.ops
class FourierFeatures(nn.Module): class FourierFeatures(nn.Module):
def __init__(self, in_features, out_features, std=1., dtype=None, device=None): def __init__(self, in_features, out_features, std=1., dtype=None, device=None):
@ -19,7 +19,7 @@ class FourierFeatures(nn.Module):
[out_features // 2, in_features], dtype=dtype, device=device)) [out_features // 2, in_features], dtype=dtype, device=device))
def forward(self, input): def forward(self, input):
f = 2 * math.pi * input @ comfy.ops.cast_to_input(self.weight.T, input) f = 2 * math.pi * input @ seap.ops.cast_to_input(self.weight.T, input)
return torch.cat([f.cos(), f.sin()], dim=-1) return torch.cat([f.cos(), f.sin()], dim=-1)
# norms # norms
@ -40,8 +40,8 @@ class LayerNorm(nn.Module):
def forward(self, x): def forward(self, x):
beta = self.beta beta = self.beta
if beta is not None: if beta is not None:
beta = comfy.ops.cast_to_input(beta, x) beta = seap.ops.cast_to_input(beta, x)
return F.layer_norm(x, x.shape[-1:], weight=comfy.ops.cast_to_input(self.gamma, x), bias=beta) return F.layer_norm(x, x.shape[-1:], weight=seap.ops.cast_to_input(self.gamma, x), bias=beta)
class GLU(nn.Module): class GLU(nn.Module):
def __init__( def __init__(
@ -164,14 +164,14 @@ class RotaryEmbedding(nn.Module):
t = t / self.interpolation_factor t = t / self.interpolation_factor
freqs = torch.einsum('i , j -> i j', t, comfy.ops.cast_to_input(self.inv_freq, t)) freqs = torch.einsum('i , j -> i j', t, seap.ops.cast_to_input(self.inv_freq, t))
freqs = torch.cat((freqs, freqs), dim = -1) freqs = torch.cat((freqs, freqs), dim = -1)
if self.scale is None: if self.scale is None:
return freqs, 1. return freqs, 1.
power = (torch.arange(seq_len, device = device) - (seq_len // 2)) / self.scale_base power = (torch.arange(seq_len, device = device) - (seq_len // 2)) / self.scale_base
scale = comfy.ops.cast_to_input(self.scale, t) ** rearrange(power, 'n -> n 1') scale = seap.ops.cast_to_input(self.scale, t) ** rearrange(power, 'n -> n 1')
scale = torch.cat((scale, scale), dim = -1) scale = torch.cat((scale, scale), dim = -1)
return freqs, scale return freqs, scale

View File

@ -6,7 +6,7 @@ from torch import Tensor, einsum
from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple, TypeVar, Union from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple, TypeVar, Union
from einops import rearrange from einops import rearrange
import math import math
import comfy.ops import seap.ops
class LearnedPositionalEmbedding(nn.Module): class LearnedPositionalEmbedding(nn.Module):
"""Used for continuous time""" """Used for continuous time"""
@ -27,7 +27,7 @@ class LearnedPositionalEmbedding(nn.Module):
def TimePositionalEmbedding(dim: int, out_features: int) -> nn.Module: def TimePositionalEmbedding(dim: int, out_features: int) -> nn.Module:
return nn.Sequential( return nn.Sequential(
LearnedPositionalEmbedding(dim), LearnedPositionalEmbedding(dim),
comfy.ops.manual_cast.Linear(in_features=dim + 1, out_features=out_features), seap.ops.manual_cast.Linear(in_features=dim + 1, out_features=out_features),
) )

View File

@ -7,9 +7,9 @@ import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from comfy.ldm.modules.attention import optimized_attention from seap.ldm.modules.attention import optimized_attention
import comfy.ops import seap.ops
import comfy.ldm.common_dit import seap.ldm.common_dit
def modulate(x, shift, scale): def modulate(x, shift, scale):
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1) return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
@ -408,7 +408,7 @@ class MMDiT(nn.Module):
def patchify(self, x): def patchify(self, x):
B, C, H, W = x.size() B, C, H, W = x.size()
x = comfy.ldm.common_dit.pad_to_patch_size(x, (self.patch_size, self.patch_size)) x = seap.ldm.common_dit.pad_to_patch_size(x, (self.patch_size, self.patch_size))
x = x.view( x = x.view(
B, B,
C, C,
@ -426,7 +426,7 @@ class MMDiT(nn.Module):
max_dim = max(h, w) max_dim = max(h, w)
cur_dim = self.h_max cur_dim = self.h_max
pos_encoding = comfy.ops.cast_to_input(self.positional_encoding.reshape(1, cur_dim, cur_dim, -1), x) pos_encoding = seap.ops.cast_to_input(self.positional_encoding.reshape(1, cur_dim, cur_dim, -1), x)
if max_dim > cur_dim: if max_dim > cur_dim:
pos_encoding = F.interpolate(pos_encoding.movedim(-1, 1), (max_dim, max_dim), mode="bilinear").movedim(1, -1) pos_encoding = F.interpolate(pos_encoding.movedim(-1, 1), (max_dim, max_dim), mode="bilinear").movedim(1, -1)
@ -454,7 +454,7 @@ class MMDiT(nn.Module):
t = timestep t = timestep
c = self.cond_seq_linear(c_seq) # B, T_c, D c = self.cond_seq_linear(c_seq) # B, T_c, D
c = torch.cat([comfy.ops.cast_to_input(self.register_tokens, c).repeat(c.size(0), 1, 1), c], dim=1) c = torch.cat([seap.ops.cast_to_input(self.register_tokens, c).repeat(c.size(0), 1, 1), c], dim=1)
global_cond = self.t_embedder(t, x.dtype) # B, D global_cond = self.t_embedder(t, x.dtype) # B, D

View File

@ -18,8 +18,8 @@
import torch import torch
import torch.nn as nn import torch.nn as nn
from comfy.ldm.modules.attention import optimized_attention from seap.ldm.modules.attention import optimized_attention
import comfy.ops import seap.ops
class OptimizedAttention(nn.Module): class OptimizedAttention(nn.Module):
def __init__(self, c, nhead, dropout=0.0, dtype=None, device=None, operations=None): def __init__(self, c, nhead, dropout=0.0, dtype=None, device=None, operations=None):
@ -77,7 +77,7 @@ class GlobalResponseNorm(nn.Module):
def forward(self, x): def forward(self, x):
Gx = torch.norm(x, p=2, dim=(1, 2), keepdim=True) Gx = torch.norm(x, p=2, dim=(1, 2), keepdim=True)
Nx = Gx / (Gx.mean(dim=-1, keepdim=True) + 1e-6) Nx = Gx / (Gx.mean(dim=-1, keepdim=True) + 1e-6)
return comfy.ops.cast_to_input(self.gamma, x) * (x * Nx) + comfy.ops.cast_to_input(self.beta, x) + x return seap.ops.cast_to_input(self.gamma, x) * (x * Nx) + seap.ops.cast_to_input(self.beta, x) + x
class ResBlock(nn.Module): class ResBlock(nn.Module):

View File

@ -1,5 +1,5 @@
import torch import torch
import comfy.ops import seap.ops
def pad_to_patch_size(img, patch_size=(2, 2), padding_mode="circular"): def pad_to_patch_size(img, patch_size=(2, 2), padding_mode="circular"):
if padding_mode == "circular" and torch.jit.is_tracing() or torch.jit.is_scripting(): if padding_mode == "circular" and torch.jit.is_tracing() or torch.jit.is_scripting():
@ -15,7 +15,7 @@ except:
def rms_norm(x, weight, eps=1e-6): def rms_norm(x, weight, eps=1e-6):
if rms_norm_torch is not None and not (torch.jit.is_tracing() or torch.jit.is_scripting()): if rms_norm_torch is not None and not (torch.jit.is_tracing() or torch.jit.is_scripting()):
return rms_norm_torch(x, weight.shape, weight=comfy.ops.cast_to(weight, dtype=x.dtype, device=x.device), eps=eps) return rms_norm_torch(x, weight.shape, weight=seap.ops.cast_to(weight, dtype=x.dtype, device=x.device), eps=eps)
else: else:
rrms = torch.rsqrt(torch.mean(x**2, dim=-1, keepdim=True) + eps) rrms = torch.rsqrt(torch.mean(x**2, dim=-1, keepdim=True) + eps)
return (x * rrms) * comfy.ops.cast_to(weight, dtype=x.dtype, device=x.device) return (x * rrms) * seap.ops.cast_to(weight, dtype=x.dtype, device=x.device)

View File

@ -11,7 +11,7 @@ from .layers import (DoubleStreamBlock, EmbedND, LastLayer,
timestep_embedding) timestep_embedding)
from .model import Flux from .model import Flux
import comfy.ldm.common_dit import seap.ldm.common_dit
class MistolineCondDownsamplBlock(nn.Module): class MistolineCondDownsamplBlock(nn.Module):
def __init__(self, dtype=None, device=None, operations=None): def __init__(self, dtype=None, device=None, operations=None):
@ -179,7 +179,7 @@ class ControlNetFlux(Flux):
def forward(self, x, timesteps, context, y, guidance=None, hint=None, **kwargs): def forward(self, x, timesteps, context, y, guidance=None, hint=None, **kwargs):
patch_size = 2 patch_size = 2
if self.latent_input: if self.latent_input:
hint = comfy.ldm.common_dit.pad_to_patch_size(hint, (patch_size, patch_size)) hint = seap.ldm.common_dit.pad_to_patch_size(hint, (patch_size, patch_size))
elif self.mistoline: elif self.mistoline:
hint = hint * 2.0 - 1.0 hint = hint * 2.0 - 1.0
hint = self.input_cond_block(hint) hint = self.input_cond_block(hint)
@ -190,7 +190,7 @@ class ControlNetFlux(Flux):
hint = rearrange(hint, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch_size, pw=patch_size) hint = rearrange(hint, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch_size, pw=patch_size)
bs, c, h, w = x.shape bs, c, h, w = x.shape
x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch_size, patch_size)) x = seap.ldm.common_dit.pad_to_patch_size(x, (patch_size, patch_size))
img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch_size, pw=patch_size) img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch_size, pw=patch_size)

View File

@ -5,8 +5,8 @@ import torch
from torch import Tensor, nn from torch import Tensor, nn
from .math import attention, rope from .math import attention, rope
import comfy.ops import seap.ops
import comfy.ldm.common_dit import seap.ldm.common_dit
class EmbedND(nn.Module): class EmbedND(nn.Module):
@ -64,7 +64,7 @@ class RMSNorm(torch.nn.Module):
self.scale = nn.Parameter(torch.empty((dim), dtype=dtype, device=device)) self.scale = nn.Parameter(torch.empty((dim), dtype=dtype, device=device))
def forward(self, x: Tensor): def forward(self, x: Tensor):
return comfy.ldm.common_dit.rms_norm(x, self.scale, 1e-6) return seap.ldm.common_dit.rms_norm(x, self.scale, 1e-6)
class QKNorm(torch.nn.Module): class QKNorm(torch.nn.Module):

View File

@ -1,8 +1,8 @@
import torch import torch
from einops import rearrange from einops import rearrange
from torch import Tensor from torch import Tensor
from comfy.ldm.modules.attention import optimized_attention from seap.ldm.modules.attention import optimized_attention
import comfy.model_management import seap.model_management
def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor) -> Tensor: def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor) -> Tensor:
q, k = apply_rope(q, k, pe) q, k = apply_rope(q, k, pe)
@ -14,7 +14,7 @@ def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor) -> Tensor:
def rope(pos: Tensor, dim: int, theta: int) -> Tensor: def rope(pos: Tensor, dim: int, theta: int) -> Tensor:
assert dim % 2 == 0 assert dim % 2 == 0
if comfy.model_management.is_device_mps(pos.device) or comfy.model_management.is_intel_xpu(): if seap.model_management.is_device_mps(pos.device) or seap.model_management.is_intel_xpu():
device = torch.device("cpu") device = torch.device("cpu")
else: else:
device = pos.device device = pos.device

View File

@ -15,7 +15,7 @@ from .layers import (
) )
from einops import rearrange, repeat from einops import rearrange, repeat
import comfy.ldm.common_dit import seap.ldm.common_dit
@dataclass @dataclass
class FluxParams: class FluxParams:
@ -144,7 +144,7 @@ class Flux(nn.Module):
def forward(self, x, timestep, context, y, guidance, control=None, **kwargs): def forward(self, x, timestep, context, y, guidance, control=None, **kwargs):
bs, c, h, w = x.shape bs, c, h, w = x.shape
patch_size = 2 patch_size = 2
x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch_size, patch_size)) x = seap.ldm.common_dit.pad_to_patch_size(x, (patch_size, patch_size))
img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch_size, pw=patch_size) img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch_size, pw=patch_size)

View File

@ -1,7 +1,7 @@
import torch import torch
import torch.nn as nn import torch.nn as nn
from typing import Tuple, Union, Optional from typing import Tuple, Union, Optional
from comfy.ldm.modules.attention import optimized_attention from seap.ldm.modules.attention import optimized_attention
def reshape_for_broadcast(freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]], x: torch.Tensor, head_first=False): def reshape_for_broadcast(freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]], x: torch.Tensor, head_first=False):

View File

@ -6,16 +6,16 @@ import torch.nn.functional as F
from torch.utils import checkpoint from torch.utils import checkpoint
from comfy.ldm.modules.diffusionmodules.mmdit import ( from seap.ldm.modules.diffusionmodules.mmdit import (
Mlp, Mlp,
TimestepEmbedder, TimestepEmbedder,
PatchEmbed, PatchEmbed,
RMSNorm, RMSNorm,
) )
from comfy.ldm.modules.diffusionmodules.util import timestep_embedding from seap.ldm.modules.diffusionmodules.util import timestep_embedding
from .poolers import AttentionPool from .poolers import AttentionPool
import comfy.latent_formats import seap.latent_formats
from .models import HunYuanDiTBlock, calc_rope from .models import HunYuanDiTBlock, calc_rope
from .posemb_layers import get_2d_rotary_pos_embed, get_fill_resize_and_crop from .posemb_layers import get_2d_rotary_pos_embed, get_fill_resize_and_crop
@ -93,7 +93,7 @@ class HunYuanControlNet(nn.Module):
self.use_style_cond = use_style_cond self.use_style_cond = use_style_cond
self.norm = norm self.norm = norm
self.dtype = dtype self.dtype = dtype
self.latent_format = comfy.latent_formats.SDXL self.latent_format = seap.latent_formats.SDXL
self.mlp_t5 = nn.Sequential( self.mlp_t5 = nn.Sequential(
nn.Linear( nn.Linear(
@ -261,7 +261,7 @@ class HunYuanControlNet(nn.Module):
b_t5, l_t5, c_t5 = text_states_t5.shape b_t5, l_t5, c_t5 = text_states_t5.shape
text_states_t5 = self.mlp_t5(text_states_t5.view(-1, c_t5)).view(b_t5, l_t5, -1) text_states_t5 = self.mlp_t5(text_states_t5.view(-1, c_t5)).view(b_t5, l_t5, -1)
padding = comfy.ops.cast_to_input(self.text_embedding_padding, text_states) padding = seap.ops.cast_to_input(self.text_embedding_padding, text_states)
text_states[:, -self.text_len :] = torch.where( text_states[:, -self.text_len :] = torch.where(
text_states_mask[:, -self.text_len :].unsqueeze(2), text_states_mask[:, -self.text_len :].unsqueeze(2),

View File

@ -4,9 +4,9 @@ import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
import comfy.ops import seap.ops
from comfy.ldm.modules.diffusionmodules.mmdit import Mlp, TimestepEmbedder, PatchEmbed, RMSNorm from seap.ldm.modules.diffusionmodules.mmdit import Mlp, TimestepEmbedder, PatchEmbed, RMSNorm
from comfy.ldm.modules.diffusionmodules.util import timestep_embedding from seap.ldm.modules.diffusionmodules.util import timestep_embedding
from torch.utils import checkpoint from torch.utils import checkpoint
from .attn_layers import Attention, CrossAttention from .attn_layers import Attention, CrossAttention
@ -325,7 +325,7 @@ class HunYuanDiT(nn.Module):
b_t5, l_t5, c_t5 = text_states_t5.shape b_t5, l_t5, c_t5 = text_states_t5.shape
text_states_t5 = self.mlp_t5(text_states_t5.view(-1, c_t5)).view(b_t5, l_t5, -1) text_states_t5 = self.mlp_t5(text_states_t5.view(-1, c_t5)).view(b_t5, l_t5, -1)
padding = comfy.ops.cast_to_input(self.text_embedding_padding, text_states) padding = seap.ops.cast_to_input(self.text_embedding_padding, text_states)
text_states[:,-self.text_len:] = torch.where(text_states_mask[:,-self.text_len:].unsqueeze(2), text_states[:,-self.text_len:], padding[:self.text_len]) text_states[:,-self.text_len:] = torch.where(text_states_mask[:,-self.text_len:].unsqueeze(2), text_states[:,-self.text_len:], padding[:self.text_len])
text_states_t5[:,-self.text_len_t5:] = torch.where(text_states_t5_mask[:,-self.text_len_t5:].unsqueeze(2), text_states_t5[:,-self.text_len_t5:], padding[self.text_len:]) text_states_t5[:,-self.text_len_t5:] = torch.where(text_states_t5_mask[:,-self.text_len_t5:].unsqueeze(2), text_states_t5[:,-self.text_len_t5:], padding[self.text_len:])

View File

@ -1,8 +1,8 @@
import torch import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from comfy.ldm.modules.attention import optimized_attention from seap.ldm.modules.attention import optimized_attention
import comfy.ops import seap.ops
class AttentionPool(nn.Module): class AttentionPool(nn.Module):
def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: int = None, dtype=None, device=None, operations=None): def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: int = None, dtype=None, device=None, operations=None):
@ -19,7 +19,7 @@ class AttentionPool(nn.Module):
x = x[:,:self.positional_embedding.shape[0] - 1] x = x[:,:self.positional_embedding.shape[0] - 1]
x = x.permute(1, 0, 2) # NLC -> LNC x = x.permute(1, 0, 2) # NLC -> LNC
x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0) # (L+1)NC x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0) # (L+1)NC
x = x + comfy.ops.cast_to_input(self.positional_embedding[:, None, :], x) # (L+1)NC x = x + seap.ops.cast_to_input(self.positional_embedding[:, None, :], x) # (L+1)NC
q = self.q_proj(x[:1]) q = self.q_proj(x[:1])
k = self.k_proj(x) k = self.k_proj(x)

View File

@ -2,11 +2,11 @@ import torch
from contextlib import contextmanager from contextlib import contextmanager
from typing import Any, Dict, List, Optional, Tuple, Union from typing import Any, Dict, List, Optional, Tuple, Union
from comfy.ldm.modules.distributions.distributions import DiagonalGaussianDistribution from seap.ldm.modules.distributions.distributions import DiagonalGaussianDistribution
from comfy.ldm.util import instantiate_from_config from seap.ldm.util import instantiate_from_config
from comfy.ldm.modules.ema import LitEma from seap.ldm.modules.ema import LitEma
import comfy.ops import seap.ops
class DiagonalGaussianRegularizer(torch.nn.Module): class DiagonalGaussianRegularizer(torch.nn.Module):
def __init__(self, sample: bool = True): def __init__(self, sample: bool = True):
@ -151,21 +151,21 @@ class AutoencodingEngineLegacy(AutoencodingEngine):
ddconfig = kwargs.pop("ddconfig") ddconfig = kwargs.pop("ddconfig")
super().__init__( super().__init__(
encoder_config={ encoder_config={
"target": "comfy.ldm.modules.diffusionmodules.model.Encoder", "target": "seap.ldm.modules.diffusionmodules.model.Encoder",
"params": ddconfig, "params": ddconfig,
}, },
decoder_config={ decoder_config={
"target": "comfy.ldm.modules.diffusionmodules.model.Decoder", "target": "seap.ldm.modules.diffusionmodules.model.Decoder",
"params": ddconfig, "params": ddconfig,
}, },
**kwargs, **kwargs,
) )
self.quant_conv = comfy.ops.disable_weight_init.Conv2d( self.quant_conv = seap.ops.disable_weight_init.Conv2d(
(1 + ddconfig["double_z"]) * ddconfig["z_channels"], (1 + ddconfig["double_z"]) * ddconfig["z_channels"],
(1 + ddconfig["double_z"]) * embed_dim, (1 + ddconfig["double_z"]) * embed_dim,
1, 1,
) )
self.post_quant_conv = comfy.ops.disable_weight_init.Conv2d(embed_dim, ddconfig["z_channels"], 1) self.post_quant_conv = seap.ops.disable_weight_init.Conv2d(embed_dim, ddconfig["z_channels"], 1)
self.embed_dim = embed_dim self.embed_dim = embed_dim
def get_autoencoder_params(self) -> list: def get_autoencoder_params(self) -> list:
@ -219,7 +219,7 @@ class AutoencoderKL(AutoencodingEngineLegacy):
super().__init__( super().__init__(
regularizer_config={ regularizer_config={
"target": ( "target": (
"comfy.ldm.models.autoencoder.DiagonalGaussianRegularizer" "seap.ldm.models.autoencoder.DiagonalGaussianRegularizer"
) )
}, },
**kwargs, **kwargs,

View File

@ -9,15 +9,15 @@ import logging
from .diffusionmodules.util import AlphaBlender, timestep_embedding from .diffusionmodules.util import AlphaBlender, timestep_embedding
from .sub_quadratic_attention import efficient_dot_product_attention from .sub_quadratic_attention import efficient_dot_product_attention
from comfy import model_management from seap import model_management
if model_management.xformers_enabled(): if model_management.xformers_enabled():
import xformers import xformers
import xformers.ops import xformers.ops
from comfy.cli_args import args from seap.cli_args import args
import comfy.ops import seap.ops
ops = comfy.ops.disable_weight_init ops = seap.ops.disable_weight_init
FORCE_UPCAST_ATTENTION_DTYPE = model_management.force_upcast_attention_dtype() FORCE_UPCAST_ATTENTION_DTYPE = model_management.force_upcast_attention_dtype()

View File

@ -8,8 +8,8 @@ import torch.nn as nn
from .. import attention from .. import attention
from einops import rearrange, repeat from einops import rearrange, repeat
from .util import timestep_embedding from .util import timestep_embedding
import comfy.ops import seap.ops
import comfy.ldm.common_dit import seap.ldm.common_dit
def default(x, y): def default(x, y):
if x is not None: if x is not None:
@ -112,7 +112,7 @@ class PatchEmbed(nn.Module):
# f"Input width ({W}) should be divisible by patch size ({self.patch_size[1]})." # f"Input width ({W}) should be divisible by patch size ({self.patch_size[1]})."
# ) # )
if self.dynamic_img_pad: if self.dynamic_img_pad:
x = comfy.ldm.common_dit.pad_to_patch_size(x, self.patch_size, padding_mode=self.padding_mode) x = seap.ldm.common_dit.pad_to_patch_size(x, self.patch_size, padding_mode=self.padding_mode)
x = self.proj(x) x = self.proj(x)
if self.flatten: if self.flatten:
x = x.flatten(2).transpose(1, 2) # NCHW -> NLC x = x.flatten(2).transpose(1, 2) # NCHW -> NLC
@ -356,7 +356,7 @@ class RMSNorm(torch.nn.Module):
self.register_parameter("weight", None) self.register_parameter("weight", None)
def forward(self, x): def forward(self, x):
return comfy.ldm.common_dit.rms_norm(x, self.weight, self.eps) return seap.ldm.common_dit.rms_norm(x, self.weight, self.eps)
@ -906,7 +906,7 @@ class MMDiT(nn.Module):
context = self.context_processor(context) context = self.context_processor(context)
hw = x.shape[-2:] hw = x.shape[-2:]
x = self.x_embedder(x) + comfy.ops.cast_to_input(self.cropped_pos_embed(hw, device=x.device), x) x = self.x_embedder(x) + seap.ops.cast_to_input(self.cropped_pos_embed(hw, device=x.device), x)
c = self.t_embedder(t, dtype=x.dtype) # (N, D) c = self.t_embedder(t, dtype=x.dtype) # (N, D)
if y is not None and self.y_embedder is not None: if y is not None and self.y_embedder is not None:
y = self.y_embedder(y) # (N, D) y = self.y_embedder(y) # (N, D)

View File

@ -6,9 +6,9 @@ import numpy as np
from typing import Optional, Any from typing import Optional, Any
import logging import logging
from comfy import model_management from seap import model_management
import comfy.ops import seap.ops
ops = comfy.ops.disable_weight_init ops = seap.ops.disable_weight_init
if model_management.xformers_enabled_vae(): if model_management.xformers_enabled_vae():
import xformers import xformers

View File

@ -14,9 +14,9 @@ from .util import (
AlphaBlender, AlphaBlender,
) )
from ..attention import SpatialTransformer, SpatialVideoTransformer, default from ..attention import SpatialTransformer, SpatialVideoTransformer, default
from comfy.ldm.util import exists from seap.ldm.util import exists
import comfy.ops import seap.ops
ops = comfy.ops.disable_weight_init ops = seap.ops.disable_weight_init
class TimestepBlock(nn.Module): class TimestepBlock(nn.Module):
""" """

View File

@ -4,7 +4,7 @@ import numpy as np
from functools import partial from functools import partial
from .util import extract_into_tensor, make_beta_schedule from .util import extract_into_tensor, make_beta_schedule
from comfy.ldm.util import default from seap.ldm.util import default
class AbstractLowScaleModel(nn.Module): class AbstractLowScaleModel(nn.Module):

View File

@ -15,7 +15,7 @@ import torch.nn as nn
import numpy as np import numpy as np
from einops import repeat, rearrange from einops import repeat, rearrange
from comfy.ldm.util import instantiate_from_config from seap.ldm.util import instantiate_from_config
class AlphaBlender(nn.Module): class AlphaBlender(nn.Module):
strategies = ["learned", "fixed", "learned_with_images"] strategies = ["learned", "fixed", "learned_with_images"]

View File

@ -25,7 +25,7 @@ except ImportError:
from torch import Tensor from torch import Tensor
from typing import List from typing import List
from comfy import model_management from seap import model_management
def dynamic_slice( def dynamic_slice(
x: Tensor, x: Tensor,

View File

@ -4,8 +4,8 @@ from typing import Callable, Iterable, Union
import torch import torch
from einops import rearrange, repeat from einops import rearrange, repeat
import comfy.ops import seap.ops
ops = comfy.ops.disable_weight_init ops = seap.ops.disable_weight_init
from .diffusionmodules.model import ( from .diffusionmodules.model import (
AttnBlock, AttnBlock,

View File

@ -17,9 +17,9 @@
""" """
from __future__ import annotations from __future__ import annotations
import comfy.utils import seap.utils
import comfy.model_management import seap.model_management
import comfy.model_base import seap.model_base
import logging import logging
import torch import torch
@ -288,7 +288,7 @@ def model_lora_keys_unet(model, key_map={}):
key_map["lora_prior_unet_{}".format(key_lora)] = k #cascade lora: TODO put lora key prefix in the model config key_map["lora_prior_unet_{}".format(key_lora)] = k #cascade lora: TODO put lora key prefix in the model config
key_map["{}".format(k[:-len(".weight")])] = k #generic lora format without any weird key names key_map["{}".format(k[:-len(".weight")])] = k #generic lora format without any weird key names
diffusers_keys = comfy.utils.unet_to_diffusers(model.model_config.unet_config) diffusers_keys = seap.utils.unet_to_diffusers(model.model_config.unet_config)
for k in diffusers_keys: for k in diffusers_keys:
if k.endswith(".weight"): if k.endswith(".weight"):
unet_key = "diffusion_model.{}".format(diffusers_keys[k]) unet_key = "diffusion_model.{}".format(diffusers_keys[k])
@ -303,8 +303,8 @@ def model_lora_keys_unet(model, key_map={}):
diffusers_lora_key = diffusers_lora_key[:-2] diffusers_lora_key = diffusers_lora_key[:-2]
key_map[diffusers_lora_key] = unet_key key_map[diffusers_lora_key] = unet_key
if isinstance(model, comfy.model_base.SD3): #Diffusers lora SD3 if isinstance(model, seap.model_base.SD3): #Diffusers lora SD3
diffusers_keys = comfy.utils.mmdit_to_diffusers(model.model_config.unet_config, output_prefix="diffusion_model.") diffusers_keys = seap.utils.mmdit_to_diffusers(model.model_config.unet_config, output_prefix="diffusion_model.")
for k in diffusers_keys: for k in diffusers_keys:
if k.endswith(".weight"): if k.endswith(".weight"):
to = diffusers_keys[k] to = diffusers_keys[k]
@ -317,22 +317,22 @@ def model_lora_keys_unet(model, key_map={}):
key_lora = "lora_transformer_{}".format(k[:-len(".weight")].replace(".", "_")) #OneTrainer lora key_lora = "lora_transformer_{}".format(k[:-len(".weight")].replace(".", "_")) #OneTrainer lora
key_map[key_lora] = to key_map[key_lora] = to
if isinstance(model, comfy.model_base.AuraFlow): #Diffusers lora AuraFlow if isinstance(model, seap.model_base.AuraFlow): #Diffusers lora AuraFlow
diffusers_keys = comfy.utils.auraflow_to_diffusers(model.model_config.unet_config, output_prefix="diffusion_model.") diffusers_keys = seap.utils.auraflow_to_diffusers(model.model_config.unet_config, output_prefix="diffusion_model.")
for k in diffusers_keys: for k in diffusers_keys:
if k.endswith(".weight"): if k.endswith(".weight"):
to = diffusers_keys[k] to = diffusers_keys[k]
key_lora = "transformer.{}".format(k[:-len(".weight")]) #simpletrainer and probably regular diffusers lora format key_lora = "transformer.{}".format(k[:-len(".weight")]) #simpletrainer and probably regular diffusers lora format
key_map[key_lora] = to key_map[key_lora] = to
if isinstance(model, comfy.model_base.HunyuanDiT): if isinstance(model, seap.model_base.HunyuanDiT):
for k in sdk: for k in sdk:
if k.startswith("diffusion_model.") and k.endswith(".weight"): if k.startswith("diffusion_model.") and k.endswith(".weight"):
key_lora = k[len("diffusion_model."):-len(".weight")] key_lora = k[len("diffusion_model."):-len(".weight")]
key_map["base_model.model.{}".format(key_lora)] = k #official hunyuan lora format key_map["base_model.model.{}".format(key_lora)] = k #official hunyuan lora format
if isinstance(model, comfy.model_base.Flux): #Diffusers lora Flux if isinstance(model, seap.model_base.Flux): #Diffusers lora Flux
diffusers_keys = comfy.utils.flux_to_diffusers(model.model_config.unet_config, output_prefix="diffusion_model.") diffusers_keys = seap.utils.flux_to_diffusers(model.model_config.unet_config, output_prefix="diffusion_model.")
for k in diffusers_keys: for k in diffusers_keys:
if k.endswith(".weight"): if k.endswith(".weight"):
to = diffusers_keys[k] to = diffusers_keys[k]
@ -344,7 +344,7 @@ def model_lora_keys_unet(model, key_map={}):
def weight_decompose(dora_scale, weight, lora_diff, alpha, strength, intermediate_dtype, function): def weight_decompose(dora_scale, weight, lora_diff, alpha, strength, intermediate_dtype, function):
dora_scale = comfy.model_management.cast_to_device(dora_scale, weight.device, intermediate_dtype) dora_scale = seap.model_management.cast_to_device(dora_scale, weight.device, intermediate_dtype)
lora_diff *= alpha lora_diff *= alpha
weight_calc = weight + function(lora_diff).type(weight.dtype) weight_calc = weight + function(lora_diff).type(weight.dtype)
weight_norm = ( weight_norm = (
@ -415,7 +415,7 @@ def calculate_weight(patches, weight, key, intermediate_dtype=torch.float32):
weight *= strength_model weight *= strength_model
if isinstance(v, list): if isinstance(v, list):
v = (calculate_weight(v[1:], comfy.model_management.cast_to_device(v[0], weight.device, intermediate_dtype, copy=True), key, intermediate_dtype=intermediate_dtype), ) v = (calculate_weight(v[1:], seap.model_management.cast_to_device(v[0], weight.device, intermediate_dtype, copy=True), key, intermediate_dtype=intermediate_dtype),)
if len(v) == 1: if len(v) == 1:
patch_type = "diff" patch_type = "diff"
@ -435,10 +435,10 @@ def calculate_weight(patches, weight, key, intermediate_dtype=torch.float32):
if diff.shape != weight.shape: if diff.shape != weight.shape:
logging.warning("WARNING SHAPE MISMATCH {} WEIGHT NOT MERGED {} != {}".format(key, diff.shape, weight.shape)) logging.warning("WARNING SHAPE MISMATCH {} WEIGHT NOT MERGED {} != {}".format(key, diff.shape, weight.shape))
else: else:
weight += function(strength * comfy.model_management.cast_to_device(diff, weight.device, weight.dtype)) weight += function(strength * seap.model_management.cast_to_device(diff, weight.device, weight.dtype))
elif patch_type == "lora": #lora/locon elif patch_type == "lora": #lora/locon
mat1 = comfy.model_management.cast_to_device(v[0], weight.device, intermediate_dtype) mat1 = seap.model_management.cast_to_device(v[0], weight.device, intermediate_dtype)
mat2 = comfy.model_management.cast_to_device(v[1], weight.device, intermediate_dtype) mat2 = seap.model_management.cast_to_device(v[1], weight.device, intermediate_dtype)
dora_scale = v[4] dora_scale = v[4]
if v[2] is not None: if v[2] is not None:
alpha = v[2] / mat2.shape[0] alpha = v[2] / mat2.shape[0]
@ -447,7 +447,7 @@ def calculate_weight(patches, weight, key, intermediate_dtype=torch.float32):
if v[3] is not None: if v[3] is not None:
#locon mid weights, hopefully the math is fine because I didn't properly test it #locon mid weights, hopefully the math is fine because I didn't properly test it
mat3 = comfy.model_management.cast_to_device(v[3], weight.device, intermediate_dtype) mat3 = seap.model_management.cast_to_device(v[3], weight.device, intermediate_dtype)
final_shape = [mat2.shape[1], mat2.shape[0], mat3.shape[2], mat3.shape[3]] final_shape = [mat2.shape[1], mat2.shape[0], mat3.shape[2], mat3.shape[3]]
mat2 = torch.mm(mat2.transpose(0, 1).flatten(start_dim=1), mat3.transpose(0, 1).flatten(start_dim=1)).reshape(final_shape).transpose(0, 1) mat2 = torch.mm(mat2.transpose(0, 1).flatten(start_dim=1), mat3.transpose(0, 1).flatten(start_dim=1)).reshape(final_shape).transpose(0, 1)
try: try:
@ -471,23 +471,23 @@ def calculate_weight(patches, weight, key, intermediate_dtype=torch.float32):
if w1 is None: if w1 is None:
dim = w1_b.shape[0] dim = w1_b.shape[0]
w1 = torch.mm(comfy.model_management.cast_to_device(w1_a, weight.device, intermediate_dtype), w1 = torch.mm(seap.model_management.cast_to_device(w1_a, weight.device, intermediate_dtype),
comfy.model_management.cast_to_device(w1_b, weight.device, intermediate_dtype)) seap.model_management.cast_to_device(w1_b, weight.device, intermediate_dtype))
else: else:
w1 = comfy.model_management.cast_to_device(w1, weight.device, intermediate_dtype) w1 = seap.model_management.cast_to_device(w1, weight.device, intermediate_dtype)
if w2 is None: if w2 is None:
dim = w2_b.shape[0] dim = w2_b.shape[0]
if t2 is None: if t2 is None:
w2 = torch.mm(comfy.model_management.cast_to_device(w2_a, weight.device, intermediate_dtype), w2 = torch.mm(seap.model_management.cast_to_device(w2_a, weight.device, intermediate_dtype),
comfy.model_management.cast_to_device(w2_b, weight.device, intermediate_dtype)) seap.model_management.cast_to_device(w2_b, weight.device, intermediate_dtype))
else: else:
w2 = torch.einsum('i j k l, j r, i p -> p r k l', w2 = torch.einsum('i j k l, j r, i p -> p r k l',
comfy.model_management.cast_to_device(t2, weight.device, intermediate_dtype), seap.model_management.cast_to_device(t2, weight.device, intermediate_dtype),
comfy.model_management.cast_to_device(w2_b, weight.device, intermediate_dtype), seap.model_management.cast_to_device(w2_b, weight.device, intermediate_dtype),
comfy.model_management.cast_to_device(w2_a, weight.device, intermediate_dtype)) seap.model_management.cast_to_device(w2_a, weight.device, intermediate_dtype))
else: else:
w2 = comfy.model_management.cast_to_device(w2, weight.device, intermediate_dtype) w2 = seap.model_management.cast_to_device(w2, weight.device, intermediate_dtype)
if len(w2.shape) == 4: if len(w2.shape) == 4:
w1 = w1.unsqueeze(2).unsqueeze(2) w1 = w1.unsqueeze(2).unsqueeze(2)
@ -519,19 +519,19 @@ def calculate_weight(patches, weight, key, intermediate_dtype=torch.float32):
t1 = v[5] t1 = v[5]
t2 = v[6] t2 = v[6]
m1 = torch.einsum('i j k l, j r, i p -> p r k l', m1 = torch.einsum('i j k l, j r, i p -> p r k l',
comfy.model_management.cast_to_device(t1, weight.device, intermediate_dtype), seap.model_management.cast_to_device(t1, weight.device, intermediate_dtype),
comfy.model_management.cast_to_device(w1b, weight.device, intermediate_dtype), seap.model_management.cast_to_device(w1b, weight.device, intermediate_dtype),
comfy.model_management.cast_to_device(w1a, weight.device, intermediate_dtype)) seap.model_management.cast_to_device(w1a, weight.device, intermediate_dtype))
m2 = torch.einsum('i j k l, j r, i p -> p r k l', m2 = torch.einsum('i j k l, j r, i p -> p r k l',
comfy.model_management.cast_to_device(t2, weight.device, intermediate_dtype), seap.model_management.cast_to_device(t2, weight.device, intermediate_dtype),
comfy.model_management.cast_to_device(w2b, weight.device, intermediate_dtype), seap.model_management.cast_to_device(w2b, weight.device, intermediate_dtype),
comfy.model_management.cast_to_device(w2a, weight.device, intermediate_dtype)) seap.model_management.cast_to_device(w2a, weight.device, intermediate_dtype))
else: else:
m1 = torch.mm(comfy.model_management.cast_to_device(w1a, weight.device, intermediate_dtype), m1 = torch.mm(seap.model_management.cast_to_device(w1a, weight.device, intermediate_dtype),
comfy.model_management.cast_to_device(w1b, weight.device, intermediate_dtype)) seap.model_management.cast_to_device(w1b, weight.device, intermediate_dtype))
m2 = torch.mm(comfy.model_management.cast_to_device(w2a, weight.device, intermediate_dtype), m2 = torch.mm(seap.model_management.cast_to_device(w2a, weight.device, intermediate_dtype),
comfy.model_management.cast_to_device(w2b, weight.device, intermediate_dtype)) seap.model_management.cast_to_device(w2b, weight.device, intermediate_dtype))
try: try:
lora_diff = (m1 * m2).reshape(weight.shape) lora_diff = (m1 * m2).reshape(weight.shape)
@ -556,10 +556,10 @@ def calculate_weight(patches, weight, key, intermediate_dtype=torch.float32):
old_glora = False old_glora = False
rank = v[1].shape[0] rank = v[1].shape[0]
a1 = comfy.model_management.cast_to_device(v[0].flatten(start_dim=1), weight.device, intermediate_dtype) a1 = seap.model_management.cast_to_device(v[0].flatten(start_dim=1), weight.device, intermediate_dtype)
a2 = comfy.model_management.cast_to_device(v[1].flatten(start_dim=1), weight.device, intermediate_dtype) a2 = seap.model_management.cast_to_device(v[1].flatten(start_dim=1), weight.device, intermediate_dtype)
b1 = comfy.model_management.cast_to_device(v[2].flatten(start_dim=1), weight.device, intermediate_dtype) b1 = seap.model_management.cast_to_device(v[2].flatten(start_dim=1), weight.device, intermediate_dtype)
b2 = comfy.model_management.cast_to_device(v[3].flatten(start_dim=1), weight.device, intermediate_dtype) b2 = seap.model_management.cast_to_device(v[3].flatten(start_dim=1), weight.device, intermediate_dtype)
if v[4] is not None: if v[4] is not None:
alpha = v[4] / rank alpha = v[4] / rank

View File

@ -18,24 +18,24 @@
import torch import torch
import logging import logging
from comfy.ldm.modules.diffusionmodules.openaimodel import UNetModel, Timestep from seap.ldm.modules.diffusionmodules.openaimodel import UNetModel, Timestep
from comfy.ldm.cascade.stage_c import StageC from seap.ldm.cascade.stage_c import StageC
from comfy.ldm.cascade.stage_b import StageB from seap.ldm.cascade.stage_b import StageB
from comfy.ldm.modules.encoders.noise_aug_modules import CLIPEmbeddingNoiseAugmentation from seap.ldm.modules.encoders.noise_aug_modules import CLIPEmbeddingNoiseAugmentation
from comfy.ldm.modules.diffusionmodules.upscaling import ImageConcatWithNoiseAugmentation from seap.ldm.modules.diffusionmodules.upscaling import ImageConcatWithNoiseAugmentation
from comfy.ldm.modules.diffusionmodules.mmdit import OpenAISignatureMMDITWrapper from seap.ldm.modules.diffusionmodules.mmdit import OpenAISignatureMMDITWrapper
import comfy.ldm.aura.mmdit import seap.ldm.aura.mmdit
import comfy.ldm.hydit.models import seap.ldm.hydit.models
import comfy.ldm.audio.dit import seap.ldm.audio.dit
import comfy.ldm.audio.embedders import seap.ldm.audio.embedders
import comfy.ldm.flux.model import seap.ldm.flux.model
import comfy.model_management import seap.model_management
import comfy.conds import seap.conds
import comfy.ops import seap.ops
from enum import Enum from enum import Enum
from . import utils from . import utils
import comfy.latent_formats import seap.latent_formats
import math import math
class ModelType(Enum): class ModelType(Enum):
@ -49,7 +49,7 @@ class ModelType(Enum):
FLUX = 8 FLUX = 8
from comfy.model_sampling import EPS, V_PREDICTION, EDM, ModelSamplingDiscrete, ModelSamplingContinuousEDM, StableCascadeSampling, ModelSamplingContinuousV from seap.model_sampling import EPS, V_PREDICTION, EDM, ModelSamplingDiscrete, ModelSamplingContinuousEDM, StableCascadeSampling, ModelSamplingContinuousV
def model_sampling(model_config, model_type): def model_sampling(model_config, model_type):
@ -63,8 +63,8 @@ def model_sampling(model_config, model_type):
c = V_PREDICTION c = V_PREDICTION
s = ModelSamplingContinuousEDM s = ModelSamplingContinuousEDM
elif model_type == ModelType.FLOW: elif model_type == ModelType.FLOW:
c = comfy.model_sampling.CONST c = seap.model_sampling.CONST
s = comfy.model_sampling.ModelSamplingDiscreteFlow s = seap.model_sampling.ModelSamplingDiscreteFlow
elif model_type == ModelType.STABLE_CASCADE: elif model_type == ModelType.STABLE_CASCADE:
c = EPS c = EPS
s = StableCascadeSampling s = StableCascadeSampling
@ -75,8 +75,8 @@ def model_sampling(model_config, model_type):
c = V_PREDICTION c = V_PREDICTION
s = ModelSamplingContinuousV s = ModelSamplingContinuousV
elif model_type == ModelType.FLUX: elif model_type == ModelType.FLUX:
c = comfy.model_sampling.CONST c = seap.model_sampling.CONST
s = comfy.model_sampling.ModelSamplingFlux s = seap.model_sampling.ModelSamplingFlux
class ModelSampling(s, c): class ModelSampling(s, c):
pass pass
@ -96,11 +96,11 @@ class BaseModel(torch.nn.Module):
if not unet_config.get("disable_unet_model_creation", False): if not unet_config.get("disable_unet_model_creation", False):
if model_config.custom_operations is None: if model_config.custom_operations is None:
operations = comfy.ops.pick_operations(unet_config.get("dtype", None), self.manual_cast_dtype, fp8_optimizations=model_config.optimizations.get("fp8", False)) operations = seap.ops.pick_operations(unet_config.get("dtype", None), self.manual_cast_dtype, fp8_optimizations=model_config.optimizations.get("fp8", False))
else: else:
operations = model_config.custom_operations operations = model_config.custom_operations
self.diffusion_model = unet_model(**unet_config, device=device, operations=operations) self.diffusion_model = unet_model(**unet_config, device=device, operations=operations)
if comfy.model_management.force_channels_last(): if seap.model_management.force_channels_last():
self.diffusion_model.to(memory_format=torch.channels_last) self.diffusion_model.to(memory_format=torch.channels_last)
logging.debug("using channels last mode for diffusion model") logging.debug("using channels last mode for diffusion model")
logging.info("model weight dtype {}, manual cast: {}".format(self.get_dtype(), self.manual_cast_dtype)) logging.info("model weight dtype {}, manual cast: {}".format(self.get_dtype(), self.manual_cast_dtype))
@ -191,23 +191,23 @@ class BaseModel(torch.nn.Module):
elif ck == "masked_image": elif ck == "masked_image":
cond_concat.append(self.blank_inpaint_image_like(noise)) cond_concat.append(self.blank_inpaint_image_like(noise))
data = torch.cat(cond_concat, dim=1) data = torch.cat(cond_concat, dim=1)
out['c_concat'] = comfy.conds.CONDNoiseShape(data) out['c_concat'] = seap.conds.CONDNoiseShape(data)
adm = self.encode_adm(**kwargs) adm = self.encode_adm(**kwargs)
if adm is not None: if adm is not None:
out['y'] = comfy.conds.CONDRegular(adm) out['y'] = seap.conds.CONDRegular(adm)
cross_attn = kwargs.get("cross_attn", None) cross_attn = kwargs.get("cross_attn", None)
if cross_attn is not None: if cross_attn is not None:
out['c_crossattn'] = comfy.conds.CONDCrossAttn(cross_attn) out['c_crossattn'] = seap.conds.CONDCrossAttn(cross_attn)
cross_attn_cnet = kwargs.get("cross_attn_controlnet", None) cross_attn_cnet = kwargs.get("cross_attn_controlnet", None)
if cross_attn_cnet is not None: if cross_attn_cnet is not None:
out['crossattn_controlnet'] = comfy.conds.CONDCrossAttn(cross_attn_cnet) out['crossattn_controlnet'] = seap.conds.CONDCrossAttn(cross_attn_cnet)
c_concat = kwargs.get("noise_concat", None) c_concat = kwargs.get("noise_concat", None)
if c_concat is not None: if c_concat is not None:
out['c_concat'] = comfy.conds.CONDNoiseShape(c_concat) out['c_concat'] = seap.conds.CONDNoiseShape(c_concat)
return out return out
@ -267,13 +267,13 @@ class BaseModel(torch.nn.Module):
self.blank_inpaint_image_like = blank_inpaint_image_like self.blank_inpaint_image_like = blank_inpaint_image_like
def memory_required(self, input_shape): def memory_required(self, input_shape):
if comfy.model_management.xformers_enabled() or comfy.model_management.pytorch_attention_flash_attention(): if seap.model_management.xformers_enabled() or seap.model_management.pytorch_attention_flash_attention():
dtype = self.get_dtype() dtype = self.get_dtype()
if self.manual_cast_dtype is not None: if self.manual_cast_dtype is not None:
dtype = self.manual_cast_dtype dtype = self.manual_cast_dtype
#TODO: this needs to be tweaked #TODO: this needs to be tweaked
area = input_shape[0] * math.prod(input_shape[2:]) area = input_shape[0] * math.prod(input_shape[2:])
return (area * comfy.model_management.dtype_size(dtype) * 0.01 * self.memory_usage_factor) * (1024 * 1024) return (area * seap.model_management.dtype_size(dtype) * 0.01 * self.memory_usage_factor) * (1024 * 1024)
else: else:
#TODO: this formula might be too aggressive since I tweaked the sub-quad and split algorithms to use less memory. #TODO: this formula might be too aggressive since I tweaked the sub-quad and split algorithms to use less memory.
area = input_shape[0] * math.prod(input_shape[2:]) area = input_shape[0] * math.prod(input_shape[2:])
@ -398,7 +398,7 @@ class SVD_img2vid(BaseModel):
out = {} out = {}
adm = self.encode_adm(**kwargs) adm = self.encode_adm(**kwargs)
if adm is not None: if adm is not None:
out['y'] = comfy.conds.CONDRegular(adm) out['y'] = seap.conds.CONDRegular(adm)
latent_image = kwargs.get("concat_latent_image", None) latent_image = kwargs.get("concat_latent_image", None)
noise = kwargs.get("noise", None) noise = kwargs.get("noise", None)
@ -412,16 +412,16 @@ class SVD_img2vid(BaseModel):
latent_image = utils.resize_to_batch_size(latent_image, noise.shape[0]) latent_image = utils.resize_to_batch_size(latent_image, noise.shape[0])
out['c_concat'] = comfy.conds.CONDNoiseShape(latent_image) out['c_concat'] = seap.conds.CONDNoiseShape(latent_image)
cross_attn = kwargs.get("cross_attn", None) cross_attn = kwargs.get("cross_attn", None)
if cross_attn is not None: if cross_attn is not None:
out['c_crossattn'] = comfy.conds.CONDCrossAttn(cross_attn) out['c_crossattn'] = seap.conds.CONDCrossAttn(cross_attn)
if "time_conditioning" in kwargs: if "time_conditioning" in kwargs:
out["time_context"] = comfy.conds.CONDCrossAttn(kwargs["time_conditioning"]) out["time_context"] = seap.conds.CONDCrossAttn(kwargs["time_conditioning"])
out['num_video_frames'] = comfy.conds.CONDConstant(noise.shape[0]) out['num_video_frames'] = seap.conds.CONDConstant(noise.shape[0])
return out return out
class SV3D_u(SVD_img2vid): class SV3D_u(SVD_img2vid):
@ -457,7 +457,7 @@ class SV3D_p(SVD_img2vid):
class Stable_Zero123(BaseModel): class Stable_Zero123(BaseModel):
def __init__(self, model_config, model_type=ModelType.EPS, device=None, cc_projection_weight=None, cc_projection_bias=None): def __init__(self, model_config, model_type=ModelType.EPS, device=None, cc_projection_weight=None, cc_projection_bias=None):
super().__init__(model_config, model_type, device=device) super().__init__(model_config, model_type, device=device)
self.cc_projection = comfy.ops.manual_cast.Linear(cc_projection_weight.shape[1], cc_projection_weight.shape[0], dtype=self.get_dtype(), device=device) self.cc_projection = seap.ops.manual_cast.Linear(cc_projection_weight.shape[1], cc_projection_weight.shape[0], dtype=self.get_dtype(), device=device)
self.cc_projection.weight.copy_(cc_projection_weight) self.cc_projection.weight.copy_(cc_projection_weight)
self.cc_projection.bias.copy_(cc_projection_bias) self.cc_projection.bias.copy_(cc_projection_bias)
@ -475,13 +475,13 @@ class Stable_Zero123(BaseModel):
latent_image = utils.resize_to_batch_size(latent_image, noise.shape[0]) latent_image = utils.resize_to_batch_size(latent_image, noise.shape[0])
out['c_concat'] = comfy.conds.CONDNoiseShape(latent_image) out['c_concat'] = seap.conds.CONDNoiseShape(latent_image)
cross_attn = kwargs.get("cross_attn", None) cross_attn = kwargs.get("cross_attn", None)
if cross_attn is not None: if cross_attn is not None:
if cross_attn.shape[-1] != 768: if cross_attn.shape[-1] != 768:
cross_attn = self.cc_projection(cross_attn) cross_attn = self.cc_projection(cross_attn)
out['c_crossattn'] = comfy.conds.CONDCrossAttn(cross_attn) out['c_crossattn'] = seap.conds.CONDCrossAttn(cross_attn)
return out return out
class SD_X4Upscaler(BaseModel): class SD_X4Upscaler(BaseModel):
@ -512,8 +512,8 @@ class SD_X4Upscaler(BaseModel):
image = utils.resize_to_batch_size(image, noise.shape[0]) image = utils.resize_to_batch_size(image, noise.shape[0])
out['c_concat'] = comfy.conds.CONDNoiseShape(image) out['c_concat'] = seap.conds.CONDNoiseShape(image)
out['y'] = comfy.conds.CONDRegular(noise_level) out['y'] = seap.conds.CONDRegular(noise_level)
return out return out
class IP2P: class IP2P:
@ -532,10 +532,10 @@ class IP2P:
image = utils.resize_to_batch_size(image, noise.shape[0]) image = utils.resize_to_batch_size(image, noise.shape[0])
out['c_concat'] = comfy.conds.CONDNoiseShape(self.process_ip2p_image_in(image)) out['c_concat'] = seap.conds.CONDNoiseShape(self.process_ip2p_image_in(image))
adm = self.encode_adm(**kwargs) adm = self.encode_adm(**kwargs)
if adm is not None: if adm is not None:
out['y'] = comfy.conds.CONDRegular(adm) out['y'] = seap.conds.CONDRegular(adm)
return out return out
class SD15_instructpix2pix(IP2P, BaseModel): class SD15_instructpix2pix(IP2P, BaseModel):
@ -547,7 +547,7 @@ class SDXL_instructpix2pix(IP2P, SDXL):
def __init__(self, model_config, model_type=ModelType.EPS, device=None): def __init__(self, model_config, model_type=ModelType.EPS, device=None):
super().__init__(model_config, model_type, device=device) super().__init__(model_config, model_type, device=device)
if model_type == ModelType.V_PREDICTION_EDM: if model_type == ModelType.V_PREDICTION_EDM:
self.process_ip2p_image_in = lambda image: comfy.latent_formats.SDXL().process_in(image) #cosxl ip2p self.process_ip2p_image_in = lambda image: seap.latent_formats.SDXL().process_in(image) #cosxl ip2p
else: else:
self.process_ip2p_image_in = lambda image: image #diffusers ip2p self.process_ip2p_image_in = lambda image: image #diffusers ip2p
@ -561,7 +561,7 @@ class StableCascade_C(BaseModel):
out = {} out = {}
clip_text_pooled = kwargs["pooled_output"] clip_text_pooled = kwargs["pooled_output"]
if clip_text_pooled is not None: if clip_text_pooled is not None:
out['clip_text_pooled'] = comfy.conds.CONDRegular(clip_text_pooled) out['clip_text_pooled'] = seap.conds.CONDRegular(clip_text_pooled)
if "unclip_conditioning" in kwargs: if "unclip_conditioning" in kwargs:
embeds = [] embeds = []
@ -571,13 +571,13 @@ class StableCascade_C(BaseModel):
clip_img = torch.cat(embeds, dim=1) clip_img = torch.cat(embeds, dim=1)
else: else:
clip_img = torch.zeros((1, 1, 768)) clip_img = torch.zeros((1, 1, 768))
out["clip_img"] = comfy.conds.CONDRegular(clip_img) out["clip_img"] = seap.conds.CONDRegular(clip_img)
out["sca"] = comfy.conds.CONDRegular(torch.zeros((1,))) out["sca"] = seap.conds.CONDRegular(torch.zeros((1,)))
out["crp"] = comfy.conds.CONDRegular(torch.zeros((1,))) out["crp"] = seap.conds.CONDRegular(torch.zeros((1,)))
cross_attn = kwargs.get("cross_attn", None) cross_attn = kwargs.get("cross_attn", None)
if cross_attn is not None: if cross_attn is not None:
out['clip_text'] = comfy.conds.CONDCrossAttn(cross_attn) out['clip_text'] = seap.conds.CONDCrossAttn(cross_attn)
return out return out
@ -592,13 +592,13 @@ class StableCascade_B(BaseModel):
clip_text_pooled = kwargs["pooled_output"] clip_text_pooled = kwargs["pooled_output"]
if clip_text_pooled is not None: if clip_text_pooled is not None:
out['clip'] = comfy.conds.CONDRegular(clip_text_pooled) out['clip'] = seap.conds.CONDRegular(clip_text_pooled)
#size of prior doesn't really matter if zeros because it gets resized but I still want it to get batched #size of prior doesn't really matter if zeros because it gets resized but I still want it to get batched
prior = kwargs.get("stable_cascade_prior", torch.zeros((1, 16, (noise.shape[2] * 4) // 42, (noise.shape[3] * 4) // 42), dtype=noise.dtype, layout=noise.layout, device=noise.device)) prior = kwargs.get("stable_cascade_prior", torch.zeros((1, 16, (noise.shape[2] * 4) // 42, (noise.shape[3] * 4) // 42), dtype=noise.dtype, layout=noise.layout, device=noise.device))
out["effnet"] = comfy.conds.CONDRegular(prior) out["effnet"] = seap.conds.CONDRegular(prior)
out["sca"] = comfy.conds.CONDRegular(torch.zeros((1,))) out["sca"] = seap.conds.CONDRegular(torch.zeros((1,)))
return out return out
@ -613,27 +613,27 @@ class SD3(BaseModel):
out = super().extra_conds(**kwargs) out = super().extra_conds(**kwargs)
cross_attn = kwargs.get("cross_attn", None) cross_attn = kwargs.get("cross_attn", None)
if cross_attn is not None: if cross_attn is not None:
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) out['c_crossattn'] = seap.conds.CONDRegular(cross_attn)
return out return out
class AuraFlow(BaseModel): class AuraFlow(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLOW, device=None): def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.aura.mmdit.MMDiT) super().__init__(model_config, model_type, device=device, unet_model=seap.ldm.aura.mmdit.MMDiT)
def extra_conds(self, **kwargs): def extra_conds(self, **kwargs):
out = super().extra_conds(**kwargs) out = super().extra_conds(**kwargs)
cross_attn = kwargs.get("cross_attn", None) cross_attn = kwargs.get("cross_attn", None)
if cross_attn is not None: if cross_attn is not None:
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) out['c_crossattn'] = seap.conds.CONDRegular(cross_attn)
return out return out
class StableAudio1(BaseModel): class StableAudio1(BaseModel):
def __init__(self, model_config, seconds_start_embedder_weights, seconds_total_embedder_weights, model_type=ModelType.V_PREDICTION_CONTINUOUS, device=None): def __init__(self, model_config, seconds_start_embedder_weights, seconds_total_embedder_weights, model_type=ModelType.V_PREDICTION_CONTINUOUS, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.audio.dit.AudioDiffusionTransformer) super().__init__(model_config, model_type, device=device, unet_model=seap.ldm.audio.dit.AudioDiffusionTransformer)
self.seconds_start_embedder = comfy.ldm.audio.embedders.NumberConditioner(768, min_val=0, max_val=512) self.seconds_start_embedder = seap.ldm.audio.embedders.NumberConditioner(768, min_val=0, max_val=512)
self.seconds_total_embedder = comfy.ldm.audio.embedders.NumberConditioner(768, min_val=0, max_val=512) self.seconds_total_embedder = seap.ldm.audio.embedders.NumberConditioner(768, min_val=0, max_val=512)
self.seconds_start_embedder.load_state_dict(seconds_start_embedder_weights) self.seconds_start_embedder.load_state_dict(seconds_start_embedder_weights)
self.seconds_total_embedder.load_state_dict(seconds_total_embedder_weights) self.seconds_total_embedder.load_state_dict(seconds_total_embedder_weights)
@ -650,12 +650,12 @@ class StableAudio1(BaseModel):
seconds_total_embed = self.seconds_total_embedder([seconds_total])[0].to(device) seconds_total_embed = self.seconds_total_embedder([seconds_total])[0].to(device)
global_embed = torch.cat([seconds_start_embed, seconds_total_embed], dim=-1).reshape((1, -1)) global_embed = torch.cat([seconds_start_embed, seconds_total_embed], dim=-1).reshape((1, -1))
out['global_embed'] = comfy.conds.CONDRegular(global_embed) out['global_embed'] = seap.conds.CONDRegular(global_embed)
cross_attn = kwargs.get("cross_attn", None) cross_attn = kwargs.get("cross_attn", None)
if cross_attn is not None: if cross_attn is not None:
cross_attn = torch.cat([cross_attn.to(device), seconds_start_embed.repeat((cross_attn.shape[0], 1, 1)), seconds_total_embed.repeat((cross_attn.shape[0], 1, 1))], dim=1) cross_attn = torch.cat([cross_attn.to(device), seconds_start_embed.repeat((cross_attn.shape[0], 1, 1)), seconds_total_embed.repeat((cross_attn.shape[0], 1, 1))], dim=1)
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) out['c_crossattn'] = seap.conds.CONDRegular(cross_attn)
return out return out
def state_dict_for_saving(self, clip_state_dict=None, vae_state_dict=None, clip_vision_state_dict=None): def state_dict_for_saving(self, clip_state_dict=None, vae_state_dict=None, clip_vision_state_dict=None):
@ -669,25 +669,25 @@ class StableAudio1(BaseModel):
class HunyuanDiT(BaseModel): class HunyuanDiT(BaseModel):
def __init__(self, model_config, model_type=ModelType.V_PREDICTION, device=None): def __init__(self, model_config, model_type=ModelType.V_PREDICTION, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.hydit.models.HunYuanDiT) super().__init__(model_config, model_type, device=device, unet_model=seap.ldm.hydit.models.HunYuanDiT)
def extra_conds(self, **kwargs): def extra_conds(self, **kwargs):
out = super().extra_conds(**kwargs) out = super().extra_conds(**kwargs)
cross_attn = kwargs.get("cross_attn", None) cross_attn = kwargs.get("cross_attn", None)
if cross_attn is not None: if cross_attn is not None:
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) out['c_crossattn'] = seap.conds.CONDRegular(cross_attn)
attention_mask = kwargs.get("attention_mask", None) attention_mask = kwargs.get("attention_mask", None)
if attention_mask is not None: if attention_mask is not None:
out['text_embedding_mask'] = comfy.conds.CONDRegular(attention_mask) out['text_embedding_mask'] = seap.conds.CONDRegular(attention_mask)
conditioning_mt5xl = kwargs.get("conditioning_mt5xl", None) conditioning_mt5xl = kwargs.get("conditioning_mt5xl", None)
if conditioning_mt5xl is not None: if conditioning_mt5xl is not None:
out['encoder_hidden_states_t5'] = comfy.conds.CONDRegular(conditioning_mt5xl) out['encoder_hidden_states_t5'] = seap.conds.CONDRegular(conditioning_mt5xl)
attention_mask_mt5xl = kwargs.get("attention_mask_mt5xl", None) attention_mask_mt5xl = kwargs.get("attention_mask_mt5xl", None)
if attention_mask_mt5xl is not None: if attention_mask_mt5xl is not None:
out['text_embedding_mask_t5'] = comfy.conds.CONDRegular(attention_mask_mt5xl) out['text_embedding_mask_t5'] = seap.conds.CONDRegular(attention_mask_mt5xl)
width = kwargs.get("width", 768) width = kwargs.get("width", 768)
height = kwargs.get("height", 768) height = kwargs.get("height", 768)
@ -696,12 +696,12 @@ class HunyuanDiT(BaseModel):
target_width = kwargs.get("target_width", width) target_width = kwargs.get("target_width", width)
target_height = kwargs.get("target_height", height) target_height = kwargs.get("target_height", height)
out['image_meta_size'] = comfy.conds.CONDRegular(torch.FloatTensor([[height, width, target_height, target_width, 0, 0]])) out['image_meta_size'] = seap.conds.CONDRegular(torch.FloatTensor([[height, width, target_height, target_width, 0, 0]]))
return out return out
class Flux(BaseModel): class Flux(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLUX, device=None): def __init__(self, model_config, model_type=ModelType.FLUX, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.flux.model.Flux) super().__init__(model_config, model_type, device=device, unet_model=seap.ldm.flux.model.Flux)
def encode_adm(self, **kwargs): def encode_adm(self, **kwargs):
return kwargs["pooled_output"] return kwargs["pooled_output"]
@ -710,6 +710,6 @@ class Flux(BaseModel):
out = super().extra_conds(**kwargs) out = super().extra_conds(**kwargs)
cross_attn = kwargs.get("cross_attn", None) cross_attn = kwargs.get("cross_attn", None)
if cross_attn is not None: if cross_attn is not None:
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) out['c_crossattn'] = seap.conds.CONDRegular(cross_attn)
out['guidance'] = comfy.conds.CONDRegular(torch.FloatTensor([kwargs.get("guidance", 3.5)])) out['guidance'] = seap.conds.CONDRegular(torch.FloatTensor([kwargs.get("guidance", 3.5)]))
return out return out

View File

@ -1,6 +1,6 @@
import comfy.supported_models import seap.supported_models
import comfy.supported_models_base import seap.supported_models_base
import comfy.utils import seap.utils
import math import math
import logging import logging
import torch import torch
@ -273,7 +273,7 @@ def detect_unet_config(state_dict, key_prefix):
return unet_config return unet_config
def model_config_from_unet_config(unet_config, state_dict=None): def model_config_from_unet_config(unet_config, state_dict=None):
for model_config in comfy.supported_models.models: for model_config in seap.supported_models.models:
if model_config.matches(unet_config, state_dict): if model_config.matches(unet_config, state_dict):
return model_config(unet_config) return model_config(unet_config)
@ -286,7 +286,7 @@ def model_config_from_unet(state_dict, unet_key_prefix, use_base_if_no_match=Fal
return None return None
model_config = model_config_from_unet_config(unet_config, state_dict) model_config = model_config_from_unet_config(unet_config, state_dict)
if model_config is None and use_base_if_no_match: if model_config is None and use_base_if_no_match:
return comfy.supported_models_base.BASE(unet_config) return seap.supported_models_base.BASE(unet_config)
else: else:
return model_config return model_config
@ -505,15 +505,15 @@ def convert_diffusers_mmdit(state_dict, output_prefix=""):
depth = count_blocks(state_dict, 'transformer_blocks.{}.') depth = count_blocks(state_dict, 'transformer_blocks.{}.')
depth_single_blocks = count_blocks(state_dict, 'single_transformer_blocks.{}.') depth_single_blocks = count_blocks(state_dict, 'single_transformer_blocks.{}.')
hidden_size = state_dict["x_embedder.bias"].shape[0] hidden_size = state_dict["x_embedder.bias"].shape[0]
sd_map = comfy.utils.flux_to_diffusers({"depth": depth, "depth_single_blocks": depth_single_blocks, "hidden_size": hidden_size}, output_prefix=output_prefix) sd_map = seap.utils.flux_to_diffusers({"depth": depth, "depth_single_blocks": depth_single_blocks, "hidden_size": hidden_size}, output_prefix=output_prefix)
elif 'transformer_blocks.0.attn.add_q_proj.weight' in state_dict: #SD3 elif 'transformer_blocks.0.attn.add_q_proj.weight' in state_dict: #SD3
num_blocks = count_blocks(state_dict, 'transformer_blocks.{}.') num_blocks = count_blocks(state_dict, 'transformer_blocks.{}.')
depth = state_dict["pos_embed.proj.weight"].shape[0] // 64 depth = state_dict["pos_embed.proj.weight"].shape[0] // 64
sd_map = comfy.utils.mmdit_to_diffusers({"depth": depth, "num_blocks": num_blocks}, output_prefix=output_prefix) sd_map = seap.utils.mmdit_to_diffusers({"depth": depth, "num_blocks": num_blocks}, output_prefix=output_prefix)
elif 'joint_transformer_blocks.0.attn.add_k_proj.weight' in state_dict: #AuraFlow elif 'joint_transformer_blocks.0.attn.add_k_proj.weight' in state_dict: #AuraFlow
num_joint = count_blocks(state_dict, 'joint_transformer_blocks.{}.') num_joint = count_blocks(state_dict, 'joint_transformer_blocks.{}.')
num_single = count_blocks(state_dict, 'single_transformer_blocks.{}.') num_single = count_blocks(state_dict, 'single_transformer_blocks.{}.')
sd_map = comfy.utils.auraflow_to_diffusers({"n_double_layers": num_joint, "n_layers": num_joint + num_single}, output_prefix=output_prefix) sd_map = seap.utils.auraflow_to_diffusers({"n_double_layers": num_joint, "n_layers": num_joint + num_single}, output_prefix=output_prefix)
else: else:
return None return None

View File

@ -19,7 +19,7 @@
import psutil import psutil
import logging import logging
from enum import Enum from enum import Enum
from comfy.cli_args import args from seap.cli_args import args
import torch import torch
import sys import sys
import platform import platform
@ -1101,7 +1101,7 @@ def unload_all_models():
def resolve_lowvram_weight(weight, model, key): #TODO: remove def resolve_lowvram_weight(weight, model, key): #TODO: remove
print("WARNING: The comfy.model_management.resolve_lowvram_weight function will be removed soon, please stop using it.") print("WARNING: The seap.model_management.resolve_lowvram_weight function will be removed soon, please stop using it.")
return weight return weight
#TODO: might be cleaner to put this somewhere else #TODO: might be cleaner to put this somewhere else

View File

@ -24,11 +24,11 @@ import uuid
import collections import collections
import math import math
import comfy.utils import seap.utils
import comfy.float import seap.float
import comfy.model_management import seap.model_management
import comfy.lora import seap.lora
from comfy.comfy_types import UnetWrapperFunction from seap.comfy_types import UnetWrapperFunction
def string_to_seed(data): def string_to_seed(data):
crc = 0xFFFFFFFF crc = 0xFFFFFFFF
@ -91,9 +91,9 @@ class LowVramPatch:
intermediate_dtype = weight.dtype intermediate_dtype = weight.dtype
if intermediate_dtype not in [torch.float32, torch.float16, torch.bfloat16]: #intermediate_dtype has to be one that is supported in math ops if intermediate_dtype not in [torch.float32, torch.float16, torch.bfloat16]: #intermediate_dtype has to be one that is supported in math ops
intermediate_dtype = torch.float32 intermediate_dtype = torch.float32
return comfy.float.stochastic_rounding(comfy.lora.calculate_weight(self.patches[self.key], weight.to(intermediate_dtype), self.key, intermediate_dtype=intermediate_dtype), weight.dtype, seed=string_to_seed(self.key)) return seap.float.stochastic_rounding(seap.lora.calculate_weight(self.patches[self.key], weight.to(intermediate_dtype), self.key, intermediate_dtype=intermediate_dtype), weight.dtype, seed=string_to_seed(self.key))
return comfy.lora.calculate_weight(self.patches[self.key], weight, self.key, intermediate_dtype=intermediate_dtype) return seap.lora.calculate_weight(self.patches[self.key], weight, self.key, intermediate_dtype=intermediate_dtype)
class ModelPatcher: class ModelPatcher:
def __init__(self, model, load_device, offload_device, size=0, weight_inplace_update=False): def __init__(self, model, load_device, offload_device, size=0, weight_inplace_update=False):
self.size = size self.size = size
@ -127,7 +127,7 @@ class ModelPatcher:
def model_size(self): def model_size(self):
if self.size > 0: if self.size > 0:
return self.size return self.size
self.size = comfy.model_management.module_size(self.model) self.size = seap.model_management.module_size(self.model)
return self.size return self.size
def loaded_size(self): def loaded_size(self):
@ -236,7 +236,7 @@ class ModelPatcher:
if name in self.object_patches_backup: if name in self.object_patches_backup:
return self.object_patches_backup[name] return self.object_patches_backup[name]
else: else:
return comfy.utils.get_attr(self.model, name) return seap.utils.get_attr(self.model, name)
def model_patches_to(self, device): def model_patches_to(self, device):
to = self.model_options["transformer_options"] to = self.model_options["transformer_options"]
@ -317,7 +317,7 @@ class ModelPatcher:
if key not in self.patches: if key not in self.patches:
return return
weight = comfy.utils.get_attr(self.model, key) weight = seap.utils.get_attr(self.model, key)
inplace_update = self.weight_inplace_update or inplace_update inplace_update = self.weight_inplace_update or inplace_update
@ -325,15 +325,15 @@ class ModelPatcher:
self.backup[key] = collections.namedtuple('Dimension', ['weight', 'inplace_update'])(weight.to(device=self.offload_device, copy=inplace_update), inplace_update) self.backup[key] = collections.namedtuple('Dimension', ['weight', 'inplace_update'])(weight.to(device=self.offload_device, copy=inplace_update), inplace_update)
if device_to is not None: if device_to is not None:
temp_weight = comfy.model_management.cast_to_device(weight, device_to, torch.float32, copy=True) temp_weight = seap.model_management.cast_to_device(weight, device_to, torch.float32, copy=True)
else: else:
temp_weight = weight.to(torch.float32, copy=True) temp_weight = weight.to(torch.float32, copy=True)
out_weight = comfy.lora.calculate_weight(self.patches[key], temp_weight, key) out_weight = seap.lora.calculate_weight(self.patches[key], temp_weight, key)
out_weight = comfy.float.stochastic_rounding(out_weight, weight.dtype, seed=string_to_seed(key)) out_weight = seap.float.stochastic_rounding(out_weight, weight.dtype, seed=string_to_seed(key))
if inplace_update: if inplace_update:
comfy.utils.copy_to_param(self.model, key, out_weight) seap.utils.copy_to_param(self.model, key, out_weight)
else: else:
comfy.utils.set_attr_param(self.model, key, out_weight) seap.utils.set_attr_param(self.model, key, out_weight)
def load(self, device_to=None, lowvram_model_memory=0, force_patch_weights=False, full_load=False): def load(self, device_to=None, lowvram_model_memory=0, force_patch_weights=False, full_load=False):
mem_counter = 0 mem_counter = 0
@ -342,7 +342,7 @@ class ModelPatcher:
loading = [] loading = []
for n, m in self.model.named_modules(): for n, m in self.model.named_modules():
if hasattr(m, "comfy_cast_weights") or hasattr(m, "weight"): if hasattr(m, "comfy_cast_weights") or hasattr(m, "weight"):
loading.append((comfy.model_management.module_size(m), n, m)) loading.append((seap.model_management.module_size(m), n, m))
load_completely = [] load_completely = []
loading.sort(reverse=True) loading.sort(reverse=True)
@ -422,7 +422,7 @@ class ModelPatcher:
def patch_model(self, device_to=None, lowvram_model_memory=0, load_weights=True, force_patch_weights=False): def patch_model(self, device_to=None, lowvram_model_memory=0, load_weights=True, force_patch_weights=False):
for k in self.object_patches: for k in self.object_patches:
old = comfy.utils.set_attr(self.model, k, self.object_patches[k]) old = seap.utils.set_attr(self.model, k, self.object_patches[k])
if k not in self.object_patches_backup: if k not in self.object_patches_backup:
self.object_patches_backup[k] = old self.object_patches_backup[k] = old
@ -449,9 +449,9 @@ class ModelPatcher:
for k in keys: for k in keys:
bk = self.backup[k] bk = self.backup[k]
if bk.inplace_update: if bk.inplace_update:
comfy.utils.copy_to_param(self.model, k, bk.weight) seap.utils.copy_to_param(self.model, k, bk.weight)
else: else:
comfy.utils.set_attr_param(self.model, k, bk.weight) seap.utils.set_attr_param(self.model, k, bk.weight)
self.backup.clear() self.backup.clear()
@ -466,7 +466,7 @@ class ModelPatcher:
keys = list(self.object_patches_backup.keys()) keys = list(self.object_patches_backup.keys())
for k in keys: for k in keys:
comfy.utils.set_attr(self.model, k, self.object_patches_backup[k]) seap.utils.set_attr(self.model, k, self.object_patches_backup[k])
self.object_patches_backup.clear() self.object_patches_backup.clear()
@ -478,7 +478,7 @@ class ModelPatcher:
for n, m in self.model.named_modules(): for n, m in self.model.named_modules():
shift_lowvram = False shift_lowvram = False
if hasattr(m, "comfy_cast_weights"): if hasattr(m, "comfy_cast_weights"):
module_mem = comfy.model_management.module_size(m) module_mem = seap.model_management.module_size(m)
unload_list.append((module_mem, n, m)) unload_list.append((module_mem, n, m))
unload_list.sort() unload_list.sort()
@ -496,9 +496,9 @@ class ModelPatcher:
bk = self.backup.get(key, None) bk = self.backup.get(key, None)
if bk is not None: if bk is not None:
if bk.inplace_update: if bk.inplace_update:
comfy.utils.copy_to_param(self.model, key, bk.weight) seap.utils.copy_to_param(self.model, key, bk.weight)
else: else:
comfy.utils.set_attr_param(self.model, key, bk.weight) seap.utils.set_attr_param(self.model, key, bk.weight)
self.backup.pop(key) self.backup.pop(key)
m.to(device_to) m.to(device_to)
@ -536,5 +536,5 @@ class ModelPatcher:
return self.model.device return self.model.device
def calculate_weight(self, patches, weight, key, intermediate_dtype=torch.float32): def calculate_weight(self, patches, weight, key, intermediate_dtype=torch.float32):
print("WARNING the ModelPatcher.calculate_weight function is deprecated, please use: comfy.lora.calculate_weight instead") print("WARNING the ModelPatcher.calculate_weight function is deprecated, please use: seap.lora.calculate_weight instead")
return comfy.lora.calculate_weight(patches, weight, key, intermediate_dtype=intermediate_dtype) return seap.lora.calculate_weight(patches, weight, key, intermediate_dtype=intermediate_dtype)

View File

@ -1,5 +1,5 @@
import torch import torch
from comfy.ldm.modules.diffusionmodules.util import make_beta_schedule from seap.ldm.modules.diffusionmodules.util import make_beta_schedule
import math import math
class EPS: class EPS:

View File

@ -17,8 +17,8 @@
""" """
import torch import torch
import comfy.model_management import seap.model_management
from comfy.cli_args import args from seap.cli_args import args
def cast_to(weight, dtype=None, device=None, non_blocking=False, copy=False): def cast_to(weight, dtype=None, device=None, non_blocking=False, copy=False):
if device is None or weight.device == device: if device is None or weight.device == device:
@ -44,7 +44,7 @@ def cast_bias_weight(s, input=None, dtype=None, device=None, bias_dtype=None):
device = input.device device = input.device
bias = None bias = None
non_blocking = comfy.model_management.device_supports_non_blocking(device) non_blocking = seap.model_management.device_supports_non_blocking(device)
if s.bias is not None: if s.bias is not None:
has_function = s.bias_function is not None has_function = s.bias_function is not None
bias = cast_to(s.bias, bias_dtype, device, non_blocking=non_blocking, copy=has_function) bias = cast_to(s.bias, bias_dtype, device, non_blocking=non_blocking, copy=has_function)
@ -300,13 +300,13 @@ class fp8_ops(manual_cast):
def pick_operations(weight_dtype, compute_dtype, load_device=None, disable_fast_fp8=False, fp8_optimizations=False): def pick_operations(weight_dtype, compute_dtype, load_device=None, disable_fast_fp8=False, fp8_optimizations=False):
if comfy.model_management.supports_fp8_compute(load_device): if seap.model_management.supports_fp8_compute(load_device):
if (fp8_optimizations or args.fast) and not disable_fast_fp8: if (fp8_optimizations or args.fast) and not disable_fast_fp8:
return fp8_ops return fp8_ops
if compute_dtype is None or weight_dtype == compute_dtype: if compute_dtype is None or weight_dtype == compute_dtype:
return disable_weight_init return disable_weight_init
if args.fast and not disable_fast_fp8: if args.fast and not disable_fast_fp8:
if comfy.model_management.supports_fp8_compute(load_device): if seap.model_management.supports_fp8_compute(load_device):
return fp8_ops return fp8_ops
return manual_cast return manual_cast

View File

@ -1,7 +1,7 @@
import torch import torch
import comfy.model_management import seap.model_management
import comfy.samplers import seap.samplers
import comfy.utils import seap.utils
import numpy as np import numpy as np
import logging import logging
@ -27,24 +27,24 @@ def prepare_noise(latent_image, seed, noise_inds=None):
def fix_empty_latent_channels(model, latent_image): def fix_empty_latent_channels(model, latent_image):
latent_channels = model.get_model_object("latent_format").latent_channels #Resize the empty latent image so it has the right number of channels latent_channels = model.get_model_object("latent_format").latent_channels #Resize the empty latent image so it has the right number of channels
if latent_channels != latent_image.shape[1] and torch.count_nonzero(latent_image) == 0: if latent_channels != latent_image.shape[1] and torch.count_nonzero(latent_image) == 0:
latent_image = comfy.utils.repeat_to_batch_size(latent_image, latent_channels, dim=1) latent_image = seap.utils.repeat_to_batch_size(latent_image, latent_channels, dim=1)
return latent_image return latent_image
def prepare_sampling(model, noise_shape, positive, negative, noise_mask): def prepare_sampling(model, noise_shape, positive, negative, noise_mask):
logging.warning("Warning: comfy.sample.prepare_sampling isn't used anymore and can be removed") logging.warning("Warning: seap.sample.prepare_sampling isn't used anymore and can be removed")
return model, positive, negative, noise_mask, [] return model, positive, negative, noise_mask, []
def cleanup_additional_models(models): def cleanup_additional_models(models):
logging.warning("Warning: comfy.sample.cleanup_additional_models isn't used anymore and can be removed") logging.warning("Warning: seap.sample.cleanup_additional_models isn't used anymore and can be removed")
def sample(model, noise, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise=1.0, disable_noise=False, start_step=None, last_step=None, force_full_denoise=False, noise_mask=None, sigmas=None, callback=None, disable_pbar=False, seed=None): def sample(model, noise, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise=1.0, disable_noise=False, start_step=None, last_step=None, force_full_denoise=False, noise_mask=None, sigmas=None, callback=None, disable_pbar=False, seed=None):
sampler = comfy.samplers.KSampler(model, steps=steps, device=model.load_device, sampler=sampler_name, scheduler=scheduler, denoise=denoise, model_options=model.model_options) sampler = seap.samplers.KSampler(model, steps=steps, device=model.load_device, sampler=sampler_name, scheduler=scheduler, denoise=denoise, model_options=model.model_options)
samples = sampler.sample(noise, positive, negative, cfg=cfg, latent_image=latent_image, start_step=start_step, last_step=last_step, force_full_denoise=force_full_denoise, denoise_mask=noise_mask, sigmas=sigmas, callback=callback, disable_pbar=disable_pbar, seed=seed) samples = sampler.sample(noise, positive, negative, cfg=cfg, latent_image=latent_image, start_step=start_step, last_step=last_step, force_full_denoise=force_full_denoise, denoise_mask=noise_mask, sigmas=sigmas, callback=callback, disable_pbar=disable_pbar, seed=seed)
samples = samples.to(comfy.model_management.intermediate_device()) samples = samples.to(seap.model_management.intermediate_device())
return samples return samples
def sample_custom(model, noise, cfg, sampler, sigmas, positive, negative, latent_image, noise_mask=None, callback=None, disable_pbar=False, seed=None): def sample_custom(model, noise, cfg, sampler, sigmas, positive, negative, latent_image, noise_mask=None, callback=None, disable_pbar=False, seed=None):
samples = comfy.samplers.sample(model, noise, positive, negative, cfg, model.load_device, sampler, sigmas, model_options=model.model_options, latent_image=latent_image, denoise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=seed) samples = seap.samplers.sample(model, noise, positive, negative, cfg, model.load_device, sampler, sigmas, model_options=model.model_options, latent_image=latent_image, denoise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=seed)
samples = samples.to(comfy.model_management.intermediate_device()) samples = samples.to(seap.model_management.intermediate_device())
return samples return samples

View File

@ -1,12 +1,12 @@
import torch import torch
import comfy.model_management import seap.model_management
import comfy.conds import seap.conds
def prepare_mask(noise_mask, shape, device): def prepare_mask(noise_mask, shape, device):
"""ensures noise mask is of proper dimensions""" """ensures noise mask is of proper dimensions"""
noise_mask = torch.nn.functional.interpolate(noise_mask.reshape((-1, 1, noise_mask.shape[-2], noise_mask.shape[-1])), size=(shape[2], shape[3]), mode="bilinear") noise_mask = torch.nn.functional.interpolate(noise_mask.reshape((-1, 1, noise_mask.shape[-2], noise_mask.shape[-1])), size=(shape[2], shape[3]), mode="bilinear")
noise_mask = torch.cat([noise_mask] * shape[1], dim=1) noise_mask = torch.cat([noise_mask] * shape[1], dim=1)
noise_mask = comfy.utils.repeat_to_batch_size(noise_mask, shape[0]) noise_mask = seap.utils.repeat_to_batch_size(noise_mask, shape[0])
noise_mask = noise_mask.to(device) noise_mask = noise_mask.to(device)
return noise_mask return noise_mask
@ -23,7 +23,7 @@ def convert_cond(cond):
temp = c[1].copy() temp = c[1].copy()
model_conds = temp.get("model_conds", {}) model_conds = temp.get("model_conds", {})
if c[0] is not None: if c[0] is not None:
model_conds["c_crossattn"] = comfy.conds.CONDCrossAttn(c[0]) #TODO: remove model_conds["c_crossattn"] = seap.conds.CONDCrossAttn(c[0]) #TODO: remove
temp["cross_attn"] = c[0] temp["cross_attn"] = c[0]
temp["model_conds"] = model_conds temp["model_conds"] = model_conds
out.append(temp) out.append(temp)
@ -63,7 +63,7 @@ def prepare_sampling(model, noise_shape, conds):
models, inference_memory = get_additional_models(conds, model.model_dtype()) models, inference_memory = get_additional_models(conds, model.model_dtype())
memory_required = model.memory_required([noise_shape[0] * 2] + list(noise_shape[1:])) + inference_memory memory_required = model.memory_required([noise_shape[0] * 2] + list(noise_shape[1:])) + inference_memory
minimum_memory_required = model.memory_required([noise_shape[0]] + list(noise_shape[1:])) + inference_memory minimum_memory_required = model.memory_required([noise_shape[0]] + list(noise_shape[1:])) + inference_memory
comfy.model_management.load_models_gpu([model] + models, memory_required=memory_required, minimum_memory_required=minimum_memory_required) seap.model_management.load_models_gpu([model] + models, memory_required=memory_required, minimum_memory_required=minimum_memory_required)
real_model = model.model real_model = model.model
return real_model, conds, models return real_model, conds, models

View File

@ -2,10 +2,10 @@ from .k_diffusion import sampling as k_diffusion_sampling
from .extra_samplers import uni_pc from .extra_samplers import uni_pc
import torch import torch
import collections import collections
from comfy import model_management from seap import model_management
import math import math
import logging import logging
import comfy.sampler_helpers import seap.sampler_helpers
import scipy.stats import scipy.stats
import numpy import numpy
@ -249,7 +249,7 @@ def calc_cond_batch(model, conds, x_in, timestep, model_options):
return out_conds return out_conds
def calc_cond_uncond_batch(model, cond, uncond, x_in, timestep, model_options): #TODO: remove def calc_cond_uncond_batch(model, cond, uncond, x_in, timestep, model_options): #TODO: remove
logging.warning("WARNING: The comfy.samplers.calc_cond_uncond_batch function is deprecated please use the calc_cond_batch one instead.") logging.warning("WARNING: The seap.samplers.calc_cond_uncond_batch function is deprecated please use the calc_cond_batch one instead.")
return tuple(calc_cond_batch(model, [cond, uncond], x_in, timestep, model_options)) return tuple(calc_cond_batch(model, [cond, uncond], x_in, timestep, model_options))
def cfg_function(model, cond_pred, uncond_pred, cond_scale, x, timestep, model_options={}, cond=None, uncond=None): def cfg_function(model, cond_pred, uncond_pred, cond_scale, x, timestep, model_options={}, cond=None, uncond=None):
@ -434,7 +434,7 @@ def resolve_areas_and_cond_masks_multidim(conditions, dims, device):
conditions[i] = modified conditions[i] = modified
def resolve_areas_and_cond_masks(conditions, h, w, device): def resolve_areas_and_cond_masks(conditions, h, w, device):
logging.warning("WARNING: The comfy.samplers.resolve_areas_and_cond_masks function is deprecated please use the resolve_areas_and_cond_masks_multidim one instead.") logging.warning("WARNING: The seap.samplers.resolve_areas_and_cond_masks function is deprecated please use the resolve_areas_and_cond_masks_multidim one instead.")
return resolve_areas_and_cond_masks_multidim(conditions, [h, w], device) return resolve_areas_and_cond_masks_multidim(conditions, [h, w], device)
def create_cond_with_same_area_if_none(conds, c): #TODO: handle dim != 2 def create_cond_with_same_area_if_none(conds, c): #TODO: handle dim != 2
@ -676,7 +676,7 @@ class CFGGuider:
def inner_set_conds(self, conds): def inner_set_conds(self, conds):
for k in conds: for k in conds:
self.original_conds[k] = comfy.sampler_helpers.convert_cond(conds[k]) self.original_conds[k] = seap.sampler_helpers.convert_cond(conds[k])
def __call__(self, *args, **kwargs): def __call__(self, *args, **kwargs):
return self.predict_noise(*args, **kwargs) return self.predict_noise(*args, **kwargs)
@ -703,11 +703,11 @@ class CFGGuider:
for k in self.original_conds: for k in self.original_conds:
self.conds[k] = list(map(lambda a: a.copy(), self.original_conds[k])) self.conds[k] = list(map(lambda a: a.copy(), self.original_conds[k]))
self.inner_model, self.conds, self.loaded_models = comfy.sampler_helpers.prepare_sampling(self.model_patcher, noise.shape, self.conds) self.inner_model, self.conds, self.loaded_models = seap.sampler_helpers.prepare_sampling(self.model_patcher, noise.shape, self.conds)
device = self.model_patcher.load_device device = self.model_patcher.load_device
if denoise_mask is not None: if denoise_mask is not None:
denoise_mask = comfy.sampler_helpers.prepare_mask(denoise_mask, noise.shape, device) denoise_mask = seap.sampler_helpers.prepare_mask(denoise_mask, noise.shape, device)
noise = noise.to(device) noise = noise.to(device)
latent_image = latent_image.to(device) latent_image = latent_image.to(device)
@ -715,7 +715,7 @@ class CFGGuider:
output = self.inner_sample(noise, latent_image, device, sampler, sigmas, denoise_mask, callback, disable_pbar, seed) output = self.inner_sample(noise, latent_image, device, sampler, sigmas, denoise_mask, callback, disable_pbar, seed)
comfy.sampler_helpers.cleanup_models(self.conds, self.loaded_models) seap.sampler_helpers.cleanup_models(self.conds, self.loaded_models)
del self.inner_model del self.inner_model
del self.conds del self.conds
del self.loaded_models del self.loaded_models

View File

@ -2,14 +2,14 @@ import torch
from enum import Enum from enum import Enum
import logging import logging
from comfy import model_management from seap import model_management
from .ldm.models.autoencoder import AutoencoderKL, AutoencodingEngine from .ldm.models.autoencoder import AutoencoderKL, AutoencodingEngine
from .ldm.cascade.stage_a import StageA from .ldm.cascade.stage_a import StageA
from .ldm.cascade.stage_c_coder import StageC_coder from .ldm.cascade.stage_c_coder import StageC_coder
from .ldm.audio.autoencoder import AudioOobleckVAE from .ldm.audio.autoencoder import AudioOobleckVAE
import yaml import yaml
import comfy.utils import seap.utils
from . import clip_vision from . import clip_vision
from . import gligen from . import gligen
@ -18,27 +18,27 @@ from . import model_detection
from . import sd1_clip from . import sd1_clip
from . import sdxl_clip from . import sdxl_clip
import comfy.text_encoders.sd2_clip import seap.text_encoders.sd2_clip
import comfy.text_encoders.sd3_clip import seap.text_encoders.sd3_clip
import comfy.text_encoders.sa_t5 import seap.text_encoders.sa_t5
import comfy.text_encoders.aura_t5 import seap.text_encoders.aura_t5
import comfy.text_encoders.hydit import seap.text_encoders.hydit
import comfy.text_encoders.flux import seap.text_encoders.flux
import comfy.text_encoders.long_clipl import seap.text_encoders.long_clipl
import comfy.model_patcher import seap.model_patcher
import comfy.lora import seap.lora
import comfy.t2i_adapter.adapter import seap.t2i_adapter.adapter
import comfy.taesd.taesd import seap.taesd.taesd
def load_lora_for_models(model, clip, lora, strength_model, strength_clip): def load_lora_for_models(model, clip, lora, strength_model, strength_clip):
key_map = {} key_map = {}
if model is not None: if model is not None:
key_map = comfy.lora.model_lora_keys_unet(model.model, key_map) key_map = seap.lora.model_lora_keys_unet(model.model, key_map)
if clip is not None: if clip is not None:
key_map = comfy.lora.model_lora_keys_clip(clip.cond_stage_model, key_map) key_map = seap.lora.model_lora_keys_clip(clip.cond_stage_model, key_map)
loaded = comfy.lora.load_lora(lora, key_map) loaded = seap.lora.load_lora(lora, key_map)
if model is not None: if model is not None:
new_modelpatcher = model.clone() new_modelpatcher = model.clone()
k = new_modelpatcher.add_patches(loaded, strength_model) k = new_modelpatcher.add_patches(loaded, strength_model)
@ -89,7 +89,7 @@ class CLIP:
logging.warning("Had to shift TE back.") logging.warning("Had to shift TE back.")
self.tokenizer = tokenizer(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data) self.tokenizer = tokenizer(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data)
self.patcher = comfy.model_patcher.ModelPatcher(self.cond_stage_model, load_device=load_device, offload_device=offload_device) self.patcher = seap.model_patcher.ModelPatcher(self.cond_stage_model, load_device=load_device, offload_device=offload_device)
if params['device'] == load_device: if params['device'] == load_device:
model_management.load_models_gpu([self.patcher], force_full_load=True) model_management.load_models_gpu([self.patcher], force_full_load=True)
self.layer_idx = None self.layer_idx = None
@ -180,12 +180,12 @@ class VAE:
decoder_config = encoder_config.copy() decoder_config = encoder_config.copy()
decoder_config["video_kernel_size"] = [3, 1, 1] decoder_config["video_kernel_size"] = [3, 1, 1]
decoder_config["alpha"] = 0.0 decoder_config["alpha"] = 0.0
self.first_stage_model = AutoencodingEngine(regularizer_config={'target': "comfy.ldm.models.autoencoder.DiagonalGaussianRegularizer"}, self.first_stage_model = AutoencodingEngine(regularizer_config={'target': "seap.ldm.models.autoencoder.DiagonalGaussianRegularizer"},
encoder_config={'target': "comfy.ldm.modules.diffusionmodules.model.Encoder", 'params': encoder_config}, encoder_config={'target': "seap.ldm.modules.diffusionmodules.model.Encoder", 'params': encoder_config},
decoder_config={'target': "comfy.ldm.modules.temporal_ae.VideoDecoder", 'params': decoder_config}) decoder_config={'target': "seap.ldm.modules.temporal_ae.VideoDecoder", 'params': decoder_config})
elif "taesd_decoder.1.weight" in sd: elif "taesd_decoder.1.weight" in sd:
self.latent_channels = sd["taesd_decoder.1.weight"].shape[1] self.latent_channels = sd["taesd_decoder.1.weight"].shape[1]
self.first_stage_model = comfy.taesd.taesd.TAESD(latent_channels=self.latent_channels) self.first_stage_model = seap.taesd.taesd.TAESD(latent_channels=self.latent_channels)
elif "vquantizer.codebook.weight" in sd: #VQGan: stage a of stable cascade elif "vquantizer.codebook.weight" in sd: #VQGan: stage a of stable cascade
self.first_stage_model = StageA() self.first_stage_model = StageA()
self.downscale_ratio = 4 self.downscale_ratio = 4
@ -227,9 +227,9 @@ class VAE:
if 'quant_conv.weight' in sd: if 'quant_conv.weight' in sd:
self.first_stage_model = AutoencoderKL(ddconfig=ddconfig, embed_dim=4) self.first_stage_model = AutoencoderKL(ddconfig=ddconfig, embed_dim=4)
else: else:
self.first_stage_model = AutoencodingEngine(regularizer_config={'target': "comfy.ldm.models.autoencoder.DiagonalGaussianRegularizer"}, self.first_stage_model = AutoencodingEngine(regularizer_config={'target': "seap.ldm.models.autoencoder.DiagonalGaussianRegularizer"},
encoder_config={'target': "comfy.ldm.modules.diffusionmodules.model.Encoder", 'params': ddconfig}, encoder_config={'target': "seap.ldm.modules.diffusionmodules.model.Encoder", 'params': ddconfig},
decoder_config={'target': "comfy.ldm.modules.diffusionmodules.model.Decoder", 'params': ddconfig}) decoder_config={'target': "seap.ldm.modules.diffusionmodules.model.Decoder", 'params': ddconfig})
elif "decoder.layers.1.layers.0.beta" in sd: elif "decoder.layers.1.layers.0.beta" in sd:
self.first_stage_model = AudioOobleckVAE() self.first_stage_model = AudioOobleckVAE()
self.memory_used_encode = lambda shape, dtype: (1000 * shape[2]) * model_management.dtype_size(dtype) self.memory_used_encode = lambda shape, dtype: (1000 * shape[2]) * model_management.dtype_size(dtype)
@ -266,7 +266,7 @@ class VAE:
self.first_stage_model.to(self.vae_dtype) self.first_stage_model.to(self.vae_dtype)
self.output_device = model_management.intermediate_device() self.output_device = model_management.intermediate_device()
self.patcher = comfy.model_patcher.ModelPatcher(self.first_stage_model, load_device=self.device, offload_device=offload_device) self.patcher = seap.model_patcher.ModelPatcher(self.first_stage_model, load_device=self.device, offload_device=offload_device)
logging.debug("VAE load device: {}, offload device: {}, dtype: {}".format(self.device, offload_device, self.vae_dtype)) logging.debug("VAE load device: {}, offload device: {}, dtype: {}".format(self.device, offload_device, self.vae_dtype))
def vae_encode_crop_pixels(self, pixels): def vae_encode_crop_pixels(self, pixels):
@ -279,39 +279,39 @@ class VAE:
return pixels return pixels
def decode_tiled_(self, samples, tile_x=64, tile_y=64, overlap = 16): def decode_tiled_(self, samples, tile_x=64, tile_y=64, overlap = 16):
steps = samples.shape[0] * comfy.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x, tile_y, overlap) steps = samples.shape[0] * seap.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x, tile_y, overlap)
steps += samples.shape[0] * comfy.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x // 2, tile_y * 2, overlap) steps += samples.shape[0] * seap.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x // 2, tile_y * 2, overlap)
steps += samples.shape[0] * comfy.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x * 2, tile_y // 2, overlap) steps += samples.shape[0] * seap.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x * 2, tile_y // 2, overlap)
pbar = comfy.utils.ProgressBar(steps) pbar = seap.utils.ProgressBar(steps)
decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).float() decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).float()
output = self.process_output( output = self.process_output(
(comfy.utils.tiled_scale(samples, decode_fn, tile_x // 2, tile_y * 2, overlap, upscale_amount = self.upscale_ratio, output_device=self.output_device, pbar = pbar) + (seap.utils.tiled_scale(samples, decode_fn, tile_x // 2, tile_y * 2, overlap, upscale_amount = self.upscale_ratio, output_device=self.output_device, pbar = pbar) +
comfy.utils.tiled_scale(samples, decode_fn, tile_x * 2, tile_y // 2, overlap, upscale_amount = self.upscale_ratio, output_device=self.output_device, pbar = pbar) + seap.utils.tiled_scale(samples, decode_fn, tile_x * 2, tile_y // 2, overlap, upscale_amount = self.upscale_ratio, output_device=self.output_device, pbar = pbar) +
comfy.utils.tiled_scale(samples, decode_fn, tile_x, tile_y, overlap, upscale_amount = self.upscale_ratio, output_device=self.output_device, pbar = pbar)) seap.utils.tiled_scale(samples, decode_fn, tile_x, tile_y, overlap, upscale_amount = self.upscale_ratio, output_device=self.output_device, pbar = pbar))
/ 3.0) / 3.0)
return output return output
def decode_tiled_1d(self, samples, tile_x=128, overlap=32): def decode_tiled_1d(self, samples, tile_x=128, overlap=32):
decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).float() decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).float()
return comfy.utils.tiled_scale_multidim(samples, decode_fn, tile=(tile_x,), overlap=overlap, upscale_amount=self.upscale_ratio, out_channels=self.output_channels, output_device=self.output_device) return seap.utils.tiled_scale_multidim(samples, decode_fn, tile=(tile_x,), overlap=overlap, upscale_amount=self.upscale_ratio, out_channels=self.output_channels, output_device=self.output_device)
def encode_tiled_(self, pixel_samples, tile_x=512, tile_y=512, overlap = 64): def encode_tiled_(self, pixel_samples, tile_x=512, tile_y=512, overlap = 64):
steps = pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x, tile_y, overlap) steps = pixel_samples.shape[0] * seap.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x, tile_y, overlap)
steps += pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x // 2, tile_y * 2, overlap) steps += pixel_samples.shape[0] * seap.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x // 2, tile_y * 2, overlap)
steps += pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x * 2, tile_y // 2, overlap) steps += pixel_samples.shape[0] * seap.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x * 2, tile_y // 2, overlap)
pbar = comfy.utils.ProgressBar(steps) pbar = seap.utils.ProgressBar(steps)
encode_fn = lambda a: self.first_stage_model.encode((self.process_input(a)).to(self.vae_dtype).to(self.device)).float() encode_fn = lambda a: self.first_stage_model.encode((self.process_input(a)).to(self.vae_dtype).to(self.device)).float()
samples = comfy.utils.tiled_scale(pixel_samples, encode_fn, tile_x, tile_y, overlap, upscale_amount = (1/self.downscale_ratio), out_channels=self.latent_channels, output_device=self.output_device, pbar=pbar) samples = seap.utils.tiled_scale(pixel_samples, encode_fn, tile_x, tile_y, overlap, upscale_amount = (1 / self.downscale_ratio), out_channels=self.latent_channels, output_device=self.output_device, pbar=pbar)
samples += comfy.utils.tiled_scale(pixel_samples, encode_fn, tile_x * 2, tile_y // 2, overlap, upscale_amount = (1/self.downscale_ratio), out_channels=self.latent_channels, output_device=self.output_device, pbar=pbar) samples += seap.utils.tiled_scale(pixel_samples, encode_fn, tile_x * 2, tile_y // 2, overlap, upscale_amount = (1 / self.downscale_ratio), out_channels=self.latent_channels, output_device=self.output_device, pbar=pbar)
samples += comfy.utils.tiled_scale(pixel_samples, encode_fn, tile_x // 2, tile_y * 2, overlap, upscale_amount = (1/self.downscale_ratio), out_channels=self.latent_channels, output_device=self.output_device, pbar=pbar) samples += seap.utils.tiled_scale(pixel_samples, encode_fn, tile_x // 2, tile_y * 2, overlap, upscale_amount = (1 / self.downscale_ratio), out_channels=self.latent_channels, output_device=self.output_device, pbar=pbar)
samples /= 3.0 samples /= 3.0
return samples return samples
def encode_tiled_1d(self, samples, tile_x=128 * 2048, overlap=32 * 2048): def encode_tiled_1d(self, samples, tile_x=128 * 2048, overlap=32 * 2048):
encode_fn = lambda a: self.first_stage_model.encode((self.process_input(a)).to(self.vae_dtype).to(self.device)).float() encode_fn = lambda a: self.first_stage_model.encode((self.process_input(a)).to(self.vae_dtype).to(self.device)).float()
return comfy.utils.tiled_scale_multidim(samples, encode_fn, tile=(tile_x,), overlap=overlap, upscale_amount=(1/self.downscale_ratio), out_channels=self.latent_channels, output_device=self.output_device) return seap.utils.tiled_scale_multidim(samples, encode_fn, tile=(tile_x,), overlap=overlap, upscale_amount=(1 / self.downscale_ratio), out_channels=self.latent_channels, output_device=self.output_device)
def decode(self, samples_in): def decode(self, samples_in):
try: try:
@ -382,10 +382,10 @@ class StyleModel:
def load_style_model(ckpt_path): def load_style_model(ckpt_path):
model_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) model_data = seap.utils.load_torch_file(ckpt_path, safe_load=True)
keys = model_data.keys() keys = model_data.keys()
if "style_embedding" in keys: if "style_embedding" in keys:
model = comfy.t2i_adapter.adapter.StyleAdapter(width=1024, context_dim=768, num_head=8, n_layes=3, num_token=8) model = seap.t2i_adapter.adapter.StyleAdapter(width=1024, context_dim=768, num_head=8, n_layes=3, num_token=8)
else: else:
raise Exception("invalid style model {}".format(ckpt_path)) raise Exception("invalid style model {}".format(ckpt_path))
model.load_state_dict(model_data) model.load_state_dict(model_data)
@ -402,7 +402,7 @@ class CLIPType(Enum):
def load_clip(ckpt_paths, embedding_directory=None, clip_type=CLIPType.STABLE_DIFFUSION, model_options={}): def load_clip(ckpt_paths, embedding_directory=None, clip_type=CLIPType.STABLE_DIFFUSION, model_options={}):
clip_data = [] clip_data = []
for p in ckpt_paths: for p in ckpt_paths:
clip_data.append(comfy.utils.load_torch_file(p, safe_load=True)) clip_data.append(seap.utils.load_torch_file(p, safe_load=True))
return load_text_encoder_state_dicts(clip_data, embedding_directory=embedding_directory, clip_type=clip_type, model_options=model_options) return load_text_encoder_state_dicts(clip_data, embedding_directory=embedding_directory, clip_type=clip_type, model_options=model_options)
@ -452,7 +452,7 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
for i in range(len(clip_data)): for i in range(len(clip_data)):
if "transformer.resblocks.0.ln_1.weight" in clip_data[i]: if "transformer.resblocks.0.ln_1.weight" in clip_data[i]:
clip_data[i] = comfy.utils.clip_text_transformers_convert(clip_data[i], "", "") clip_data[i] = seap.utils.clip_text_transformers_convert(clip_data[i], "", "")
else: else:
if "text_projection" in clip_data[i]: if "text_projection" in clip_data[i]:
clip_data[i]["text_projection.weight"] = clip_data[i]["text_projection"].transpose(0, 1) #old models saved with the CLIPSave node clip_data[i]["text_projection.weight"] = clip_data[i]["text_projection"].transpose(0, 1) #old models saved with the CLIPSave node
@ -466,53 +466,53 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
clip_target.clip = sdxl_clip.StableCascadeClipModel clip_target.clip = sdxl_clip.StableCascadeClipModel
clip_target.tokenizer = sdxl_clip.StableCascadeTokenizer clip_target.tokenizer = sdxl_clip.StableCascadeTokenizer
elif clip_type == CLIPType.SD3: elif clip_type == CLIPType.SD3:
clip_target.clip = comfy.text_encoders.sd3_clip.sd3_clip(clip_l=False, clip_g=True, t5=False) clip_target.clip = seap.text_encoders.sd3_clip.sd3_clip(clip_l=False, clip_g=True, t5=False)
clip_target.tokenizer = comfy.text_encoders.sd3_clip.SD3Tokenizer clip_target.tokenizer = seap.text_encoders.sd3_clip.SD3Tokenizer
else: else:
clip_target.clip = sdxl_clip.SDXLRefinerClipModel clip_target.clip = sdxl_clip.SDXLRefinerClipModel
clip_target.tokenizer = sdxl_clip.SDXLTokenizer clip_target.tokenizer = sdxl_clip.SDXLTokenizer
elif te_model == TEModel.CLIP_H: elif te_model == TEModel.CLIP_H:
clip_target.clip = comfy.text_encoders.sd2_clip.SD2ClipModel clip_target.clip = seap.text_encoders.sd2_clip.SD2ClipModel
clip_target.tokenizer = comfy.text_encoders.sd2_clip.SD2Tokenizer clip_target.tokenizer = seap.text_encoders.sd2_clip.SD2Tokenizer
elif te_model == TEModel.T5_XXL: elif te_model == TEModel.T5_XXL:
clip_target.clip = comfy.text_encoders.sd3_clip.sd3_clip(clip_l=False, clip_g=False, t5=True, dtype_t5=t5xxl_weight_dtype(clip_data)) clip_target.clip = seap.text_encoders.sd3_clip.sd3_clip(clip_l=False, clip_g=False, t5=True, dtype_t5=t5xxl_weight_dtype(clip_data))
clip_target.tokenizer = comfy.text_encoders.sd3_clip.SD3Tokenizer clip_target.tokenizer = seap.text_encoders.sd3_clip.SD3Tokenizer
elif te_model == TEModel.T5_XL: elif te_model == TEModel.T5_XL:
clip_target.clip = comfy.text_encoders.aura_t5.AuraT5Model clip_target.clip = seap.text_encoders.aura_t5.AuraT5Model
clip_target.tokenizer = comfy.text_encoders.aura_t5.AuraT5Tokenizer clip_target.tokenizer = seap.text_encoders.aura_t5.AuraT5Tokenizer
elif te_model == TEModel.T5_BASE: elif te_model == TEModel.T5_BASE:
clip_target.clip = comfy.text_encoders.sa_t5.SAT5Model clip_target.clip = seap.text_encoders.sa_t5.SAT5Model
clip_target.tokenizer = comfy.text_encoders.sa_t5.SAT5Tokenizer clip_target.tokenizer = seap.text_encoders.sa_t5.SAT5Tokenizer
else: else:
if clip_type == CLIPType.SD3: if clip_type == CLIPType.SD3:
clip_target.clip = comfy.text_encoders.sd3_clip.sd3_clip(clip_l=True, clip_g=False, t5=False) clip_target.clip = seap.text_encoders.sd3_clip.sd3_clip(clip_l=True, clip_g=False, t5=False)
clip_target.tokenizer = comfy.text_encoders.sd3_clip.SD3Tokenizer clip_target.tokenizer = seap.text_encoders.sd3_clip.SD3Tokenizer
else: else:
clip_target.clip = sd1_clip.SD1ClipModel clip_target.clip = sd1_clip.SD1ClipModel
clip_target.tokenizer = sd1_clip.SD1Tokenizer clip_target.tokenizer = sd1_clip.SD1Tokenizer
elif len(clip_data) == 2: elif len(clip_data) == 2:
if clip_type == CLIPType.SD3: if clip_type == CLIPType.SD3:
te_models = [detect_te_model(clip_data[0]), detect_te_model(clip_data[1])] te_models = [detect_te_model(clip_data[0]), detect_te_model(clip_data[1])]
clip_target.clip = comfy.text_encoders.sd3_clip.sd3_clip(clip_l=TEModel.CLIP_L in te_models, clip_g=TEModel.CLIP_G in te_models, t5=TEModel.T5_XXL in te_models, dtype_t5=t5xxl_weight_dtype(clip_data)) clip_target.clip = seap.text_encoders.sd3_clip.sd3_clip(clip_l=TEModel.CLIP_L in te_models, clip_g=TEModel.CLIP_G in te_models, t5=TEModel.T5_XXL in te_models, dtype_t5=t5xxl_weight_dtype(clip_data))
clip_target.tokenizer = comfy.text_encoders.sd3_clip.SD3Tokenizer clip_target.tokenizer = seap.text_encoders.sd3_clip.SD3Tokenizer
elif clip_type == CLIPType.HUNYUAN_DIT: elif clip_type == CLIPType.HUNYUAN_DIT:
clip_target.clip = comfy.text_encoders.hydit.HyditModel clip_target.clip = seap.text_encoders.hydit.HyditModel
clip_target.tokenizer = comfy.text_encoders.hydit.HyditTokenizer clip_target.tokenizer = seap.text_encoders.hydit.HyditTokenizer
elif clip_type == CLIPType.FLUX: elif clip_type == CLIPType.FLUX:
clip_target.clip = comfy.text_encoders.flux.flux_clip(dtype_t5=t5xxl_weight_dtype(clip_data)) clip_target.clip = seap.text_encoders.flux.flux_clip(dtype_t5=t5xxl_weight_dtype(clip_data))
clip_target.tokenizer = comfy.text_encoders.flux.FluxTokenizer clip_target.tokenizer = seap.text_encoders.flux.FluxTokenizer
else: else:
clip_target.clip = sdxl_clip.SDXLClipModel clip_target.clip = sdxl_clip.SDXLClipModel
clip_target.tokenizer = sdxl_clip.SDXLTokenizer clip_target.tokenizer = sdxl_clip.SDXLTokenizer
elif len(clip_data) == 3: elif len(clip_data) == 3:
clip_target.clip = comfy.text_encoders.sd3_clip.sd3_clip(dtype_t5=t5xxl_weight_dtype(clip_data)) clip_target.clip = seap.text_encoders.sd3_clip.sd3_clip(dtype_t5=t5xxl_weight_dtype(clip_data))
clip_target.tokenizer = comfy.text_encoders.sd3_clip.SD3Tokenizer clip_target.tokenizer = seap.text_encoders.sd3_clip.SD3Tokenizer
parameters = 0 parameters = 0
tokenizer_data = {} tokenizer_data = {}
for c in clip_data: for c in clip_data:
parameters += comfy.utils.calculate_parameters(c) parameters += seap.utils.calculate_parameters(c)
tokenizer_data, model_options = comfy.text_encoders.long_clipl.model_options_long_clip(c, tokenizer_data, model_options) tokenizer_data, model_options = seap.text_encoders.long_clipl.model_options_long_clip(c, tokenizer_data, model_options)
clip = CLIP(clip_target, embedding_directory=embedding_directory, parameters=parameters, tokenizer_data=tokenizer_data, model_options=model_options) clip = CLIP(clip_target, embedding_directory=embedding_directory, parameters=parameters, tokenizer_data=tokenizer_data, model_options=model_options)
for c in clip_data: for c in clip_data:
@ -525,11 +525,11 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
return clip return clip
def load_gligen(ckpt_path): def load_gligen(ckpt_path):
data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) data = seap.utils.load_torch_file(ckpt_path, safe_load=True)
model = gligen.load_gligen(data) model = gligen.load_gligen(data)
if model_management.should_use_fp16(): if model_management.should_use_fp16():
model = model.half() model = model.half()
return comfy.model_patcher.ModelPatcher(model, load_device=model_management.get_torch_device(), offload_device=model_management.unet_offload_device()) return seap.model_patcher.ModelPatcher(model, load_device=model_management.get_torch_device(), offload_device=model_management.unet_offload_device())
def load_checkpoint(config_path=None, ckpt_path=None, output_vae=True, output_clip=True, embedding_directory=None, state_dict=None, config=None): def load_checkpoint(config_path=None, ckpt_path=None, output_vae=True, output_clip=True, embedding_directory=None, state_dict=None, config=None):
logging.warning("Warning: The load checkpoint with config function is deprecated and will eventually be removed, please use the other one.") logging.warning("Warning: The load checkpoint with config function is deprecated and will eventually be removed, please use the other one.")
@ -545,7 +545,7 @@ def load_checkpoint(config_path=None, ckpt_path=None, output_vae=True, output_cl
if "parameterization" in model_config_params: if "parameterization" in model_config_params:
if model_config_params["parameterization"] == "v": if model_config_params["parameterization"] == "v":
m = model.clone() m = model.clone()
class ModelSamplingAdvanced(comfy.model_sampling.ModelSamplingDiscrete, comfy.model_sampling.V_PREDICTION): class ModelSamplingAdvanced(seap.model_sampling.ModelSamplingDiscrete, seap.model_sampling.V_PREDICTION):
pass pass
m.add_object_patch("model_sampling", ModelSamplingAdvanced(model.model.model_config)) m.add_object_patch("model_sampling", ModelSamplingAdvanced(model.model.model_config))
model = m model = m
@ -557,7 +557,7 @@ def load_checkpoint(config_path=None, ckpt_path=None, output_vae=True, output_cl
return (model, clip, vae) return (model, clip, vae)
def load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, output_clipvision=False, embedding_directory=None, output_model=True, model_options={}, te_model_options={}): def load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, output_clipvision=False, embedding_directory=None, output_model=True, model_options={}, te_model_options={}):
sd = comfy.utils.load_torch_file(ckpt_path) sd = seap.utils.load_torch_file(ckpt_path)
out = load_state_dict_guess_config(sd, output_vae, output_clip, output_clipvision, embedding_directory, output_model, model_options, te_model_options=te_model_options) out = load_state_dict_guess_config(sd, output_vae, output_clip, output_clipvision, embedding_directory, output_model, model_options, te_model_options=te_model_options)
if out is None: if out is None:
raise RuntimeError("ERROR: Could not detect model type of: {}".format(ckpt_path)) raise RuntimeError("ERROR: Could not detect model type of: {}".format(ckpt_path))
@ -571,8 +571,8 @@ def load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_c
model_patcher = None model_patcher = None
diffusion_model_prefix = model_detection.unet_prefix_from_state_dict(sd) diffusion_model_prefix = model_detection.unet_prefix_from_state_dict(sd)
parameters = comfy.utils.calculate_parameters(sd, diffusion_model_prefix) parameters = seap.utils.calculate_parameters(sd, diffusion_model_prefix)
weight_dtype = comfy.utils.weight_dtype(sd, diffusion_model_prefix) weight_dtype = seap.utils.weight_dtype(sd, diffusion_model_prefix)
load_device = model_management.get_torch_device() load_device = model_management.get_torch_device()
model_config = model_detection.model_config_from_unet(sd, diffusion_model_prefix) model_config = model_detection.model_config_from_unet(sd, diffusion_model_prefix)
@ -602,7 +602,7 @@ def load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_c
model.load_model_weights(sd, diffusion_model_prefix) model.load_model_weights(sd, diffusion_model_prefix)
if output_vae: if output_vae:
vae_sd = comfy.utils.state_dict_prefix_replace(sd, {k: "" for k in model_config.vae_key_prefix}, filter_keys=True) vae_sd = seap.utils.state_dict_prefix_replace(sd, {k: "" for k in model_config.vae_key_prefix}, filter_keys=True)
vae_sd = model_config.process_vae_state_dict(vae_sd) vae_sd = model_config.process_vae_state_dict(vae_sd)
vae = VAE(sd=vae_sd) vae = VAE(sd=vae_sd)
@ -611,7 +611,7 @@ def load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_c
if clip_target is not None: if clip_target is not None:
clip_sd = model_config.process_clip_state_dict(sd) clip_sd = model_config.process_clip_state_dict(sd)
if len(clip_sd) > 0: if len(clip_sd) > 0:
parameters = comfy.utils.calculate_parameters(clip_sd) parameters = seap.utils.calculate_parameters(clip_sd)
clip = CLIP(clip_target, embedding_directory=embedding_directory, tokenizer_data=clip_sd, parameters=parameters, model_options=te_model_options) clip = CLIP(clip_target, embedding_directory=embedding_directory, tokenizer_data=clip_sd, parameters=parameters, model_options=te_model_options)
m, u = clip.load_sd(clip_sd, full_model=True) m, u = clip.load_sd(clip_sd, full_model=True)
if len(m) > 0: if len(m) > 0:
@ -631,7 +631,7 @@ def load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_c
logging.debug("left over keys: {}".format(left_over)) logging.debug("left over keys: {}".format(left_over))
if output_model: if output_model:
model_patcher = comfy.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=model_management.unet_offload_device()) model_patcher = seap.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=model_management.unet_offload_device())
if inital_load_device != torch.device("cpu"): if inital_load_device != torch.device("cpu"):
logging.info("loaded straight to GPU") logging.info("loaded straight to GPU")
model_management.load_models_gpu([model_patcher], force_full_load=True) model_management.load_models_gpu([model_patcher], force_full_load=True)
@ -644,11 +644,11 @@ def load_diffusion_model_state_dict(sd, model_options={}): #load unet in diffuse
#Allow loading unets from checkpoint files #Allow loading unets from checkpoint files
diffusion_model_prefix = model_detection.unet_prefix_from_state_dict(sd) diffusion_model_prefix = model_detection.unet_prefix_from_state_dict(sd)
temp_sd = comfy.utils.state_dict_prefix_replace(sd, {diffusion_model_prefix: ""}, filter_keys=True) temp_sd = seap.utils.state_dict_prefix_replace(sd, {diffusion_model_prefix: ""}, filter_keys=True)
if len(temp_sd) > 0: if len(temp_sd) > 0:
sd = temp_sd sd = temp_sd
parameters = comfy.utils.calculate_parameters(sd) parameters = seap.utils.calculate_parameters(sd)
load_device = model_management.get_torch_device() load_device = model_management.get_torch_device()
model_config = model_detection.model_config_from_unet(sd, "") model_config = model_detection.model_config_from_unet(sd, "")
@ -665,7 +665,7 @@ def load_diffusion_model_state_dict(sd, model_options={}): #load unet in diffuse
if model_config is None: if model_config is None:
return None return None
diffusers_keys = comfy.utils.unet_to_diffusers(model_config.unet_config) diffusers_keys = seap.utils.unet_to_diffusers(model_config.unet_config)
new_sd = {} new_sd = {}
for k in diffusers_keys: for k in diffusers_keys:
@ -692,11 +692,11 @@ def load_diffusion_model_state_dict(sd, model_options={}): #load unet in diffuse
left_over = sd.keys() left_over = sd.keys()
if len(left_over) > 0: if len(left_over) > 0:
logging.info("left over keys in unet: {}".format(left_over)) logging.info("left over keys in unet: {}".format(left_over))
return comfy.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=offload_device) return seap.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=offload_device)
def load_diffusion_model(unet_path, model_options={}): def load_diffusion_model(unet_path, model_options={}):
sd = comfy.utils.load_torch_file(unet_path) sd = seap.utils.load_torch_file(unet_path)
model = load_diffusion_model_state_dict(sd, model_options=model_options) model = load_diffusion_model_state_dict(sd, model_options=model_options)
if model is None: if model is None:
logging.error("ERROR UNSUPPORTED UNET {}".format(unet_path)) logging.error("ERROR UNSUPPORTED UNET {}".format(unet_path))
@ -732,4 +732,4 @@ def save_checkpoint(output_path, model, clip=None, vae=None, clip_vision=None, m
if not t.is_contiguous(): if not t.is_contiguous():
sd[k] = t.contiguous() sd[k] = t.contiguous()
comfy.utils.save_torch_file(sd, output_path, metadata=metadata) seap.utils.save_torch_file(sd, output_path, metadata=metadata)

View File

@ -1,12 +1,12 @@
import os import os
from transformers import CLIPTokenizer from transformers import CLIPTokenizer
import comfy.ops import seap.ops
import torch import torch
import traceback import traceback
import zipfile import zipfile
from . import model_management from . import model_management
import comfy.clip_model import seap.clip_model
import json import json
import logging import logging
import numbers import numbers
@ -81,7 +81,7 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
"hidden" "hidden"
] ]
def __init__(self, device="cpu", max_length=77, def __init__(self, device="cpu", max_length=77,
freeze=True, layer="last", layer_idx=None, textmodel_json_config=None, dtype=None, model_class=comfy.clip_model.CLIPTextModel, freeze=True, layer="last", layer_idx=None, textmodel_json_config=None, dtype=None, model_class=seap.clip_model.CLIPTextModel,
special_tokens={"start": 49406, "end": 49407, "pad": 49407}, layer_norm_hidden_state=True, enable_attention_masks=False, zero_out_masked=False, special_tokens={"start": 49406, "end": 49407, "pad": 49407}, layer_norm_hidden_state=True, enable_attention_masks=False, zero_out_masked=False,
return_projected_pooled=True, return_attention_masks=False, model_options={}): # clip-vit-base-patch32 return_projected_pooled=True, return_attention_masks=False, model_options={}): # clip-vit-base-patch32
super().__init__() super().__init__()
@ -95,7 +95,7 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
operations = model_options.get("custom_operations", None) operations = model_options.get("custom_operations", None)
if operations is None: if operations is None:
operations = comfy.ops.manual_cast operations = seap.ops.manual_cast
self.operations = operations self.operations = operations
self.transformer = model_class(config, dtype, device, self.operations) self.transformer = model_class(config, dtype, device, self.operations)

View File

@ -1,4 +1,4 @@
from comfy import sd1_clip from seap import sd1_clip
import torch import torch
import os import os

View File

@ -4,12 +4,12 @@ from . import utils
from . import sd1_clip from . import sd1_clip
from . import sdxl_clip from . import sdxl_clip
import comfy.text_encoders.sd2_clip import seap.text_encoders.sd2_clip
import comfy.text_encoders.sd3_clip import seap.text_encoders.sd3_clip
import comfy.text_encoders.sa_t5 import seap.text_encoders.sa_t5
import comfy.text_encoders.aura_t5 import seap.text_encoders.aura_t5
import comfy.text_encoders.hydit import seap.text_encoders.hydit
import comfy.text_encoders.flux import seap.text_encoders.flux
from . import supported_models_base from . import supported_models_base
from . import latent_formats from . import latent_formats
@ -104,7 +104,7 @@ class SD20(supported_models_base.BASE):
return state_dict return state_dict
def clip_target(self, state_dict={}): def clip_target(self, state_dict={}):
return supported_models_base.ClipTarget(comfy.text_encoders.sd2_clip.SD2Tokenizer, comfy.text_encoders.sd2_clip.SD2ClipModel) return supported_models_base.ClipTarget(seap.text_encoders.sd2_clip.SD2Tokenizer, seap.text_encoders.sd2_clip.SD2ClipModel)
class SD21UnclipL(SD20): class SD21UnclipL(SD20):
unet_config = { unet_config = {
@ -534,7 +534,7 @@ class SD3(supported_models_base.BASE):
t5 = True t5 = True
dtype_t5 = state_dict[t5_key].dtype dtype_t5 = state_dict[t5_key].dtype
return supported_models_base.ClipTarget(comfy.text_encoders.sd3_clip.SD3Tokenizer, comfy.text_encoders.sd3_clip.sd3_clip(clip_l=clip_l, clip_g=clip_g, t5=t5, dtype_t5=dtype_t5)) return supported_models_base.ClipTarget(seap.text_encoders.sd3_clip.SD3Tokenizer, seap.text_encoders.sd3_clip.sd3_clip(clip_l=clip_l, clip_g=clip_g, t5=t5, dtype_t5=dtype_t5))
class StableAudio(supported_models_base.BASE): class StableAudio(supported_models_base.BASE):
unet_config = { unet_config = {
@ -565,7 +565,7 @@ class StableAudio(supported_models_base.BASE):
return utils.state_dict_prefix_replace(state_dict, replace_prefix) return utils.state_dict_prefix_replace(state_dict, replace_prefix)
def clip_target(self, state_dict={}): def clip_target(self, state_dict={}):
return supported_models_base.ClipTarget(comfy.text_encoders.sa_t5.SAT5Tokenizer, comfy.text_encoders.sa_t5.SAT5Model) return supported_models_base.ClipTarget(seap.text_encoders.sa_t5.SAT5Tokenizer, seap.text_encoders.sa_t5.SAT5Model)
class AuraFlow(supported_models_base.BASE): class AuraFlow(supported_models_base.BASE):
unet_config = { unet_config = {
@ -588,7 +588,7 @@ class AuraFlow(supported_models_base.BASE):
return out return out
def clip_target(self, state_dict={}): def clip_target(self, state_dict={}):
return supported_models_base.ClipTarget(comfy.text_encoders.aura_t5.AuraT5Tokenizer, comfy.text_encoders.aura_t5.AuraT5Model) return supported_models_base.ClipTarget(seap.text_encoders.aura_t5.AuraT5Tokenizer, seap.text_encoders.aura_t5.AuraT5Model)
class HunyuanDiT(supported_models_base.BASE): class HunyuanDiT(supported_models_base.BASE):
unet_config = { unet_config = {
@ -614,7 +614,7 @@ class HunyuanDiT(supported_models_base.BASE):
return out return out
def clip_target(self, state_dict={}): def clip_target(self, state_dict={}):
return supported_models_base.ClipTarget(comfy.text_encoders.hydit.HyditTokenizer, comfy.text_encoders.hydit.HyditModel) return supported_models_base.ClipTarget(seap.text_encoders.hydit.HyditTokenizer, seap.text_encoders.hydit.HyditModel)
class HunyuanDiT1(HunyuanDiT): class HunyuanDiT1(HunyuanDiT):
unet_config = { unet_config = {
@ -657,7 +657,7 @@ class Flux(supported_models_base.BASE):
dtype_t5 = None dtype_t5 = None
if t5_key in state_dict: if t5_key in state_dict:
dtype_t5 = state_dict[t5_key].dtype dtype_t5 = state_dict[t5_key].dtype
return supported_models_base.ClipTarget(comfy.text_encoders.flux.FluxTokenizer, comfy.text_encoders.flux.flux_clip(dtype_t5=dtype_t5)) return supported_models_base.ClipTarget(seap.text_encoders.flux.FluxTokenizer, seap.text_encoders.flux.flux_clip(dtype_t5=dtype_t5))
class FluxSchnell(Flux): class FluxSchnell(Flux):
unet_config = { unet_config = {

View File

@ -6,11 +6,11 @@ Tiny AutoEncoder for Stable Diffusion
import torch import torch
import torch.nn as nn import torch.nn as nn
import comfy.utils import seap.utils
import comfy.ops import seap.ops
def conv(n_in, n_out, **kwargs): def conv(n_in, n_out, **kwargs):
return comfy.ops.disable_weight_init.Conv2d(n_in, n_out, 3, padding=1, **kwargs) return seap.ops.disable_weight_init.Conv2d(n_in, n_out, 3, padding=1, **kwargs)
class Clamp(nn.Module): class Clamp(nn.Module):
def forward(self, x): def forward(self, x):
@ -20,7 +20,7 @@ class Block(nn.Module):
def __init__(self, n_in, n_out): def __init__(self, n_in, n_out):
super().__init__() super().__init__()
self.conv = nn.Sequential(conv(n_in, n_out), nn.ReLU(), conv(n_out, n_out), nn.ReLU(), conv(n_out, n_out)) self.conv = nn.Sequential(conv(n_in, n_out), nn.ReLU(), conv(n_out, n_out), nn.ReLU(), conv(n_out, n_out))
self.skip = comfy.ops.disable_weight_init.Conv2d(n_in, n_out, 1, bias=False) if n_in != n_out else nn.Identity() self.skip = seap.ops.disable_weight_init.Conv2d(n_in, n_out, 1, bias=False) if n_in != n_out else nn.Identity()
self.fuse = nn.ReLU() self.fuse = nn.ReLU()
def forward(self, x): def forward(self, x):
return self.fuse(self.conv(x) + self.skip(x)) return self.fuse(self.conv(x) + self.skip(x))
@ -56,9 +56,9 @@ class TAESD(nn.Module):
self.vae_scale = torch.nn.Parameter(torch.tensor(1.0)) self.vae_scale = torch.nn.Parameter(torch.tensor(1.0))
self.vae_shift = torch.nn.Parameter(torch.tensor(0.0)) self.vae_shift = torch.nn.Parameter(torch.tensor(0.0))
if encoder_path is not None: if encoder_path is not None:
self.taesd_encoder.load_state_dict(comfy.utils.load_torch_file(encoder_path, safe_load=True)) self.taesd_encoder.load_state_dict(seap.utils.load_torch_file(encoder_path, safe_load=True))
if decoder_path is not None: if decoder_path is not None:
self.taesd_decoder.load_state_dict(comfy.utils.load_torch_file(decoder_path, safe_load=True)) self.taesd_decoder.load_state_dict(seap.utils.load_torch_file(decoder_path, safe_load=True))
@staticmethod @staticmethod
def scale_latents(x): def scale_latents(x):

View File

@ -1,12 +1,12 @@
from comfy import sd1_clip from seap import sd1_clip
from .spiece_tokenizer import SPieceTokenizer from .spiece_tokenizer import SPieceTokenizer
import comfy.text_encoders.t5 import seap.text_encoders.t5
import os import os
class PT5XlModel(sd1_clip.SDClipModel): class PT5XlModel(sd1_clip.SDClipModel):
def __init__(self, device="cpu", layer="last", layer_idx=None, dtype=None, model_options={}): def __init__(self, device="cpu", layer="last", layer_idx=None, dtype=None, model_options={}):
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "t5_pile_config_xl.json") textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "t5_pile_config_xl.json")
super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, dtype=dtype, special_tokens={"end": 2, "pad": 1}, model_class=comfy.text_encoders.t5.T5, enable_attention_masks=True, zero_out_masked=True, model_options=model_options) super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, dtype=dtype, special_tokens={"end": 2, "pad": 1}, model_class=seap.text_encoders.t5.T5, enable_attention_masks=True, zero_out_masked=True, model_options=model_options)
class PT5XlTokenizer(sd1_clip.SDTokenizer): class PT5XlTokenizer(sd1_clip.SDTokenizer):
def __init__(self, embedding_directory=None, tokenizer_data={}): def __init__(self, embedding_directory=None, tokenizer_data={}):

View File

@ -1,6 +1,6 @@
import torch import torch
from comfy.ldm.modules.attention import optimized_attention_for_device from seap.ldm.modules.attention import optimized_attention_for_device
import comfy.ops import seap.ops
class BertAttention(torch.nn.Module): class BertAttention(torch.nn.Module):
def __init__(self, embed_dim, heads, dtype, device, operations): def __init__(self, embed_dim, heads, dtype, device, operations):
@ -95,11 +95,11 @@ class BertEmbeddings(torch.nn.Module):
def forward(self, input_tokens, token_type_ids=None, dtype=None): def forward(self, input_tokens, token_type_ids=None, dtype=None):
x = self.word_embeddings(input_tokens, out_dtype=dtype) x = self.word_embeddings(input_tokens, out_dtype=dtype)
x += comfy.ops.cast_to_input(self.position_embeddings.weight[:x.shape[1]], x) x += seap.ops.cast_to_input(self.position_embeddings.weight[:x.shape[1]], x)
if token_type_ids is not None: if token_type_ids is not None:
x += self.token_type_embeddings(token_type_ids, out_dtype=x.dtype) x += self.token_type_embeddings(token_type_ids, out_dtype=x.dtype)
else: else:
x += comfy.ops.cast_to_input(self.token_type_embeddings.weight[0], x) x += seap.ops.cast_to_input(self.token_type_embeddings.weight[0], x)
x = self.LayerNorm(x) x = self.LayerNorm(x)
return x return x

View File

@ -1,6 +1,6 @@
from comfy import sd1_clip from seap import sd1_clip
import comfy.text_encoders.t5 import seap.text_encoders.t5
import comfy.model_management import seap.model_management
from transformers import T5TokenizerFast from transformers import T5TokenizerFast
import torch import torch
import os import os
@ -8,7 +8,7 @@ import os
class T5XXLModel(sd1_clip.SDClipModel): class T5XXLModel(sd1_clip.SDClipModel):
def __init__(self, device="cpu", layer="last", layer_idx=None, dtype=None, model_options={}): def __init__(self, device="cpu", layer="last", layer_idx=None, dtype=None, model_options={}):
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "t5_config_xxl.json") textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "t5_config_xxl.json")
super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, dtype=dtype, special_tokens={"end": 1, "pad": 0}, model_class=comfy.text_encoders.t5.T5, model_options=model_options) super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, dtype=dtype, special_tokens={"end": 1, "pad": 0}, model_class=seap.text_encoders.t5.T5, model_options=model_options)
class T5XXLTokenizer(sd1_clip.SDTokenizer): class T5XXLTokenizer(sd1_clip.SDTokenizer):
def __init__(self, embedding_directory=None, tokenizer_data={}): def __init__(self, embedding_directory=None, tokenizer_data={}):
@ -38,7 +38,7 @@ class FluxTokenizer:
class FluxClipModel(torch.nn.Module): class FluxClipModel(torch.nn.Module):
def __init__(self, dtype_t5=None, device="cpu", dtype=None, model_options={}): def __init__(self, dtype_t5=None, device="cpu", dtype=None, model_options={}):
super().__init__() super().__init__()
dtype_t5 = comfy.model_management.pick_weight_dtype(dtype_t5, dtype, device) dtype_t5 = seap.model_management.pick_weight_dtype(dtype_t5, dtype, device)
clip_l_class = model_options.get("clip_l_class", sd1_clip.SDClipModel) clip_l_class = model_options.get("clip_l_class", sd1_clip.SDClipModel)
self.clip_l = clip_l_class(device=device, dtype=dtype, return_projected_pooled=False, model_options=model_options) self.clip_l = clip_l_class(device=device, dtype=dtype, return_projected_pooled=False, model_options=model_options)
self.t5xxl = T5XXLModel(device=device, dtype=dtype_t5, model_options=model_options) self.t5xxl = T5XXLModel(device=device, dtype=dtype_t5, model_options=model_options)

View File

@ -1,8 +1,8 @@
from comfy import sd1_clip from seap import sd1_clip
from transformers import BertTokenizer from transformers import BertTokenizer
from .spiece_tokenizer import SPieceTokenizer from .spiece_tokenizer import SPieceTokenizer
from .bert import BertModel from .bert import BertModel
import comfy.text_encoders.t5 import seap.text_encoders.t5
import os import os
import torch import torch
@ -20,7 +20,7 @@ class HyditBertTokenizer(sd1_clip.SDTokenizer):
class MT5XLModel(sd1_clip.SDClipModel): class MT5XLModel(sd1_clip.SDClipModel):
def __init__(self, device="cpu", layer="last", layer_idx=None, dtype=None, model_options={}): def __init__(self, device="cpu", layer="last", layer_idx=None, dtype=None, model_options={}):
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "mt5_config_xl.json") textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "mt5_config_xl.json")
super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, dtype=dtype, special_tokens={"end": 1, "pad": 0}, model_class=comfy.text_encoders.t5.T5, enable_attention_masks=True, return_attention_masks=True, model_options=model_options) super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, dtype=dtype, special_tokens={"end": 1, "pad": 0}, model_class=seap.text_encoders.t5.T5, enable_attention_masks=True, return_attention_masks=True, model_options=model_options)
class MT5XLTokenizer(sd1_clip.SDTokenizer): class MT5XLTokenizer(sd1_clip.SDTokenizer):
def __init__(self, embedding_directory=None, tokenizer_data={}): def __init__(self, embedding_directory=None, tokenizer_data={}):

Some files were not shown because too many files have changed in this diff Show More