feat: create architecture with diffusers pipeline
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
from .src.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
|
||||
+34
-3
@@ -18,7 +18,9 @@ dependencies = [
|
||||
]
|
||||
|
||||
[tool.poetry]
|
||||
packages = [{include = "comfyui_diffusers", from = "src"}]
|
||||
packages = [
|
||||
{include = "src"}
|
||||
]
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 100
|
||||
@@ -45,7 +47,7 @@ ignore = [
|
||||
"__init__.py" = ["F401"] # unused imports in __init__.py
|
||||
|
||||
[tool.ruff.lint.isort]
|
||||
known-first-party = ["comfyui_diffusers"]
|
||||
known-first-party = ["src"]
|
||||
|
||||
[tool.ruff.format]
|
||||
quote-style = "double"
|
||||
@@ -53,7 +55,8 @@ indent-style = "space"
|
||||
|
||||
[tool.mypy]
|
||||
python_version = "3.10"
|
||||
namespace_packages = false
|
||||
namespace_packages = true
|
||||
explicit_package_bases = true
|
||||
strict = true
|
||||
warn_return_any = true
|
||||
warn_unused_configs = true
|
||||
@@ -67,6 +70,34 @@ warn_no_return = true
|
||||
warn_unreachable = true
|
||||
strict_equality = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = [
|
||||
"folder_paths",
|
||||
"folder_paths.*",
|
||||
]
|
||||
ignore_missing_imports = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = [
|
||||
"diffusers",
|
||||
"diffusers.*",
|
||||
]
|
||||
ignore_missing_imports = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = [
|
||||
"dependency_injector",
|
||||
"dependency_injector.*",
|
||||
]
|
||||
ignore_missing_imports = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = [
|
||||
"transformers",
|
||||
"transformers.*",
|
||||
]
|
||||
ignore_missing_imports = true
|
||||
|
||||
[build-system]
|
||||
requires = ["poetry-core>=2.0.0,<3.0.0"]
|
||||
build-backend = "poetry.core.masonry.api"
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from .di import Container
|
||||
|
||||
__all__ = ["Container"]
|
||||
@@ -0,0 +1,22 @@
|
||||
from dependency_injector import containers, providers
|
||||
|
||||
from ..infra.diffusers.pipeline_repository import DiffusersPipelineRepository
|
||||
from ..service.create_pipeline_service import CreatePipelineService
|
||||
from ..ui.pipeline_handler import PipelineHandler
|
||||
|
||||
|
||||
class Container(containers.DeclarativeContainer):
|
||||
# 1. Repositories
|
||||
pipeline_repository_provider = providers.Factory(DiffusersPipelineRepository)
|
||||
|
||||
# 2. Services
|
||||
create_pipeline_service_provider = providers.Factory(
|
||||
CreatePipelineService,
|
||||
pipeline_repo=pipeline_repository_provider,
|
||||
)
|
||||
|
||||
# 3. Handler
|
||||
pipeline_handler_provider = providers.Factory(
|
||||
PipelineHandler,
|
||||
create_pipeline_service=create_pipeline_service_provider,
|
||||
)
|
||||
@@ -0,0 +1,3 @@
|
||||
from .pipeline import PipelineModel
|
||||
|
||||
__all__ = ["PipelineModel"]
|
||||
@@ -0,0 +1,9 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from diffusers import StableDiffusionPipeline
|
||||
|
||||
|
||||
@dataclass
|
||||
class PipelineModel:
|
||||
pipeline: StableDiffusionPipeline
|
||||
path: str
|
||||
@@ -0,0 +1,3 @@
|
||||
from .i_pipeline_repository import IPipelineRepository
|
||||
|
||||
__all__ = ["IPipelineRepository"]
|
||||
@@ -0,0 +1,16 @@
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
import torch
|
||||
from diffusers import StableDiffusionPipeline
|
||||
|
||||
|
||||
class IPipelineRepository(ABC):
|
||||
@abstractmethod
|
||||
def convert_from_single_file(self, checkpoint_path: str, dtype: torch.dtype) -> str:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def load_pipeline_from_path(
|
||||
self, model_path: str, dtype: torch.dtype
|
||||
) -> StableDiffusionPipeline:
|
||||
pass
|
||||
@@ -0,0 +1,3 @@
|
||||
from .pipeline_repository import DiffusersPipelineRepository
|
||||
|
||||
__all__ = ["DiffusersPipelineRepository"]
|
||||
@@ -0,0 +1,33 @@
|
||||
import os
|
||||
|
||||
import folder_paths # pyright: ignore[reportMissingImports]
|
||||
import torch
|
||||
from diffusers import StableDiffusionPipeline
|
||||
|
||||
from ...domain.repositories import IPipelineRepository
|
||||
|
||||
|
||||
class DiffusersPipelineRepository(IPipelineRepository):
|
||||
def __init__(self) -> None:
|
||||
self.tmp_dir = folder_paths.get_temp_directory()
|
||||
|
||||
def convert_from_single_file(self, checkpoint_path: str, dtype: torch.dtype) -> str:
|
||||
checkpoint_name = os.path.basename(checkpoint_path)
|
||||
ckpt_cache_path = os.path.join(self.tmp_dir, checkpoint_name)
|
||||
|
||||
StableDiffusionPipeline.from_single_file(
|
||||
pretrained_model_link_or_path=checkpoint_path,
|
||||
torch_dtype=dtype,
|
||||
cache_dir=self.tmp_dir,
|
||||
).save_pretrained(ckpt_cache_path, safe_serialization=True)
|
||||
return ckpt_cache_path
|
||||
|
||||
def load_pipeline_from_path(
|
||||
self, model_path: str, dtype: torch.dtype
|
||||
) -> StableDiffusionPipeline:
|
||||
pipe: StableDiffusionPipeline = StableDiffusionPipeline.from_pretrained( # type: ignore[no-untyped-call]
|
||||
pretrained_model_name_or_path=model_path,
|
||||
torch_dtype=dtype,
|
||||
cache_dir=self.tmp_dir,
|
||||
)
|
||||
return pipe
|
||||
@@ -0,0 +1,10 @@
|
||||
from ..di import Container
|
||||
from .diffusers_pipeline_loader import DiffusersPipelineLoader
|
||||
|
||||
container = Container()
|
||||
container.wire(modules=[__name__])
|
||||
|
||||
NODE_CLASS_MAPPINGS = {DiffusersPipelineLoader.__name__: DiffusersPipelineLoader}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {DiffusersPipelineLoader.__name__: "Diffusers Pipeline Loader"}
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
@@ -0,0 +1,32 @@
|
||||
import folder_paths # pyright: ignore[reportMissingImports]
|
||||
from dependency_injector.wiring import Provide, inject
|
||||
|
||||
from ..di import Container
|
||||
from ..domain.model import PipelineModel
|
||||
from ..ui.pipeline_handler import PipelineHandler
|
||||
|
||||
|
||||
class DiffusersPipelineLoader:
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s) -> dict[str, dict[str, tuple[tuple[str, ...], ...]]]:
|
||||
return {
|
||||
"required": {
|
||||
"checkpoint_name": (folder_paths.get_filename_list("checkpoints"),),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PIPELINE",)
|
||||
FUNCTION = "execute"
|
||||
CATEGORY = "Diffusers"
|
||||
|
||||
@inject
|
||||
def execute(
|
||||
self,
|
||||
checkpoint_name: str,
|
||||
handler: PipelineHandler = Provide[Container.pipeline_handler_provider],
|
||||
) -> tuple[PipelineModel]:
|
||||
pipeline_model = handler.create(checkpoint_name)
|
||||
return (pipeline_model,)
|
||||
@@ -0,0 +1,3 @@
|
||||
from .create_pipeline_service import CreatePipelineService
|
||||
|
||||
__all__ = ["CreatePipelineService"]
|
||||
@@ -0,0 +1,23 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.model import PipelineModel
|
||||
from ..domain.repositories import IPipelineRepository
|
||||
|
||||
|
||||
class CreatePipelineService:
|
||||
def __init__(self, pipeline_repo: IPipelineRepository) -> None:
|
||||
self.pipeline_repo = pipeline_repo
|
||||
self.dtype = torch.float32
|
||||
|
||||
def execute(self, checkpoint_full_path: str) -> PipelineModel:
|
||||
path = self.pipeline_repo.convert_from_single_file(checkpoint_full_path, self.dtype)
|
||||
if not os.path.exists(path):
|
||||
raise RuntimeError(f"Failed to convert pipeline checkpoint: {checkpoint_full_path}")
|
||||
|
||||
pipe = self.pipeline_repo.load_pipeline_from_path(path, self.dtype)
|
||||
if pipe is None:
|
||||
raise RuntimeError(f"Failed to create pipeline from checkpoint: {checkpoint_full_path}")
|
||||
|
||||
return PipelineModel(pipeline=pipe, path=path)
|
||||
@@ -0,0 +1,3 @@
|
||||
from .pipeline_handler import PipelineHandler
|
||||
|
||||
__all__ = ["PipelineHandler"]
|
||||
@@ -0,0 +1,17 @@
|
||||
import os
|
||||
|
||||
import folder_paths # pyright: ignore[reportMissingImports]
|
||||
|
||||
from ..domain.model import PipelineModel
|
||||
from ..service.create_pipeline_service import CreatePipelineService
|
||||
|
||||
|
||||
class PipelineHandler:
|
||||
def __init__(self, create_pipeline_service: CreatePipelineService) -> None:
|
||||
self.create_pipeline_service = create_pipeline_service
|
||||
|
||||
def create(self, checkpoint_name: str) -> PipelineModel:
|
||||
checkpoint_path = folder_paths.get_full_path("checkpoints", checkpoint_name)
|
||||
if not os.path.exists(checkpoint_path):
|
||||
raise FileNotFoundError(f"Checkpoint file not found: {checkpoint_name}")
|
||||
return self.create_pipeline_service.execute(checkpoint_path)
|
||||
Reference in New Issue
Block a user