feat: create architecture with diffusers pipeline

This commit is contained in:
Limitex
2025-10-31 20:53:33 +09:00
parent 40a179f5d2
commit 59c94c6439
18 changed files with 217 additions and 3 deletions
+3
View File
@@ -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
View File
@@ -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"
+3
View File
@@ -0,0 +1,3 @@
from .di import Container
__all__ = ["Container"]
+22
View File
@@ -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,
)
+3
View File
@@ -0,0 +1,3 @@
from .pipeline import PipelineModel
__all__ = ["PipelineModel"]
+9
View File
@@ -0,0 +1,9 @@
from dataclasses import dataclass
from diffusers import StableDiffusionPipeline
@dataclass
class PipelineModel:
pipeline: StableDiffusionPipeline
path: str
+3
View File
@@ -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
View File
+3
View File
@@ -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
+10
View File
@@ -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"]
+32
View File
@@ -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,)
+3
View File
@@ -0,0 +1,3 @@
from .create_pipeline_service import CreatePipelineService
__all__ = ["CreatePipelineService"]
+23
View File
@@ -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)
+3
View File
@@ -0,0 +1,3 @@
from .pipeline_handler import PipelineHandler
__all__ = ["PipelineHandler"]
+17
View File
@@ -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)