Files
Artificial-Sweetener-Simple…/simple_syrup/runtime/diffusion_model_loader.py
T

81 lines
2.7 KiB
Python

# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Load standalone diffusion models through ComfyUI's native policy."""
from __future__ import annotations
import importlib
from pathlib import Path
from types import ModuleType
from typing import Any
import torch
DIFFUSION_WEIGHT_DTYPES = (
"default",
"fp8_e4m3fn",
"fp8_e4m3fn_fast",
"fp8_e5m2",
)
class DiffusionModelLoader:
"""Adapt ComfyUI's standalone diffusion loader behind a typed boundary."""
def __init__(self, folder_paths_module: ModuleType | None = None) -> None:
"""Create a loader with injectable ComfyUI folder paths."""
self._folder_paths_module = folder_paths_module
def load(self, diffusion_model: str, weight_dtype: str) -> object:
"""Load one diffusion model using a validated weight dtype."""
return self.load_path(self.resolve_path(diffusion_model), weight_dtype)
def resolve_path(self, diffusion_model: str) -> Path:
"""Resolve a workflow diffusion model name through ComfyUI folders."""
model_path = self._folder_paths().get_full_path_or_raise(
"diffusion_models",
diffusion_model,
)
return Path(str(model_path))
def load_path(self, model_path: Path, weight_dtype: str) -> object:
"""Load an already resolved diffusion checkpoint path."""
model_options = diffusion_model_options(weight_dtype)
comfy_sd: Any = importlib.import_module("comfy.sd")
return comfy_sd.load_diffusion_model(
str(model_path),
model_options=model_options,
)
def _folder_paths(self) -> ModuleType:
"""Return the configured ComfyUI folder paths module."""
if self._folder_paths_module is not None:
return self._folder_paths_module
module: Any = importlib.import_module("folder_paths")
if not isinstance(module, ModuleType):
raise TypeError("folder_paths import did not return a module.")
self._folder_paths_module = module
return module
def diffusion_model_options(weight_dtype: str) -> dict[str, object]:
"""Return ComfyUI model options for a supported diffusion weight dtype."""
if weight_dtype not in DIFFUSION_WEIGHT_DTYPES:
valid = ", ".join(DIFFUSION_WEIGHT_DTYPES)
raise ValueError(f"diffusion_weight_dtype must be one of: {valid}.")
if weight_dtype == "fp8_e4m3fn":
return {"dtype": torch.float8_e4m3fn}
if weight_dtype == "fp8_e4m3fn_fast":
return {"dtype": torch.float8_e4m3fn, "fp8_optimizations": True}
if weight_dtype == "fp8_e5m2":
return {"dtype": torch.float8_e5m2}
return {}