From f031c285899a8fded489367d897fc3f415a1f249 Mon Sep 17 00:00:00 2001 From: Artificial Sweetener Date: Mon, 10 Aug 2026 00:47:13 -0400 Subject: [PATCH] feat(anima): add cached quantization profiles --- __init__.py | 8 + simple_syrup/domain/anima_quantization.py | 118 +++++ simple_syrup/domain/model_quantization.py | 82 ++++ simple_syrup/domain/quant_cache.py | 310 +++++++++++++ simple_syrup/nodes/simple_load_anima.py | 34 +- simple_syrup/runtime/checkpoint_quantizer.py | 286 ++++++++++++ .../runtime/comfy_safetensors_dtypes.py | 132 ++++++ .../runtime/diffusion_model_loader.py | 15 +- simple_syrup/runtime/model_choices.py | 3 +- simple_syrup/runtime/quant_cache_leases.py | 94 ++++ simple_syrup/runtime/quant_cache_lock.py | 97 ++++ .../runtime/quant_cache_repository.py | 331 ++++++++++++++ simple_syrup/runtime/quant_cache_routes.py | 143 ++++++ simple_syrup/runtime/quant_cache_settings.py | 28 ++ .../runtime/quantization_capabilities.py | 128 ++++++ simple_syrup/runtime/quantization_progress.py | 79 ++++ simple_syrup/runtime/settings.py | 115 +---- simple_syrup/runtime/settings_repository.py | 103 +++++ simple_syrup/runtime/settings_routes.py | 25 +- .../services/anima_diffusion_model_service.py | 89 ++++ .../anima_loader_service.py} | 62 +-- .../services/external_llm_prompt_service.py | 2 +- simple_syrup/services/quant_cache_service.py | 127 ++++++ .../services/quantized_model_boundaries.py | 83 ++++ .../services/quantized_model_resolver.py | 252 ++++++++++ tests/test_anima_diffusion_model_service.py | 143 ++++++ tests/test_anima_loader.py | 27 +- tests/test_anima_quantization_workflow.py | 271 +++++++++++ tests/test_checkpoint_quantizer.py | 431 ++++++++++++++++++ tests/test_comfy_safetensors_dtypes.py | 60 +++ tests/test_model_quantization.py | 107 +++++ tests/test_persisted_widget_order_contract.py | 1 + tests/test_quant_cache.py | 247 ++++++++++ tests/test_quant_cache_routes.py | 160 +++++++ tests/test_quantization_capabilities.py | 129 ++++++ tests/test_quantized_model_resolver.py | 360 +++++++++++++++ tests/test_settings.py | 17 +- tests/test_settings_routes.py | 52 ++- tests/test_simple_load_anima_node.py | 34 ++ web/dist/simple-syrup.js | 405 +++++++++++----- web/src/api.ts | 97 +++- web/src/downloadableModelsSetting.ts | 50 ++ .../{settings.ts => externalLlmSettings.ts} | 151 +----- web/src/main.ts | 2 +- web/src/quantCacheSetting.ts | 141 ++++++ web/src/settingsRegistration.ts | 86 ++++ web/src/settingsUi.ts | 83 ++++ web/tests/api.test.ts | 67 ++- web/tests/settings.test.ts | 168 ++++++- 49 files changed, 5594 insertions(+), 441 deletions(-) create mode 100644 simple_syrup/domain/anima_quantization.py create mode 100644 simple_syrup/domain/model_quantization.py create mode 100644 simple_syrup/domain/quant_cache.py create mode 100644 simple_syrup/runtime/checkpoint_quantizer.py create mode 100644 simple_syrup/runtime/comfy_safetensors_dtypes.py create mode 100644 simple_syrup/runtime/quant_cache_leases.py create mode 100644 simple_syrup/runtime/quant_cache_lock.py create mode 100644 simple_syrup/runtime/quant_cache_repository.py create mode 100644 simple_syrup/runtime/quant_cache_routes.py create mode 100644 simple_syrup/runtime/quant_cache_settings.py create mode 100644 simple_syrup/runtime/quantization_capabilities.py create mode 100644 simple_syrup/runtime/quantization_progress.py create mode 100644 simple_syrup/runtime/settings_repository.py create mode 100644 simple_syrup/services/anima_diffusion_model_service.py rename simple_syrup/{runtime/anima_loader.py => services/anima_loader_service.py} (73%) create mode 100644 simple_syrup/services/quant_cache_service.py create mode 100644 simple_syrup/services/quantized_model_boundaries.py create mode 100644 simple_syrup/services/quantized_model_resolver.py create mode 100644 tests/test_anima_diffusion_model_service.py create mode 100644 tests/test_anima_quantization_workflow.py create mode 100644 tests/test_checkpoint_quantizer.py create mode 100644 tests/test_comfy_safetensors_dtypes.py create mode 100644 tests/test_model_quantization.py create mode 100644 tests/test_quant_cache.py create mode 100644 tests/test_quant_cache_routes.py create mode 100644 tests/test_quantization_capabilities.py create mode 100644 tests/test_quantized_model_resolver.py create mode 100644 web/src/downloadableModelsSetting.ts rename web/src/{settings.ts => externalLlmSettings.ts} (70%) create mode 100644 web/src/quantCacheSetting.ts create mode 100644 web/src/settingsRegistration.ts create mode 100644 web/src/settingsUi.ts diff --git a/__init__.py b/__init__.py index 84358d3..673cb00 100644 --- a/__init__.py +++ b/__init__.py @@ -12,12 +12,18 @@ from . import simple_syrup as _simple_syrup_package sys.modules.setdefault("simple_syrup", _simple_syrup_package) +from .simple_syrup.runtime.comfy_safetensors_dtypes import ( # noqa: E402 + register_comfy_safetensors_dtypes, +) from .simple_syrup.runtime.external_llm_routes import ( # noqa: E402 register_external_llm_routes, ) from .simple_syrup.runtime.mask_batch_preview_routes import ( # noqa: E402 register_mask_batch_preview_routes, ) +from .simple_syrup.runtime.quant_cache_routes import ( # noqa: E402 + register_quant_cache_routes, +) from .simple_syrup.runtime.settings_routes import register_settings_routes # noqa: E402 WEB_DIRECTORY = "./web/dist" @@ -42,6 +48,8 @@ async def comfy_entrypoint() -> object: register_settings_routes() +register_comfy_safetensors_dtypes() +register_quant_cache_routes() register_external_llm_routes() register_mask_batch_preview_routes() diff --git a/simple_syrup/domain/anima_quantization.py b/simple_syrup/domain/anima_quantization.py new file mode 100644 index 0000000..63524d9 --- /dev/null +++ b/simple_syrup/domain/anima_quantization.py @@ -0,0 +1,118 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Define versioned, quality-aware Anima quantization profiles.""" + +from __future__ import annotations + +import re +from dataclasses import dataclass + +from .model_quantization import ( + QuantizationFormat, + QuantizationProfile, + TensorDescriptor, +) + +ORIGINAL_PROFILE = QuantizationProfile("original", "Original", 2, frozenset()) +FP8_E4M3_PROFILE = QuantizationProfile( + "fp8-e4m3", + "FP8 E4M3", + 2, + frozenset({QuantizationFormat.FP8_E4M3}), +) +FP8_E5M2_PROFILE = QuantizationProfile( + "fp8-e5m2", + "FP8 E5M2", + 2, + frozenset({QuantizationFormat.FP8_E5M2}), +) +MXFP8_PROFILE = QuantizationProfile( + "mxfp8", + "MXFP8", + 2, + frozenset({QuantizationFormat.MXFP8}), +) +NVFP4_MIXED_PROFILE = QuantizationProfile( + "nvfp4-mixed", + "NVFP4 (Mixed)", + 3, + frozenset({QuantizationFormat.FP8_E4M3, QuantizationFormat.NVFP4}), +) +_PROFILES = ( + ORIGINAL_PROFILE, + FP8_E4M3_PROFILE, + FP8_E5M2_PROFILE, + MXFP8_PROFILE, + NVFP4_MIXED_PROFILE, +) +_MAIN_BLOCK_PATTERN = re.compile( + r"(?:^|\.)(?:net|diffusion_model)\.blocks\.(?P\d+)\." +) +_PROTECTED_BLOCKS = {0, 1, 27} +_FLOAT_DTYPES = {"F16", "BF16", "F32", "F64"} + + +@dataclass(frozen=True) +class AnimaQuantizationRecipe: + """Assign formats only within Anima's quality-safe DiT block envelope.""" + + model_family: str = "Anima" + version: int = 2 + + @property + def profiles(self) -> tuple[QuantizationProfile, ...]: + """Return Anima's stable workflow-facing profile order.""" + + return _PROFILES + + def profile_from_selection(self, selection: str) -> QuantizationProfile: + """Parse one current workflow selection into its profile.""" + + for profile in self.profiles: + if selection in (profile.label, profile.profile_id): + return profile + valid = ", ".join(profile.label for profile in self.profiles) + raise ValueError(f"quantization profile must be one of: {valid}.") + + def policy_for( + self, + tensor: TensorDescriptor, + profile: QuantizationProfile, + ) -> QuantizationFormat | None: + """Return Anima's per-tensor format while preserving sensitive layers.""" + + if profile not in self.profiles: + raise ValueError( + f"Unknown Anima quantization profile '{profile.profile_id}'." + ) + if profile.is_original or not _is_matrix_weight(tensor): + return None + if "llm_adapter" in tensor.name or "adaln_modulation" in tensor.name: + return None + block_match = _MAIN_BLOCK_PATTERN.search(tensor.name) + if block_match is None: + return None + if int(block_match.group("index")) in _PROTECTED_BLOCKS: + return None + if profile.profile_id == NVFP4_MIXED_PROFILE.profile_id: + if "v_proj" in tensor.name or ".mlp." in tensor.name: + return QuantizationFormat.FP8_E4M3 + if any( + projection in tensor.name + for projection in ("q_proj", "k_proj", "output_proj") + ): + return QuantizationFormat.NVFP4 + return None + return next(iter(profile.required_formats)) + + +def _is_matrix_weight(tensor: TensorDescriptor) -> bool: + """Return whether a tensor is an eligible floating-point matrix weight.""" + + return ( + tensor.dtype_name in _FLOAT_DTYPES + and len(tensor.shape) == 2 + and tensor.name.endswith(".weight") + ) diff --git a/simple_syrup/domain/model_quantization.py b/simple_syrup/domain/model_quantization.py new file mode 100644 index 0000000..730b658 --- /dev/null +++ b/simple_syrup/domain/model_quantization.py @@ -0,0 +1,82 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Define model-independent checkpoint quantization profile contracts.""" + +from __future__ import annotations + +from dataclasses import dataclass +from enum import StrEnum +from typing import Protocol + + +class QuantizationFormat(StrEnum): + """Identify one reusable ComfyUI tensor quantization format.""" + + FP8_E4M3 = "float8_e4m3fn" + FP8_E5M2 = "float8_e5m2" + NVFP4 = "nvfp4" + MXFP8 = "mxfp8" + + @property + def label(self) -> str: + """Return the concise format label used in diagnostics.""" + + return { + QuantizationFormat.FP8_E4M3: "FP8 E4M3", + QuantizationFormat.FP8_E5M2: "FP8 E5M2", + QuantizationFormat.NVFP4: "NVFP4", + QuantizationFormat.MXFP8: "MXFP8", + }[self] + + +@dataclass(frozen=True) +class QuantizationProfile: + """Describe one workflow-facing, versioned per-tensor policy profile.""" + + profile_id: str + label: str + version: int + required_formats: frozenset[QuantizationFormat] + + @property + def is_original(self) -> bool: + """Return whether this profile loads the source checkpoint unchanged.""" + + return not self.required_formats + + +@dataclass(frozen=True) +class TensorDescriptor: + """Describe a checkpoint tensor without coupling policy to PyTorch.""" + + name: str + shape: tuple[int, ...] + dtype_name: str + + +class ModelQuantizationRecipe(Protocol): + """Assign model-specific per-tensor formats for named profiles.""" + + @property + def model_family(self) -> str: + """Return the stable family identifier used in cache identity.""" + + @property + def version(self) -> int: + """Return the recipe version used in cache invalidation.""" + + @property + def profiles(self) -> tuple[QuantizationProfile, ...]: + """Return deterministic workflow profiles owned by this recipe.""" + + def profile_from_selection(self, selection: str) -> QuantizationProfile: + """Parse a workflow selection into a recipe-owned profile.""" + + def policy_for( + self, + tensor: TensorDescriptor, + profile: QuantizationProfile, + ) -> QuantizationFormat | None: + """Return the tensor format or ``None`` to preserve source precision.""" diff --git a/simple_syrup/domain/quant_cache.py b/simple_syrup/domain/quant_cache.py new file mode 100644 index 0000000..32d8307 --- /dev/null +++ b/simple_syrup/domain/quant_cache.py @@ -0,0 +1,310 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Define validated identities and records for quantized profile artifacts.""" + +from __future__ import annotations + +import hashlib +from dataclasses import dataclass +from datetime import UTC, datetime +from pathlib import Path +from typing import cast + +from .model_quantization import QuantizationProfile + +MANIFEST_SCHEMA_VERSION = 2 + + +@dataclass(frozen=True) +class SourceCheckpointIdentity: + """Identify an authoritative source checkpoint and its current file state.""" + + display_name: str + path: Path + size_bytes: int + modified_ns: int + sha256: str + + +@dataclass(frozen=True) +class QuantCacheIdentity: + """Identify one profile and recipe derivative of a source checkpoint.""" + + source: SourceCheckpointIdentity + profile: QuantizationProfile + model_family: str + recipe_version: int + + @property + def stable_key(self) -> str: + """Return the collision-resistant cache key.""" + + identity = "\0".join( + ( + self.source.sha256, + self.profile.profile_id, + str(self.profile.version), + self.model_family, + str(self.recipe_version), + ) + ) + return hashlib.sha256(identity.encode("utf-8")).hexdigest() + + +@dataclass(frozen=True) +class QuantCacheManifest: + """Describe one complete SimpleSyrup-managed cache artifact.""" + + source_model: str + source_path: str + source_sha256: str + source_size_bytes: int + source_modified_ns: int + profile_id: str + profile_label: str + profile_version: int + quantization_formats: tuple[str, ...] + model_family: str + recipe_version: int + artifact_file: str + artifact_size_bytes: int + created_at: str + last_used_at: str + schema_version: int = MANIFEST_SCHEMA_VERSION + managed_by: str = "SimpleSyrup" + + @classmethod + def create( + cls, + identity: QuantCacheIdentity, + artifact_file: str, + artifact_size_bytes: int, + ) -> QuantCacheManifest: + """Create a current manifest for a completed artifact.""" + + now = datetime.now(UTC).isoformat() + return cls( + source_model=identity.source.display_name, + source_path=str(identity.source.path), + source_sha256=identity.source.sha256, + source_size_bytes=identity.source.size_bytes, + source_modified_ns=identity.source.modified_ns, + profile_id=identity.profile.profile_id, + profile_label=identity.profile.label, + profile_version=identity.profile.version, + quantization_formats=tuple( + sorted(item.value for item in identity.profile.required_formats) + ), + model_family=identity.model_family, + recipe_version=identity.recipe_version, + artifact_file=artifact_file, + artifact_size_bytes=artifact_size_bytes, + created_at=now, + last_used_at=now, + ) + + def matches_current_source( + self, + source_model: str, + source_path: Path, + source_size_bytes: int, + source_modified_ns: int, + profile: QuantizationProfile, + model_family: str, + recipe_version: int, + ) -> bool: + """Return whether this artifact derives from the unchanged source file.""" + + return ( + self.source_model == source_model + and Path(self.source_path) == source_path + and self.source_size_bytes == source_size_bytes + and self.source_modified_ns == source_modified_ns + and self.profile_id == profile.profile_id + and self.profile_version == profile.version + and self.model_family == model_family + and self.recipe_version == recipe_version + ) + + def matches_identity(self, identity: QuantCacheIdentity) -> bool: + """Return whether this v2 manifest exactly describes an identity.""" + + return ( + self.schema_version == MANIFEST_SCHEMA_VERSION + and self.source_sha256 == identity.source.sha256 + and self.profile_id == identity.profile.profile_id + and self.profile_version == identity.profile.version + and self.model_family == identity.model_family + and self.recipe_version == identity.recipe_version + ) + + def touched(self) -> QuantCacheManifest: + """Return a copy with a current explicit LRU timestamp.""" + + payload = self.to_payload() + payload["last_used_at"] = datetime.now(UTC).isoformat() + return QuantCacheManifest.from_payload(payload) + + def to_payload(self) -> dict[str, object]: + """Return the human-readable JSON representation.""" + + if self.schema_version == 1: + return { + "schema_version": 1, + "managed_by": self.managed_by, + "source_model": self.source_model, + "source_path": self.source_path, + "source_sha256": self.source_sha256, + "source_size_bytes": self.source_size_bytes, + "source_modified_ns": self.source_modified_ns, + "quantization_format": self.quantization_formats[0], + "model_family": self.model_family, + "recipe_version": self.recipe_version, + "artifact_file": self.artifact_file, + "artifact_size_bytes": self.artifact_size_bytes, + "created_at": self.created_at, + "last_used_at": self.last_used_at, + } + return { + "schema_version": self.schema_version, + "managed_by": self.managed_by, + "source_model": self.source_model, + "source_path": self.source_path, + "source_sha256": self.source_sha256, + "source_size_bytes": self.source_size_bytes, + "source_modified_ns": self.source_modified_ns, + "profile_id": self.profile_id, + "profile_label": self.profile_label, + "profile_version": self.profile_version, + "quantization_formats": list(self.quantization_formats), + "model_family": self.model_family, + "recipe_version": self.recipe_version, + "artifact_file": self.artifact_file, + "artifact_size_bytes": self.artifact_size_bytes, + "created_at": self.created_at, + "last_used_at": self.last_used_at, + } + + @classmethod + def from_payload(cls, payload: object) -> QuantCacheManifest: + """Validate managed manifests, retaining v1 only for cache cleanup.""" + + if not isinstance(payload, dict): + raise ValueError("Quant cache manifest must be a JSON object.") + if payload.get("schema_version") == 1: + return cls._from_legacy_payload(payload) + required_strings = ( + "managed_by", + "source_model", + "source_path", + "source_sha256", + "profile_id", + "profile_label", + "model_family", + "artifact_file", + "created_at", + "last_used_at", + ) + for key in required_strings: + if not isinstance(payload.get(key), str): + raise ValueError( + f"Quant cache manifest field '{key}' must be a string." + ) + required_integers = ( + "schema_version", + "source_size_bytes", + "source_modified_ns", + "profile_version", + "recipe_version", + "artifact_size_bytes", + ) + for key in required_integers: + value = payload.get(key) + if not isinstance(value, int) or isinstance(value, bool): + raise ValueError( + f"Quant cache manifest field '{key}' must be an integer." + ) + raw_formats = payload.get("quantization_formats") + if not isinstance(raw_formats, list) or not all( + isinstance(item, str) for item in raw_formats + ): + raise ValueError( + "Quant cache manifest field 'quantization_formats' must be a " + "string list." + ) + if payload["managed_by"] != "SimpleSyrup": + raise ValueError("Quant cache manifest is not managed by SimpleSyrup.") + if payload["schema_version"] != MANIFEST_SCHEMA_VERSION: + raise ValueError("Quant cache manifest schema version is unsupported.") + return cls( + schema_version=payload["schema_version"], + managed_by=payload["managed_by"], + source_model=payload["source_model"], + source_path=payload["source_path"], + source_sha256=payload["source_sha256"], + source_size_bytes=payload["source_size_bytes"], + source_modified_ns=payload["source_modified_ns"], + profile_id=payload["profile_id"], + profile_label=payload["profile_label"], + profile_version=payload["profile_version"], + quantization_formats=tuple(raw_formats), + model_family=payload["model_family"], + recipe_version=payload["recipe_version"], + artifact_file=payload["artifact_file"], + artifact_size_bytes=payload["artifact_size_bytes"], + created_at=payload["created_at"], + last_used_at=payload["last_used_at"], + ) + + @classmethod + def _from_legacy_payload(cls, payload: dict[object, object]) -> QuantCacheManifest: + """Decode v1 solely so ordinary LRU and clearing can remove it.""" + + required_strings = ( + "managed_by", + "source_model", + "source_path", + "source_sha256", + "quantization_format", + "model_family", + "artifact_file", + "created_at", + "last_used_at", + ) + required_integers = ( + "source_size_bytes", + "source_modified_ns", + "recipe_version", + "artifact_size_bytes", + ) + if any(not isinstance(payload.get(key), str) for key in required_strings): + raise ValueError("Legacy quant cache manifest has invalid string fields.") + if any( + not isinstance(payload.get(key), int) or isinstance(payload.get(key), bool) + for key in required_integers + ): + raise ValueError("Legacy quant cache manifest has invalid integer fields.") + if payload["managed_by"] != "SimpleSyrup": + raise ValueError("Quant cache manifest is not managed by SimpleSyrup.") + quantization_format = str(payload["quantization_format"]) + return cls( + schema_version=1, + managed_by=str(payload["managed_by"]), + source_model=str(payload["source_model"]), + source_path=str(payload["source_path"]), + source_sha256=str(payload["source_sha256"]), + source_size_bytes=cast(int, payload["source_size_bytes"]), + source_modified_ns=cast(int, payload["source_modified_ns"]), + profile_id=f"legacy-v1-{quantization_format}", + profile_label=f"Legacy v1 {quantization_format}", + profile_version=1, + quantization_formats=(quantization_format,), + model_family=str(payload["model_family"]), + recipe_version=cast(int, payload["recipe_version"]), + artifact_file=str(payload["artifact_file"]), + artifact_size_bytes=cast(int, payload["artifact_size_bytes"]), + created_at=str(payload["created_at"]), + last_used_at=str(payload["last_used_at"]), + ) diff --git a/simple_syrup/nodes/simple_load_anima.py b/simple_syrup/nodes/simple_load_anima.py index 250b702..46aa698 100644 --- a/simple_syrup/nodes/simple_load_anima.py +++ b/simple_syrup/nodes/simple_load_anima.py @@ -10,14 +10,17 @@ import importlib from types import ModuleType from typing import Any -from ..runtime.anima_loader import ( +from ..domain.anima_quantization import AnimaQuantizationRecipe +from ..runtime.diffusion_model_loader import DIFFUSION_WEIGHT_DTYPES +from ..runtime.model_downloads import ComfyProgressReporter +from ..runtime.quantization_capabilities import QuantizationCapabilityCatalog +from ..runtime.quantization_progress import ComfyQuantizationProgressReporter +from ..runtime.vae_loader import vae_choices +from ..services.anima_loader_service import ( AUTO_CHOICE, CLIP_DEVICES, - DIFFUSION_WEIGHT_DTYPES, AnimaLoaderService, ) -from ..runtime.model_downloads import ComfyProgressReporter -from ..runtime.vae_loader import vae_choices from . import tooltips @@ -25,6 +28,8 @@ class SimpleLoadAnima: """Expose Anima diffusion, text encoder, and VAE loading as one node.""" _service = AnimaLoaderService() + _quantization_recipe = AnimaQuantizationRecipe() + _quantization_capabilities = QuantizationCapabilityCatalog() RETURN_TYPES = ("MODEL", "CLIP", "VAE") RETURN_NAMES = ("model", "clip", "vae") @@ -53,14 +58,28 @@ class SimpleLoadAnima: ) }, ), + "quantization": ( + cls._quantization_capabilities.selection_labels( + cls._quantization_recipe.profiles + ), + { + "default": "Original", + "advanced": True, + "tooltip": ( + "Creates or reuses a GPU-supported quantized copy in the " + "global models/SyrupQuants cache; Original loads the " + "selected model unchanged." + ), + }, + ), "diffusion_weight_dtype": ( list(DIFFUSION_WEIGHT_DTYPES), { "default": "default", "advanced": True, "tooltip": ( - "Weight precision for Anima. Lower precision can reduce " - "memory use but may slightly change results." + "Load-time weight precision used with Original; cached " + "quantized copies use their stored quantization format." ), }, ), @@ -103,6 +122,7 @@ class SimpleLoadAnima: def load_models( self, diffusion_model: str, + quantization: str, diffusion_weight_dtype: str, text_encoder: str, text_encoder_device: str, @@ -112,11 +132,13 @@ class SimpleLoadAnima: return self._service.load_models( diffusion_model=diffusion_model, + quantization=quantization, diffusion_weight_dtype=diffusion_weight_dtype, text_encoder=text_encoder, text_encoder_device=text_encoder_device, vae=vae, progress=ComfyProgressReporter(), + quantization_progress=ComfyQuantizationProgressReporter(), ) diff --git a/simple_syrup/runtime/checkpoint_quantizer.py b/simple_syrup/runtime/checkpoint_quantizer.py new file mode 100644 index 0000000..12a8fb5 --- /dev/null +++ b/simple_syrup/runtime/checkpoint_quantizer.py @@ -0,0 +1,286 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Stream source safetensors through ComfyUI-native quantization layouts.""" + +from __future__ import annotations + +import importlib +import json +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import torch +from safetensors import safe_open +from safetensors.torch import save_file + +from ..domain.model_quantization import ( + ModelQuantizationRecipe, + QuantizationFormat, + QuantizationProfile, + TensorDescriptor, +) +from ..domain.quant_cache import SourceCheckpointIdentity +from .comfy_safetensors_dtypes import ComfySafetensorsDtypeRegistry +from .quantization_progress import QuantizationProgressReporter + +_DTYPE_BYTES = { + "BOOL": 1, + "U8": 1, + "I8": 1, + "F8_E4M3": 1, + "F8_E5M2": 1, + "F8_E8M0": 1, + "U16": 2, + "I16": 2, + "F16": 2, + "BF16": 2, + "U32": 4, + "I32": 4, + "F32": 4, + "U64": 8, + "I64": 8, + "F64": 8, +} + + +@dataclass(frozen=True) +class CheckpointQuantizationResult: + """Summarize a completed checkpoint conversion.""" + + quantized_tensor_count: int + preserved_tensor_count: int + output_size_bytes: int + + +class SafetensorsCheckpointQuantizer: + """Quantize eligible tensors without constructing the source model graph.""" + + def __init__( + self, + dtype_registry: ComfySafetensorsDtypeRegistry | None = None, + ) -> None: + """Create a quantizer with an injectable checkpoint compatibility owner.""" + + self._dtype_registry = dtype_registry or ComfySafetensorsDtypeRegistry() + + def quantize( + self, + *, + source: SourceCheckpointIdentity, + destination_path: Path, + profile: QuantizationProfile, + recipe: ModelQuantizationRecipe, + progress: QuantizationProgressReporter, + progress_base: int, + progress_total: int, + ) -> CheckpointQuantizationResult: + """Convert one safetensors checkpoint tensor by tensor.""" + + if source.path.suffix.lower() not in (".safetensors", ".sft"): + raise ValueError( + "On-demand quantization requires a safetensors diffusion model. " + f"Selected source: '{source.path.name}'." + ) + if profile.is_original: + raise ValueError("Original profiles do not require quantization.") + + quant_ops: Any = importlib.import_module("comfy.quant_ops") + model_management: Any = importlib.import_module("comfy.model_management") + layouts = _resolve_layouts(quant_ops, profile) + + destination_path.parent.mkdir(parents=True, exist_ok=True) + device = model_management.get_torch_device() + output: dict[str, torch.Tensor] = {} + quantized_count = 0 + preserved_count = 0 + processed_bytes = 0 + + with safe_open( + str(source.path), framework="pt", device="cpu" + ) as source_checkpoint: + keys = list(source_checkpoint.keys()) + if any(key.endswith(".comfy_quant") for key in keys): + raise ValueError( + "The selected source already contains ComfyUI quantization " + "metadata. Choose the original full-precision Anima checkpoint " + "as the source." + ) + metadata = dict(source_checkpoint.metadata() or {}) + with torch.inference_mode(): + for key in keys: + tensor_slice = source_checkpoint.get_slice(key) + descriptor = TensorDescriptor( + name=key, + shape=tuple(tensor_slice.get_shape()), + dtype_name=tensor_slice.get_dtype(), + ) + tensor_bytes = _storage_bytes(descriptor) + tensor = source_checkpoint.get_tensor(key) + quantization_format = recipe.policy_for(descriptor, profile) + if quantization_format is not None: + device_tensor = tensor.to(device=device).contiguous() + quantized = quant_ops.QuantizedTensor.from_float( + device_tensor, + layouts[quantization_format], + scale="recalculate", + ) + state_tensors = quantized.state_dict(key) + for output_key, output_tensor in state_tensors.items(): + output[output_key] = ( + output_tensor.detach().to(device="cpu").contiguous() + ) + layer_name = key.removesuffix(".weight") + output[f"{layer_name}.comfy_quant"] = torch.tensor( + list( + json.dumps( + {"format": quantization_format.value}, + separators=(",", ":"), + ).encode("utf-8") + ), + dtype=torch.uint8, + ) + quantized_count += 1 + del quantized, device_tensor + else: + output[key] = tensor + preserved_count += 1 + processed_bytes += tensor_bytes + progress.advance( + min(progress_base + processed_bytes, progress_total), + progress_total, + ) + + metadata.update( + { + "simple_syrup.derived_model": "true", + "simple_syrup.model_family": recipe.model_family, + "simple_syrup.quantization_profile": profile.profile_id, + "simple_syrup.profile_label": profile.label, + "simple_syrup.profile_version": str(profile.version), + "simple_syrup.quantization_formats": ",".join( + sorted(item.value for item in profile.required_formats) + ), + "simple_syrup.recipe_version": str(recipe.version), + "simple_syrup.source_model": source.display_name, + "simple_syrup.source_sha256": source.sha256, + } + ) + save_file(output, str(destination_path), metadata=metadata) + + self._validate_output(destination_path, profile, quantized_count) + return CheckpointQuantizationResult( + quantized_tensor_count=quantized_count, + preserved_tensor_count=preserved_count, + output_size_bytes=destination_path.stat().st_size, + ) + + def _validate_output( + self, + path: Path, + profile: QuantizationProfile, + expected_quantized_tensors: int, + ) -> None: + """Validate generated dtypes and Comfy markers before cache publication.""" + + if expected_quantized_tensors <= 0: + raise ValueError( + "The selected checkpoint contained no tensors eligible for the model's " + "quantization recipe." + ) + self._dtype_registry.validate_checkpoint_header(path) + with safe_open(str(path), framework="pt", device="cpu") as checkpoint: + marker_keys = tuple( + key for key in checkpoint.keys() if key.endswith(".comfy_quant") + ) + if len(marker_keys) != expected_quantized_tensors: + raise ValueError( + "Quantized checkpoint validation found incomplete layer metadata." + ) + expected_formats = ",".join( + sorted(item.value for item in profile.required_formats) + ) + metadata = checkpoint.metadata() or {} + if ( + metadata.get("simple_syrup.quantization_profile") != profile.profile_id + or metadata.get("simple_syrup.profile_version") != str(profile.version) + or metadata.get("simple_syrup.quantization_formats") != expected_formats + ): + raise ValueError( + "Quantized checkpoint validation found inconsistent " + "profile metadata." + ) + for marker_key in marker_keys: + quantization_format = _marker_format( + checkpoint.get_tensor(marker_key), marker_key + ) + if quantization_format not in profile.required_formats: + raise ValueError( + "Quantized checkpoint contains a layer format outside the " + f"selected profile: {quantization_format.label}." + ) + if not self._dtype_registry.supports_format(quantization_format): + raise ValueError( + "Generated checkpoint quantization metadata requires a " + f"format unsupported by the active Comfy loader: " + f"{quantization_format.label}." + ) + + +def _resolve_layouts( + quant_ops: Any, + profile: QuantizationProfile, +) -> dict[QuantizationFormat, str]: + """Resolve every layout required by one profile before conversion starts.""" + + layouts: dict[QuantizationFormat, str] = {} + for quantization_format in profile.required_formats: + algorithm = quant_ops.QUANT_ALGOS.get(quantization_format.value) + if not isinstance(algorithm, dict): + raise ValueError( + f"ComfyUI did not register {quantization_format.label} quantization." + ) + layout_name = algorithm.get("comfy_tensor_layout") + if not isinstance(layout_name, str): + raise ValueError( + f"ComfyUI's {quantization_format.label} layout is invalid." + ) + layouts[quantization_format] = layout_name + return layouts + + +def _storage_bytes(tensor: TensorDescriptor) -> int: + """Return safetensors storage bytes represented by a descriptor.""" + + element_bytes = _DTYPE_BYTES.get(tensor.dtype_name) + if element_bytes is None: + raise ValueError( + f"Unsupported safetensors dtype '{tensor.dtype_name}' for '{tensor.name}'." + ) + elements = 1 + for dimension in tensor.shape: + elements *= dimension + return elements * element_bytes + + +def _marker_format(marker: torch.Tensor, marker_key: str) -> QuantizationFormat: + """Parse one Comfy quant marker into a supported domain format.""" + + try: + payload: object = json.loads(bytes(marker.tolist()).decode("utf-8")) + except (UnicodeDecodeError, json.JSONDecodeError, TypeError, ValueError) as error: + raise ValueError( + f"Quantized checkpoint marker '{marker_key}' is malformed." + ) from error + if not isinstance(payload, dict) or not isinstance(payload.get("format"), str): + raise ValueError( + f"Quantized checkpoint marker '{marker_key}' has no valid format." + ) + try: + return QuantizationFormat(payload["format"]) + except ValueError as error: + raise ValueError( + f"Quantized checkpoint marker '{marker_key}' names an unknown format." + ) from error diff --git a/simple_syrup/runtime/comfy_safetensors_dtypes.py b/simple_syrup/runtime/comfy_safetensors_dtypes.py new file mode 100644 index 0000000..cbd215e --- /dev/null +++ b/simple_syrup/runtime/comfy_safetensors_dtypes.py @@ -0,0 +1,132 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own SimpleSyrup's guarded Comfy safetensors dtype compatibility.""" + +from __future__ import annotations + +import importlib +import json +import struct +from pathlib import Path +from types import ModuleType +from typing import Any + +import torch + +from ..domain.model_quantization import QuantizationFormat +from ..shared.logging import get_logger + +LOGGER = get_logger(__name__) +_FORMAT_DTYPES = { + QuantizationFormat.FP8_E4M3: frozenset({"F32", "F8_E4M3"}), + QuantizationFormat.FP8_E5M2: frozenset({"F32", "F8_E5M2"}), + QuantizationFormat.NVFP4: frozenset({"F32", "F8_E4M3", "U8"}), + QuantizationFormat.MXFP8: frozenset({"F8_E4M3", "F8_E8M0"}), +} + + +class ComfySafetensorsDtypeRegistry: + """Register and validate dtypes used by Comfy's AIMDO mmap loader.""" + + def __init__( + self, + comfy_utils_module: ModuleType | None = None, + torch_module: ModuleType | None = None, + ) -> None: + """Create a registry with injectable host boundaries.""" + + self._comfy_utils_module = comfy_utils_module + self._torch_module = torch_module or torch + + def register_extension_dtypes(self) -> bool: + """Register E8M0 when safe, returning whether MXFP8 is loadable.""" + + dtype = getattr(self._torch_module, "float8_e8m0fnu", None) + if dtype is None: + return False + types = self._types() + existing = types.get("F8_E8M0") + if existing is None: + types["F8_E8M0"] = dtype + return True + if existing == dtype: + return True + LOGGER.warning( + "Comfy safetensors E8M0 dtype mapping conflicts with PyTorch", + extra={"registered_dtype": str(existing), "expected_dtype": str(dtype)}, + ) + return False + + def supports_format(self, quantization_format: QuantizationFormat) -> bool: + """Return whether the active Comfy loader maps every format dtype.""" + + if quantization_format is QuantizationFormat.MXFP8: + if not self.register_extension_dtypes(): + return False + return _FORMAT_DTYPES[quantization_format].issubset(self._types()) + + def validate_checkpoint_header(self, checkpoint_path: Path) -> None: + """Reject generated checkpoints containing host-unreadable dtypes.""" + + header = _read_safetensors_header(checkpoint_path) + used_dtypes = { + value.get("dtype") + for key, value in header.items() + if key != "__metadata__" and isinstance(value, dict) + } + unknown = sorted( + dtype + for dtype in used_dtypes + if isinstance(dtype, str) and dtype not in self._types() + ) + if unknown: + joined = ", ".join(unknown) + raise ValueError( + "Generated checkpoint uses safetensors dtypes unsupported by the " + f"active Comfy loader: {joined}." + ) + + def _types(self) -> dict[str, object]: + """Return Comfy's loader dtype map through one isolated private boundary.""" + + comfy_utils = self._comfy_utils_module or importlib.import_module("comfy.utils") + types: Any = getattr(comfy_utils, "_TYPES", None) + if not isinstance(types, dict): + raise RuntimeError("Comfy's safetensors dtype registry is unavailable.") + return types + + +def register_comfy_safetensors_dtypes() -> None: + """Register SimpleSyrup's lightweight host dtype compatibility at import.""" + + try: + ComfySafetensorsDtypeRegistry().register_extension_dtypes() + except (ImportError, RuntimeError) as error: + LOGGER.warning( + "Comfy safetensors dtype compatibility registration failed", + extra={"reason": str(error)}, + ) + + +def _read_safetensors_header(path: Path) -> dict[str, object]: + """Read only a safetensors JSON header without materializing tensor data.""" + + with path.open("rb") as checkpoint: + header_length_bytes = checkpoint.read(8) + if len(header_length_bytes) != 8: + raise ValueError( + "Generated checkpoint has an incomplete safetensors header." + ) + header_length = struct.unpack(" object: """Load one diffusion model using a validated weight dtype.""" - model_options = diffusion_model_options(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( - model_path, + str(model_path), model_options=model_options, ) diff --git a/simple_syrup/runtime/model_choices.py b/simple_syrup/runtime/model_choices.py index b3fd232..0a451a4 100644 --- a/simple_syrup/runtime/model_choices.py +++ b/simple_syrup/runtime/model_choices.py @@ -21,7 +21,8 @@ from .model_catalog import ( wd14_tagger_choices, ) from .model_folders import resolve_model_file -from .settings import SimpleSyrupSettings, SimpleSyrupSettingsRepository +from .settings import SimpleSyrupSettings +from .settings_repository import SimpleSyrupSettingsRepository from .vitmatte_loader import ViTMatteLoaderService NO_LOCAL_SAM_MODELS = "No local SAM models found" diff --git a/simple_syrup/runtime/quant_cache_leases.py b/simple_syrup/runtime/quant_cache_leases.py new file mode 100644 index 0000000..0394fcd --- /dev/null +++ b/simple_syrup/runtime/quant_cache_leases.py @@ -0,0 +1,94 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Protect cache artifacts while loaded ComfyUI model objects reference them.""" + +from __future__ import annotations + +import threading +import weakref +from collections.abc import Callable +from dataclasses import dataclass, field +from pathlib import Path + + +@dataclass +class QuantCacheReservation: + """Hold a temporary artifact lease across resolution and model loading.""" + + artifact_path: Path + _release_callback: Callable[[Path], None] + _released: bool = field(default=False, init=False) + + def release(self) -> None: + """Release the reservation exactly once.""" + + if self._released: + return + self._released = True + self._release_callback(self.artifact_path) + + +class QuantCacheLeaseRegistry: + """Track artifact leases against the lifetime of underlying model objects.""" + + def __init__(self) -> None: + """Create an empty thread-safe lease registry.""" + + self._counts: dict[Path, int] = {} + self._owners: weakref.WeakKeyDictionary[object, set[Path]] = ( + weakref.WeakKeyDictionary() + ) + self._lock = threading.RLock() + + def lease(self, artifact_path: Path, loaded_model: object) -> None: + """Protect an artifact until the underlying model is garbage collected.""" + + owner = getattr(loaded_model, "model", loaded_model) + normalized_path = artifact_path.resolve() + with self._lock: + existing = self._owners.setdefault(owner, set()) + if normalized_path in existing: + return + existing.add(normalized_path) + self._acquire(normalized_path) + weakref.finalize(owner, self._release, normalized_path) + + def reserve(self, artifact_path: Path) -> QuantCacheReservation: + """Protect an artifact during work that precedes a model-owned lease.""" + + normalized_path = artifact_path.resolve() + with self._lock: + self._acquire(normalized_path) + return QuantCacheReservation(normalized_path, self._release) + + def is_leased(self, artifact_path: Path) -> bool: + """Return whether a live model currently references the artifact.""" + + with self._lock: + return self._counts.get(artifact_path.resolve(), 0) > 0 + + def active_count(self) -> int: + """Return the number of distinct protected artifact paths.""" + + with self._lock: + return sum(count > 0 for count in self._counts.values()) + + def _release(self, artifact_path: Path) -> None: + """Release one model-owned artifact lease.""" + + with self._lock: + remaining = self._counts.get(artifact_path, 0) - 1 + if remaining > 0: + self._counts[artifact_path] = remaining + else: + self._counts.pop(artifact_path, None) + + def _acquire(self, artifact_path: Path) -> None: + """Increment one normalized path while the registry lock is held.""" + + self._counts[artifact_path] = self._counts.get(artifact_path, 0) + 1 + + +GLOBAL_QUANT_CACHE_LEASES = QuantCacheLeaseRegistry() diff --git a/simple_syrup/runtime/quant_cache_lock.py b/simple_syrup/runtime/quant_cache_lock.py new file mode 100644 index 0000000..285f524 --- /dev/null +++ b/simple_syrup/runtime/quant_cache_lock.py @@ -0,0 +1,97 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Coordinate quant cache generation across threads and ComfyUI processes.""" + +from __future__ import annotations + +import os +import threading +import time +from collections.abc import Iterator +from contextlib import contextmanager +from pathlib import Path + +LOCK_POLL_SECONDS = 0.2 +LOCK_TIMEOUT_SECONDS = 60 * 60 +STALE_LOCK_SECONDS = 2 * 60 * 60 + + +class QuantCacheBuildCoordinator: + """Serialize builders for the same derived checkpoint identity.""" + + _thread_locks: dict[str, threading.Lock] = {} + _thread_locks_guard = threading.Lock() + + def __init__(self, lock_directory: Path) -> None: + """Create a coordinator rooted in the global cache directory.""" + + self._lock_directory = lock_directory + + @contextmanager + def acquire(self, stable_key: str) -> Iterator[None]: + """Acquire in-process and cross-process ownership for one cache key.""" + + thread_lock = self._thread_lock(stable_key) + with thread_lock: + self._lock_directory.mkdir(parents=True, exist_ok=True) + lock_path = self._lock_directory / f"{stable_key}.lock" + self._acquire_file_lock(lock_path) + try: + yield + finally: + try: + lock_path.unlink(missing_ok=True) + except OSError: + pass + + @classmethod + def _thread_lock(cls, stable_key: str) -> threading.Lock: + """Return the shared in-process lock for a stable cache key.""" + + with cls._thread_locks_guard: + return cls._thread_locks.setdefault(stable_key, threading.Lock()) + + @staticmethod + def _acquire_file_lock(lock_path: Path) -> None: + """Create an exclusive lock file, recovering only demonstrably stale locks.""" + + started = time.monotonic() + while True: + try: + descriptor = os.open( + lock_path, + os.O_CREAT | os.O_EXCL | os.O_WRONLY, + ) + with os.fdopen(descriptor, "w", encoding="utf-8") as lock_file: + lock_file.write(f"pid={os.getpid()}\ncreated={time.time()}\n") + return + except FileExistsError: + try: + age = time.time() - lock_path.stat().st_mtime + if age > STALE_LOCK_SECONDS and not _lock_owner_is_alive(lock_path): + lock_path.unlink() + continue + except FileNotFoundError: + continue + if time.monotonic() - started >= LOCK_TIMEOUT_SECONDS: + raise TimeoutError( + "Timed out waiting for another SimpleSyrup process to finish " + "the same quantized checkpoint." + ) from None + time.sleep(LOCK_POLL_SECONDS) + + +def _lock_owner_is_alive(lock_path: Path) -> bool: + """Return whether a well-formed lock still belongs to a running process.""" + + try: + first_line = lock_path.read_text(encoding="utf-8").splitlines()[0] + process_id = int(first_line.removeprefix("pid=")) + os.kill(process_id, 0) + except PermissionError: + return True + except (IndexError, OSError, ValueError): + return False + return True diff --git a/simple_syrup/runtime/quant_cache_repository.py b/simple_syrup/runtime/quant_cache_repository.py new file mode 100644 index 0000000..a73f190 --- /dev/null +++ b/simple_syrup/runtime/quant_cache_repository.py @@ -0,0 +1,331 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Persist readable, globally shared quantized checkpoint cache artifacts.""" + +from __future__ import annotations + +import importlib +import json +import re +import shutil +import uuid +from dataclasses import dataclass +from json import JSONDecodeError +from pathlib import Path +from types import ModuleType +from typing import Any + +from ..domain.model_quantization import ModelQuantizationRecipe, QuantizationProfile +from ..domain.quant_cache import QuantCacheIdentity, QuantCacheManifest +from ..shared.logging import get_logger + +LOGGER = get_logger(__name__) +CACHE_DIRECTORY_NAME = "SyrupQuants" +ARTIFACT_FILENAME_SUFFIX = ".safetensors" +MANIFEST_FILENAME = "manifest.json" +README_FILENAME = "README.txt" +README_CONTENT = """SimpleSyrup Quantized Model Cache +=================================== + +This folder contains quantized copies generated from models selected in +SimpleSyrup loader nodes. Your original models remain in their normal folders +and are the authoritative models recorded in workflows. + +SimpleSyrup manages this folder as one global least-recently-used cache. You can +change its size limit or clear inactive cached models in the SimpleSyrup section +of ComfyUI settings. + +It is safe to delete this entire folder while ComfyUI is stopped. Missing +quantized copies will be generated again when requested. +""" +_SAFE_COMPONENT_PATTERN = re.compile(r"[^A-Za-z0-9._-]+") + + +@dataclass(frozen=True) +class QuantCacheArtifact: + """Pair a complete cached model path with its validated manifest.""" + + path: Path + manifest_path: Path + manifest: QuantCacheManifest + + +class QuantCacheRepository: + """Own filesystem layout and persistence for SimpleSyrup quant artifacts.""" + + def __init__( + self, + cache_root: Path | None = None, + folder_paths_module: ModuleType | None = None, + ) -> None: + """Create a repository with injectable model-directory discovery.""" + + self._cache_root = cache_root + self._folder_paths_module = folder_paths_module + + @property + def root(self) -> Path: + """Return the unregistered cache directory below ComfyUI models.""" + + if self._cache_root is not None: + return self._cache_root + folder_paths = self._folder_paths_module or _folder_paths() + models_dir: Any = folder_paths.models_dir + self._cache_root = Path(str(models_dir)) / CACHE_DIRECTORY_NAME + return self._cache_root + + @property + def lock_directory(self) -> Path: + """Return the internal cross-process lock directory.""" + + return self.root / ".locks" + + def ensure_root(self) -> None: + """Create the cache root and its plain-language ownership notice.""" + + self.root.mkdir(parents=True, exist_ok=True) + readme_path = self.root / README_FILENAME + if not readme_path.is_file(): + readme_path.write_text(README_CONTENT, encoding="utf-8") + + def find_current( + self, + *, + source_model: str, + source_path: Path, + source_size_bytes: int, + source_modified_ns: int, + profile: QuantizationProfile, + recipe: ModelQuantizationRecipe, + ) -> QuantCacheArtifact | None: + """Return a cached artifact matching the source's current file state.""" + + if not self.root.is_dir(): + return None + resolved_source = source_path.resolve() + for artifact in self.list_artifacts(): + if not artifact.manifest.matches_current_source( + source_model, + resolved_source, + source_size_bytes, + source_modified_ns, + profile, + recipe.model_family, + recipe.version, + ): + continue + return self.touch(artifact) + return None + + def find_identity(self, identity: QuantCacheIdentity) -> QuantCacheArtifact | None: + """Return the completed artifact for an exact hashed identity.""" + + directory = self.artifact_directory(identity) + manifest_path = directory / MANIFEST_FILENAME + artifact = self._load_artifact(manifest_path) + if artifact is None or not artifact.manifest.matches_identity(identity): + return None + return self.touch(artifact) + + def create_build_directory(self, identity: QuantCacheIdentity) -> Path: + """Create an isolated temporary directory for one atomic build.""" + + self.ensure_root() + building_root = self.root / ".building" + building_root.mkdir(parents=True, exist_ok=True) + directory = building_root / f"{identity.stable_key}-{uuid.uuid4().hex}" + directory.mkdir() + return directory + + def artifact_filename(self, identity: QuantCacheIdentity) -> str: + """Return a recognizable filename within a cache artifact directory.""" + + stem = _safe_component(Path(identity.source.display_name).stem) + return f"{stem}--{identity.profile.profile_id}{ARTIFACT_FILENAME_SUFFIX}" + + def commit( + self, + identity: QuantCacheIdentity, + build_directory: Path, + ) -> QuantCacheArtifact: + """Atomically publish a validated build directory and its manifest.""" + + self._require_within(build_directory, self.root / ".building") + artifact_file = self.artifact_filename(identity) + artifact_path = build_directory / artifact_file + if not artifact_path.is_file() or artifact_path.stat().st_size <= 0: + raise ValueError( + "Quantized checkpoint build produced no valid artifact file." + ) + manifest = QuantCacheManifest.create( + identity, + artifact_file, + artifact_path.stat().st_size, + ) + self._write_manifest(build_directory / MANIFEST_FILENAME, manifest) + + final_directory = self.artifact_directory(identity) + final_directory.parent.mkdir(parents=True, exist_ok=True) + if final_directory.exists(): + existing = self.find_identity(identity) + if existing is not None: + shutil.rmtree(build_directory) + return existing + self._require_within(final_directory, self.root) + LOGGER.warning( + "replacing invalid quant cache artifact", + extra={"artifact_directory": str(final_directory)}, + ) + shutil.rmtree(final_directory) + build_directory.replace(final_directory) + committed = self._load_artifact(final_directory / MANIFEST_FILENAME) + if committed is None: + raise RuntimeError("Published quant cache artifact could not be read back.") + return committed + + def discard_build(self, build_directory: Path) -> None: + """Remove a failed temporary build without touching completed artifacts.""" + + try: + self._require_within(build_directory, self.root / ".building") + except ValueError: + return + if build_directory.is_dir(): + shutil.rmtree(build_directory, ignore_errors=True) + + def artifact_directory(self, identity: QuantCacheIdentity) -> Path: + """Return the readable directory for an exact derived artifact.""" + + family = _safe_component(identity.model_family) + source = _safe_component(Path(identity.source.display_name).stem) + profile = _safe_component(identity.profile.profile_id) + version = ( + f"{identity.source.sha256[:12]}-profile-{identity.profile.version}" + f"-recipe-{identity.recipe_version}" + ) + return self.root / family / source / profile / version + + def list_artifacts(self) -> tuple[QuantCacheArtifact, ...]: + """Return every valid completed SimpleSyrup-managed artifact.""" + + if not self.root.is_dir(): + return () + artifacts: list[QuantCacheArtifact] = [] + for manifest_path in self.root.rglob(MANIFEST_FILENAME): + if ".building" in manifest_path.parts: + continue + artifact = self._load_artifact(manifest_path) + if artifact is not None: + artifacts.append(artifact) + return tuple(artifacts) + + def touch(self, artifact: QuantCacheArtifact) -> QuantCacheArtifact: + """Record explicit last use for portable LRU behavior.""" + + touched_manifest = artifact.manifest.touched() + self._write_manifest(artifact.manifest_path, touched_manifest) + return QuantCacheArtifact( + path=artifact.path, + manifest_path=artifact.manifest_path, + manifest=touched_manifest, + ) + + def remove(self, artifact: QuantCacheArtifact) -> bool: + """Remove one validated managed artifact directory.""" + + directory = artifact.manifest_path.parent + self._require_within(directory, self.root) + try: + shutil.rmtree(directory) + except OSError as error: + LOGGER.warning( + "quant cache artifact eviction deferred", + extra={"artifact": str(artifact.path), "reason": str(error)}, + ) + return False + self._remove_empty_parents(directory.parent) + return True + + def relative_display_path(self) -> str: + """Return the user-facing location below ComfyUI's model directory.""" + + return f"models/{CACHE_DIRECTORY_NAME}" + + def _load_artifact(self, manifest_path: Path) -> QuantCacheArtifact | None: + """Load one valid managed artifact, ignoring unrelated or corrupt files.""" + + try: + payload = json.loads(manifest_path.read_text(encoding="utf-8")) + manifest = QuantCacheManifest.from_payload(payload) + if Path(manifest.artifact_file).name != manifest.artifact_file: + raise ValueError( + "Quant cache artifact filename must not contain a path." + ) + artifact_path = manifest_path.parent / manifest.artifact_file + if not artifact_path.is_file(): + raise ValueError("Quant cache artifact file is missing.") + if artifact_path.stat().st_size != manifest.artifact_size_bytes: + raise ValueError( + "Quant cache artifact size does not match its manifest." + ) + return QuantCacheArtifact(artifact_path, manifest_path, manifest) + except (JSONDecodeError, OSError, ValueError) as error: + LOGGER.warning( + "ignoring invalid quant cache manifest", + extra={"manifest": str(manifest_path), "reason": str(error)}, + ) + return None + + @staticmethod + def _write_manifest(path: Path, manifest: QuantCacheManifest) -> None: + """Persist a manifest through an atomic same-directory replacement.""" + + temporary_path = path.with_name(f"{path.name}.{uuid.uuid4().hex}.tmp") + try: + temporary_path.write_text( + json.dumps(manifest.to_payload(), indent=2, sort_keys=True) + "\n", + encoding="utf-8", + ) + temporary_path.replace(path) + finally: + temporary_path.unlink(missing_ok=True) + + def _remove_empty_parents(self, directory: Path) -> None: + """Remove empty readable grouping folders without removing the cache root.""" + + current = directory + while current != self.root: + try: + current.rmdir() + except OSError: + return + current = current.parent + + @staticmethod + def _require_within(path: Path, expected_root: Path) -> None: + """Reject destructive operations outside the intended cache subtree.""" + + try: + path.resolve().relative_to(expected_root.resolve()) + except ValueError as error: + raise ValueError( + f"Quant cache path '{path}' is outside '{expected_root}'." + ) from error + + +def _safe_component(value: str) -> str: + """Return a readable path component with unsafe characters replaced.""" + + normalized = _SAFE_COMPONENT_PATTERN.sub("_", value).strip("._") + return normalized[:120] or "model" + + +def _folder_paths() -> ModuleType: + """Import ComfyUI folder paths lazily.""" + + module: Any = importlib.import_module("folder_paths") + if not isinstance(module, ModuleType): + raise TypeError("folder_paths import did not return a module.") + return module diff --git a/simple_syrup/runtime/quant_cache_routes.py b/simple_syrup/runtime/quant_cache_routes.py new file mode 100644 index 0000000..f48eaae --- /dev/null +++ b/simple_syrup/runtime/quant_cache_routes.py @@ -0,0 +1,143 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Expose global quant cache status and safe inactive-artifact clearing.""" + +from __future__ import annotations + +import sys +from collections.abc import Callable, Coroutine +from typing import Any, Protocol, cast + +from aiohttp import web + +from ..services.quant_cache_service import ( + QuantCacheEvictionResult, + QuantCacheService, + QuantCacheStatus, +) +from ..services.quantized_model_boundaries import QuantCacheLimitProvider +from .quant_cache_settings import SettingsQuantCacheLimitProvider + +QUANT_CACHE_ROUTE = "/simple-syrup/quant-cache" +Handler = Callable[[Any], Coroutine[Any, Any, web.Response]] +_REGISTERED_PROMPT_SERVERS: set[int] = set() + + +class QuantCacheRoutesProtocol(Protocol): + """Describe the route decorators used by quant cache endpoints.""" + + def get(self, path: str) -> Callable[[Handler], Handler]: + """Return a GET route decorator.""" + + def delete(self, path: str) -> Callable[[Handler], Handler]: + """Return a DELETE route decorator.""" + + def post(self, path: str) -> Callable[[Handler], Handler]: + """Return a POST route decorator.""" + + +class QuantCachePromptServerProtocol(Protocol): + """Describe the PromptServer state required by cache routes.""" + + routes: QuantCacheRoutesProtocol + + +class QuantCacheServiceBoundary(Protocol): + """Describe global cache operations exposed through HTTP.""" + + def status(self) -> QuantCacheStatus: + """Return current cache state.""" + + def clear_inactive(self) -> QuantCacheEvictionResult: + """Remove every inactive managed artifact.""" + + def enforce_limit(self, limit_bytes: int) -> QuantCacheEvictionResult: + """Apply the current global LRU budget.""" + + +class QuantCacheHandlers: + """Serve global quant cache state and explicit clear requests.""" + + def __init__( + self, + cache_service: QuantCacheServiceBoundary, + limit_provider: QuantCacheLimitProvider, + ) -> None: + """Create handlers with explicit authoritative collaborators.""" + + self._cache_service = cache_service + self._limit_provider = limit_provider + + async def get_status(self, _request: Any) -> web.Response: + """Return current global cache usage and configured limit.""" + + status = self._cache_service.status() + return web.json_response(status.to_payload(self._limit_provider.limit_bytes())) + + async def clear_inactive(self, _request: Any) -> web.Response: + """Clear inactive artifacts and return updated global cache status.""" + + eviction = self._cache_service.clear_inactive() + status = self._cache_service.status() + payload = status.to_payload(self._limit_provider.limit_bytes()) + payload.update( + { + "removed_artifacts": eviction.removed_artifacts, + "removed_bytes": eviction.removed_bytes, + } + ) + return web.json_response(payload) + + async def enforce_limit(self, _request: Any) -> web.Response: + """Apply the persisted limit and return updated global cache status.""" + + limit_bytes = self._limit_provider.limit_bytes() + eviction = self._cache_service.enforce_limit(limit_bytes) + status = self._cache_service.status() + payload = status.to_payload(limit_bytes) + payload.update( + { + "removed_artifacts": eviction.removed_artifacts, + "removed_bytes": eviction.removed_bytes, + } + ) + return web.json_response(payload) + + +def register_quant_cache_routes( + cache_service: QuantCacheServiceBoundary | None = None, + limit_provider: QuantCacheLimitProvider | None = None, + prompt_server: QuantCachePromptServerProtocol | None = None, +) -> bool: + """Register global cache routes with ComfyUI when PromptServer is available.""" + + server_instance = prompt_server or _prompt_server_instance() + if server_instance is None: + return False + server_key = id(server_instance) + if prompt_server is None and server_key in _REGISTERED_PROMPT_SERVERS: + return True + handlers = QuantCacheHandlers( + cache_service or QuantCacheService(), + limit_provider or SettingsQuantCacheLimitProvider(), + ) + server_instance.routes.get(QUANT_CACHE_ROUTE)(handlers.get_status) + server_instance.routes.post(QUANT_CACHE_ROUTE)(handlers.enforce_limit) + server_instance.routes.delete(QUANT_CACHE_ROUTE)(handlers.clear_inactive) + if prompt_server is None: + _REGISTERED_PROMPT_SERVERS.add(server_key) + return True + + +def _prompt_server_instance() -> QuantCachePromptServerProtocol | None: + """Return ComfyUI's PromptServer instance when available.""" + + try: + server_module = sys.modules["server"] + prompt_server = server_module.PromptServer + instance = prompt_server.instance + except (KeyError, AttributeError): + return None + return cast(QuantCachePromptServerProtocol, instance) diff --git a/simple_syrup/runtime/quant_cache_settings.py b/simple_syrup/runtime/quant_cache_settings.py new file mode 100644 index 0000000..15ff037 --- /dev/null +++ b/simple_syrup/runtime/quant_cache_settings.py @@ -0,0 +1,28 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Adapt persisted SimpleSyrup settings to the quant cache byte budget.""" + +from __future__ import annotations + +from .settings_repository import SimpleSyrupSettingsRepository + +BYTES_PER_GIB = 1024**3 + + +class SettingsQuantCacheLimitProvider: + """Read the authoritative global quant cache limit from backend settings.""" + + def __init__( + self, + repository: SimpleSyrupSettingsRepository | None = None, + ) -> None: + """Create a provider with injectable settings persistence.""" + + self._repository = repository or SimpleSyrupSettingsRepository() + + def limit_bytes(self) -> int: + """Return the configured integer GiB limit in bytes.""" + + return self._repository.load().quant_cache_limit_gib * BYTES_PER_GIB diff --git a/simple_syrup/runtime/quantization_capabilities.py b/simple_syrup/runtime/quantization_capabilities.py new file mode 100644 index 0000000..205a460 --- /dev/null +++ b/simple_syrup/runtime/quantization_capabilities.py @@ -0,0 +1,128 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Discover quantization profiles executable and reloadable by active ComfyUI.""" + +from __future__ import annotations + +import importlib +from collections.abc import Iterable +from types import ModuleType +from typing import Any + +from ..domain.model_quantization import QuantizationFormat, QuantizationProfile +from ..shared.logging import get_logger +from .comfy_safetensors_dtypes import ComfySafetensorsDtypeRegistry + +LOGGER = get_logger(__name__) + + +class QuantizationCapabilityCatalog: + """Intersect profile requirements with GPU, kernels, and checkpoint loading.""" + + def __init__( + self, + quant_ops_module: ModuleType | None = None, + model_management_module: ModuleType | None = None, + dtype_registry: ComfySafetensorsDtypeRegistry | None = None, + ) -> None: + """Create a catalog with injectable ComfyUI runtime boundaries.""" + + self._quant_ops_module = quant_ops_module + self._model_management_module = model_management_module + self._dtype_registry = dtype_registry or ComfySafetensorsDtypeRegistry() + self._cached_formats: frozenset[QuantizationFormat] | None = None + + def available_formats(self) -> frozenset[QuantizationFormat]: + """Return formats supported through compute and checkpoint reload.""" + + if self._cached_formats is not None: + return self._cached_formats + formats: set[QuantizationFormat] = set() + try: + quant_ops = self._quant_ops() + model_management = self._model_management() + registered = quant_ops.QUANT_ALGOS + device = model_management.get_torch_device() + if model_management.supports_fp8_compute(device): + self._append_registered( + formats, registered, QuantizationFormat.FP8_E4M3 + ) + self._append_registered( + formats, registered, QuantizationFormat.FP8_E5M2 + ) + if model_management.supports_nvfp4_compute(device): + self._append_registered(formats, registered, QuantizationFormat.NVFP4) + if model_management.supports_mxfp8_compute(device): + self._append_registered(formats, registered, QuantizationFormat.MXFP8) + except ( + AttributeError, + ImportError, + OSError, + RuntimeError, + ValueError, + ) as error: + LOGGER.warning( + "quantization capability discovery failed", + extra={"reason": str(error)}, + ) + self._cached_formats = frozenset(formats) + return self._cached_formats + + def available_profiles( + self, + profiles: Iterable[QuantizationProfile], + ) -> tuple[QuantizationProfile, ...]: + """Return profiles whose complete format set is available.""" + + available = self.available_formats() + return tuple( + profile + for profile in profiles + if profile.is_original or profile.required_formats.issubset(available) + ) + + def selection_labels(self, profiles: Iterable[QuantizationProfile]) -> list[str]: + """Return workflow labels in recipe-defined preference order.""" + + return [profile.label for profile in self.available_profiles(profiles)] + + def require_available(self, profile: QuantizationProfile) -> None: + """Reject a profile unavailable to the current ComfyUI runtime.""" + + missing = profile.required_formats.difference(self.available_formats()) + if not missing: + return + missing_labels = ", ".join(sorted(item.label for item in missing)) + raise ValueError( + f"{profile.label} is unavailable on the current ComfyUI GPU/runtime. " + f"Missing format support: {missing_labels}." + ) + + def _quant_ops(self) -> Any: + """Return ComfyUI's quantization operation module.""" + + return self._quant_ops_module or importlib.import_module("comfy.quant_ops") + + def _model_management(self) -> Any: + """Return ComfyUI's device capability module.""" + + return self._model_management_module or importlib.import_module( + "comfy.model_management" + ) + + def _append_registered( + self, + formats: set[QuantizationFormat], + registered: object, + candidate: QuantizationFormat, + ) -> None: + """Append only when Comfy registers, computes, and reloads the format.""" + + if ( + isinstance(registered, dict) + and candidate.value in registered + and self._dtype_registry.supports_format(candidate) + ): + formats.add(candidate) diff --git a/simple_syrup/runtime/quantization_progress.py b/simple_syrup/runtime/quantization_progress.py new file mode 100644 index 0000000..89e6ef0 --- /dev/null +++ b/simple_syrup/runtime/quantization_progress.py @@ -0,0 +1,79 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Publish quiet byte-weighted quantization progress through ComfyUI nodes.""" + +from __future__ import annotations + +import importlib +from typing import Protocol, cast + + +class QuantizationProgressReporter(Protocol): + """Report absolute byte-weighted quantization progress.""" + + def start(self, label: str, total: int) -> None: + """Begin one quantization operation.""" + + def advance(self, current: int, total: int) -> None: + """Report absolute completed work.""" + + def finish(self) -> None: + """Mark the quantization operation complete.""" + + +class _ComfyProgressBar(Protocol): + """Describe the ComfyUI progress method used by this adapter.""" + + def update_absolute(self, value: int, total: int | None = None) -> None: + """Publish absolute progress.""" + + +class NullQuantizationProgressReporter: + """Ignore quantization progress outside a ComfyUI execution context.""" + + def start(self, label: str, total: int) -> None: + """Ignore the operation start.""" + + def advance(self, current: int, total: int) -> None: + """Ignore one progress update.""" + + def finish(self) -> None: + """Ignore operation completion.""" + + +class ComfyQuantizationProgressReporter: + """Attach quantization progress to the currently executing ComfyUI node.""" + + def __init__(self) -> None: + """Create an adapter that allocates its bar only when work starts.""" + + self._progress_bar: object | None = None + self._total = 1 + + def start(self, label: str, total: int) -> None: + """Create a standard ComfyUI progress bar without console chatter.""" + + del label + self._total = max(total, 1) + comfy_utils = importlib.import_module("comfy.utils") + self._progress_bar = comfy_utils.ProgressBar(self._total) + self.advance(0, self._total) + + def advance(self, current: int, total: int) -> None: + """Publish bounded absolute progress.""" + + if self._progress_bar is None: + return + self._total = max(total, 1) + progress_bar = cast(_ComfyProgressBar, self._progress_bar) + progress_bar.update_absolute(min(max(current, 0), self._total), self._total) + + def finish(self) -> None: + """Publish completion when an operation was started.""" + + if self._progress_bar is None: + return + progress_bar = cast(_ComfyProgressBar, self._progress_bar) + progress_bar.update_absolute(self._total, self._total) diff --git a/simple_syrup/runtime/settings.py b/simple_syrup/runtime/settings.py index 9054b04..6d2897f 100644 --- a/simple_syrup/runtime/settings.py +++ b/simple_syrup/runtime/settings.py @@ -6,13 +6,7 @@ from __future__ import annotations -import importlib -import json from dataclasses import dataclass, field -from json import JSONDecodeError -from pathlib import Path -from types import ModuleType -from typing import Any, Final from ..domain.external_llm import ( ExternalLLMConfigError, @@ -22,7 +16,9 @@ from ..domain.external_llm import ( from ..shared.logging import get_logger LOGGER = get_logger(__name__) -SETTINGS_FILENAME: Final = "settings.json" +DEFAULT_QUANT_CACHE_LIMIT_GIB = 20 +MIN_QUANT_CACHE_LIMIT_GIB = 1 +MAX_QUANT_CACHE_LIMIT_GIB = 2048 class SimpleSyrupSettingsError(ValueError): @@ -97,6 +93,7 @@ class SimpleSyrupSettings: """User-configurable SimpleSyrup runtime settings.""" show_downloadable_models: bool = True + quant_cache_limit_gib: int = DEFAULT_QUANT_CACHE_LIMIT_GIB external_llm: ExternalLLMSettings = field(default_factory=ExternalLLMSettings) def to_payload(self) -> dict[str, object]: @@ -104,6 +101,7 @@ class SimpleSyrupSettings: return { "show_downloadable_models": self.show_downloadable_models, + "quant_cache_limit_gib": self.quant_cache_limit_gib, "external_llm": self.external_llm.to_payload(), } @@ -123,6 +121,22 @@ class SimpleSyrupSettings: "show_downloadable_models to be a boolean." ) + quant_cache_limit = payload.get( + "quant_cache_limit_gib", DEFAULT_QUANT_CACHE_LIMIT_GIB + ) + if ( + not isinstance(quant_cache_limit, int) + or isinstance(quant_cache_limit, bool) + or not MIN_QUANT_CACHE_LIMIT_GIB + <= quant_cache_limit + <= MAX_QUANT_CACHE_LIMIT_GIB + ): + raise SimpleSyrupSettingsError( + "SimpleSyrup settings payload is invalid. Expected " + f"quant_cache_limit_gib to be an integer from " + f"{MIN_QUANT_CACHE_LIMIT_GIB} to {MAX_QUANT_CACHE_LIMIT_GIB}." + ) + try: external_llm = ExternalLLMSettings.from_payload(payload.get("external_llm")) except SimpleSyrupSettingsError as error: @@ -132,87 +146,8 @@ class SimpleSyrupSettings: ) external_llm = ExternalLLMSettings() - return cls(show_downloadable_models=value, external_llm=external_llm) - - -class SimpleSyrupSettingsRepository: - """Load and save SimpleSyrup settings from Comfy's user directory.""" - - def __init__( - self, - settings_path: Path | None = None, - folder_paths_module: ModuleType | None = None, - ) -> None: - """Create a repository with injectable filesystem and Comfy boundaries.""" - - self._settings_path = settings_path - self._folder_paths_module = folder_paths_module - - def load(self) -> SimpleSyrupSettings: - """Load settings or return defaults for missing/malformed files.""" - - path = self.settings_path() - if not path.is_file(): - return SimpleSyrupSettings() - - try: - payload = json.loads(path.read_text(encoding="utf-8")) - return SimpleSyrupSettings.from_payload(payload) - except (JSONDecodeError, OSError, SimpleSyrupSettingsError) as error: - LOGGER.warning( - "using default settings after failed load", - extra={"settings_path": str(path), "reason": str(error)}, - ) - return SimpleSyrupSettings() - - def save(self, settings: SimpleSyrupSettings) -> SimpleSyrupSettings: - """Persist validated settings and return the saved value.""" - - path = self.settings_path() - path.parent.mkdir(parents=True, exist_ok=True) - temporary_path = path.with_name(f"{path.name}.tmp") - temporary_path.write_text( - json.dumps(settings.to_payload(), indent=2, sort_keys=True) + "\n", - encoding="utf-8", + return cls( + show_downloadable_models=value, + quant_cache_limit_gib=quant_cache_limit, + external_llm=external_llm, ) - temporary_path.replace(path) - return settings - - def settings_path(self) -> Path: - """Return the resolved settings path.""" - - if self._settings_path is not None: - return self._settings_path - - folder_paths = self._folder_paths_module or _folder_paths() - return ( - _user_directory(folder_paths) - / "default" - / "SimpleSyrup" - / SETTINGS_FILENAME - ) - - -def _user_directory(folder_paths: ModuleType) -> Path: - """Return Comfy's user directory from stable APIs or conservative fallback.""" - - get_user_directory = getattr(folder_paths, "get_user_directory", None) - if callable(get_user_directory): - user_directory = get_user_directory() - return Path(str(user_directory)) - - user_directory_attribute = getattr(folder_paths, "user_directory", None) - if user_directory_attribute is not None: - return Path(str(user_directory_attribute)) - - models_dir: Any = folder_paths.models_dir - return Path(str(models_dir)).parent / "user" - - -def _folder_paths() -> ModuleType: - """Import ComfyUI folder paths lazily.""" - - module: Any = importlib.import_module("folder_paths") - if not isinstance(module, ModuleType): - raise TypeError("folder_paths import did not return a module.") - return module diff --git a/simple_syrup/runtime/settings_repository.py b/simple_syrup/runtime/settings_repository.py new file mode 100644 index 0000000..895c481 --- /dev/null +++ b/simple_syrup/runtime/settings_repository.py @@ -0,0 +1,103 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Persist validated SimpleSyrup settings in ComfyUI's user directory.""" + +from __future__ import annotations + +import importlib +import json +from json import JSONDecodeError +from pathlib import Path +from types import ModuleType +from typing import Any, Final + +from ..shared.logging import get_logger +from .settings import SimpleSyrupSettings, SimpleSyrupSettingsError + +LOGGER = get_logger(__name__) +SETTINGS_FILENAME: Final = "settings.json" + + +class SimpleSyrupSettingsRepository: + """Load and save SimpleSyrup settings from Comfy's user directory.""" + + def __init__( + self, + settings_path: Path | None = None, + folder_paths_module: ModuleType | None = None, + ) -> None: + """Create a repository with injectable filesystem and Comfy boundaries.""" + + self._settings_path = settings_path + self._folder_paths_module = folder_paths_module + + def load(self) -> SimpleSyrupSettings: + """Load settings or return defaults for missing/malformed files.""" + + path = self.settings_path() + if not path.is_file(): + return SimpleSyrupSettings() + + try: + payload = json.loads(path.read_text(encoding="utf-8")) + return SimpleSyrupSettings.from_payload(payload) + except (JSONDecodeError, OSError, SimpleSyrupSettingsError) as error: + LOGGER.warning( + "using default settings after failed load", + extra={"settings_path": str(path), "reason": str(error)}, + ) + return SimpleSyrupSettings() + + def save(self, settings: SimpleSyrupSettings) -> SimpleSyrupSettings: + """Persist validated settings and return the saved value.""" + + path = self.settings_path() + path.parent.mkdir(parents=True, exist_ok=True) + temporary_path = path.with_name(f"{path.name}.tmp") + temporary_path.write_text( + json.dumps(settings.to_payload(), indent=2, sort_keys=True) + "\n", + encoding="utf-8", + ) + temporary_path.replace(path) + return settings + + def settings_path(self) -> Path: + """Return the resolved settings path.""" + + if self._settings_path is not None: + return self._settings_path + + folder_paths = self._folder_paths_module or _folder_paths() + return ( + _user_directory(folder_paths) + / "default" + / "SimpleSyrup" + / SETTINGS_FILENAME + ) + + +def _user_directory(folder_paths: ModuleType) -> Path: + """Return Comfy's user directory from stable APIs or conservative fallback.""" + + get_user_directory = getattr(folder_paths, "get_user_directory", None) + if callable(get_user_directory): + user_directory = get_user_directory() + return Path(str(user_directory)) + + user_directory_attribute = getattr(folder_paths, "user_directory", None) + if user_directory_attribute is not None: + return Path(str(user_directory_attribute)) + + models_dir: Any = folder_paths.models_dir + return Path(str(models_dir)).parent / "user" + + +def _folder_paths() -> ModuleType: + """Import ComfyUI folder paths lazily.""" + + module: Any = importlib.import_module("folder_paths") + if not isinstance(module, ModuleType): + raise TypeError("folder_paths import did not return a module.") + return module diff --git a/simple_syrup/runtime/settings_routes.py b/simple_syrup/runtime/settings_routes.py index 123abc7..d6f1b23 100644 --- a/simple_syrup/runtime/settings_routes.py +++ b/simple_syrup/runtime/settings_routes.py @@ -16,8 +16,8 @@ from ..shared.logging import get_logger from .settings import ( SimpleSyrupSettings, SimpleSyrupSettingsError, - SimpleSyrupSettingsRepository, ) +from .settings_repository import SimpleSyrupSettingsRepository LOGGER = get_logger(__name__) SETTINGS_ROUTE = "/simple-syrup/settings" @@ -80,13 +80,22 @@ class SettingsHandlers: """Return validated settings while preserving omitted nested config.""" settings = SimpleSyrupSettings.from_payload(payload) - if isinstance(payload, dict) and "external_llm" not in payload: - current = self._repository.load() - return SimpleSyrupSettings( - show_downloadable_models=settings.show_downloadable_models, - external_llm=current.external_llm, - ) - return settings + if not isinstance(payload, dict): + return settings + current = self._repository.load() + return SimpleSyrupSettings( + show_downloadable_models=settings.show_downloadable_models, + quant_cache_limit_gib=( + settings.quant_cache_limit_gib + if "quant_cache_limit_gib" in payload + else current.quant_cache_limit_gib + ), + external_llm=( + settings.external_llm + if "external_llm" in payload + else current.external_llm + ), + ) def register_settings_routes( diff --git a/simple_syrup/services/anima_diffusion_model_service.py b/simple_syrup/services/anima_diffusion_model_service.py new file mode 100644 index 0000000..1c72216 --- /dev/null +++ b/simple_syrup/services/anima_diffusion_model_service.py @@ -0,0 +1,89 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Resolve and load Anima's original or recipe-quantized diffusion model.""" + +from __future__ import annotations + +from types import ModuleType +from typing import Any + +from ..domain.anima_quantization import AnimaQuantizationRecipe +from ..runtime.diffusion_model_loader import DiffusionModelLoader +from ..runtime.quant_cache_leases import ( + GLOBAL_QUANT_CACHE_LEASES, + QuantCacheLeaseRegistry, +) +from ..runtime.quant_cache_repository import QuantCacheRepository +from ..runtime.quantization_progress import QuantizationProgressReporter +from .quantized_model_boundaries import ( + DiffusionModelPathLoaderBoundary, + QuantizedModelResolverBoundary, +) +from .quantized_model_resolver import QuantizedModelResolver + + +class AnimaDiffusionModelService: + """Apply Anima policy while reusing global quantization infrastructure.""" + + def __init__( + self, + diffusion_loader: DiffusionModelPathLoaderBoundary | None = None, + resolver: QuantizedModelResolverBoundary | None = None, + leases: QuantCacheLeaseRegistry | None = None, + folder_paths_module: ModuleType | None = None, + ) -> None: + """Create a service with injectable loader, resolver, and lease owners.""" + + self._diffusion_loader = diffusion_loader or DiffusionModelLoader( + folder_paths_module + ) + self._resolver = resolver or QuantizedModelResolver( + repository=QuantCacheRepository(folder_paths_module=folder_paths_module) + ) + self._leases = leases or GLOBAL_QUANT_CACHE_LEASES + self._recipe = AnimaQuantizationRecipe() + + def load( + self, + *, + diffusion_model: str, + diffusion_weight_dtype: str, + quantization: str, + progress: QuantizationProgressReporter | None = None, + ) -> object: + """Load the selected source or a globally cached Anima derivative.""" + + source_path = self._diffusion_loader.resolve_path(diffusion_model) + resolved = self._resolver.resolve( + source_model=diffusion_model, + source_path=source_path, + quantization=quantization, + recipe=self._recipe, + progress=progress, + ) + effective_weight_dtype = ( + diffusion_weight_dtype if resolved.profile.is_original else "default" + ) + try: + model = self._diffusion_loader.load_path( + resolved.path, + effective_weight_dtype, + ) + if resolved.cache_artifact is not None: + self._leases.lease(resolved.cache_artifact.path, model) + finally: + if resolved.reservation is not None: + resolved.reservation.release() + model_boundary: Any = model + model_boundary.simple_syrup_model_provenance = { + "source_model": diffusion_model, + "quantization_profile": resolved.profile.profile_id, + "quantization_profile_label": resolved.profile.label, + "quantization_profile_version": resolved.profile.version, + "derived_path": str(resolved.path) + if resolved.cache_artifact is not None + else None, + } + return model diff --git a/simple_syrup/runtime/anima_loader.py b/simple_syrup/services/anima_loader_service.py similarity index 73% rename from simple_syrup/runtime/anima_loader.py rename to simple_syrup/services/anima_loader_service.py index 9013e9b..decf687 100644 --- a/simple_syrup/runtime/anima_loader.py +++ b/simple_syrup/services/anima_loader_service.py @@ -13,19 +13,15 @@ from typing import Any, Protocol import torch -from .anima_artifacts import ANIMA_QWEN_TEXT_ENCODER, ANIMA_QWEN_VAE -from .auto_model_artifact import AutoModelArtifact -from .auto_model_resolver import AutoModelResolution, AutoModelResolver -from .model_downloads import ProgressReporter -from .vae_loader import VaeLoaderService, load_vae_path +from ..runtime.anima_artifacts import ANIMA_QWEN_TEXT_ENCODER, ANIMA_QWEN_VAE +from ..runtime.auto_model_artifact import AutoModelArtifact +from ..runtime.auto_model_resolver import AutoModelResolution, AutoModelResolver +from ..runtime.model_downloads import ProgressReporter +from ..runtime.quantization_progress import QuantizationProgressReporter +from ..runtime.vae_loader import VaeLoaderService, load_vae_path +from .anima_diffusion_model_service import AnimaDiffusionModelService AUTO_CHOICE = "auto" -DIFFUSION_WEIGHT_DTYPES = ( - "default", - "fp8_e4m3fn", - "fp8_e4m3fn_fast", - "fp8_e5m2", -) CLIP_DEVICES = ("default", "cpu") DEFAULT_CLIP_TYPE = "stable_diffusion" @@ -36,6 +32,7 @@ class AnimaLoaderService: def __init__( self, resolver: AutoModelResolverBoundary | None = None, + diffusion_service: AnimaDiffusionModelService | None = None, folder_paths_module: ModuleType | None = None, ) -> None: """Create a loader with injectable auto-resolution boundaries.""" @@ -44,21 +41,31 @@ class AnimaLoaderService: self._resolver = resolver or AutoModelResolver( folder_paths_module=folder_paths_module ) + self._diffusion_service = diffusion_service or AnimaDiffusionModelService( + folder_paths_module=folder_paths_module + ) self._vae_loader = VaeLoaderService(folder_paths_module) def load_models( self, diffusion_model: str, + quantization: str, diffusion_weight_dtype: str, text_encoder: str, text_encoder_device: str, vae: str, progress: ProgressReporter | None = None, + quantization_progress: QuantizationProgressReporter | None = None, ) -> tuple[object, object, object]: """Return ComfyUI MODEL, CLIP, and VAE objects.""" return ( - self._load_diffusion_model(diffusion_model, diffusion_weight_dtype), + self._diffusion_service.load( + diffusion_model=diffusion_model, + diffusion_weight_dtype=diffusion_weight_dtype, + quantization=quantization, + progress=quantization_progress, + ), self._load_clip( text_encoder, text_encoder_device, @@ -67,37 +74,6 @@ class AnimaLoaderService: self._load_vae(vae, progress), ) - def _load_diffusion_model( - self, - diffusion_model: str, - diffusion_weight_dtype: str, - ) -> object: - """Load a diffusion model using ComfyUI's diffusion model loader policy.""" - - if diffusion_weight_dtype not in DIFFUSION_WEIGHT_DTYPES: - valid = ", ".join(DIFFUSION_WEIGHT_DTYPES) - raise ValueError(f"diffusion_weight_dtype must be one of: {valid}.") - - model_options: dict[str, object] = {} - if diffusion_weight_dtype == "fp8_e4m3fn": - model_options["dtype"] = torch.float8_e4m3fn - elif diffusion_weight_dtype == "fp8_e4m3fn_fast": - model_options["dtype"] = torch.float8_e4m3fn - model_options["fp8_optimizations"] = True - elif diffusion_weight_dtype == "fp8_e5m2": - model_options["dtype"] = torch.float8_e5m2 - - folder_paths = self._folder_paths() - unet_path = folder_paths.get_full_path_or_raise( - "diffusion_models", - diffusion_model, - ) - comfy_sd = _comfy_sd() - return comfy_sd.load_diffusion_model( - unet_path, - model_options=model_options, - ) - def _load_clip( self, text_encoder: str, diff --git a/simple_syrup/services/external_llm_prompt_service.py b/simple_syrup/services/external_llm_prompt_service.py index 75e82c1..b3835b2 100644 --- a/simple_syrup/services/external_llm_prompt_service.py +++ b/simple_syrup/services/external_llm_prompt_service.py @@ -24,8 +24,8 @@ from ..runtime.external_llm_keyring import ExternalLLMKeyringStore from ..runtime.settings import ( ExternalLLMSettings, SimpleSyrupSettings, - SimpleSyrupSettingsRepository, ) +from ..runtime.settings_repository import SimpleSyrupSettingsRepository CONFIGURE_EXTERNAL_LLM = "Configure external LLM endpoint" diff --git a/simple_syrup/services/quant_cache_service.py b/simple_syrup/services/quant_cache_service.py new file mode 100644 index 0000000..548fafe --- /dev/null +++ b/simple_syrup/services/quant_cache_service.py @@ -0,0 +1,127 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Apply global LRU policy to inactive SimpleSyrup quant cache artifacts.""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path + +from ..runtime.quant_cache_leases import ( + GLOBAL_QUANT_CACHE_LEASES, + QuantCacheLeaseRegistry, +) +from ..runtime.quant_cache_repository import QuantCacheArtifact, QuantCacheRepository +from ..shared.logging import get_logger + +LOGGER = get_logger(__name__) + + +@dataclass(frozen=True) +class QuantCacheStatus: + """Summarize the global cache for settings and diagnostics.""" + + path: str + usage_bytes: int + artifact_count: int + active_artifact_count: int + + def to_payload(self, limit_bytes: int) -> dict[str, object]: + """Return the settings-route JSON representation.""" + + return { + "path": self.path, + "usage_bytes": self.usage_bytes, + "limit_bytes": limit_bytes, + "artifact_count": self.artifact_count, + "active_artifact_count": self.active_artifact_count, + } + + +@dataclass(frozen=True) +class QuantCacheEvictionResult: + """Summarize one bounded eviction attempt.""" + + removed_artifacts: int + removed_bytes: int + remaining_bytes: int + + +class QuantCacheService: + """Own global cache accounting, LRU eviction, and safe clearing.""" + + def __init__( + self, + repository: QuantCacheRepository | None = None, + leases: QuantCacheLeaseRegistry | None = None, + ) -> None: + """Create a cache service with injectable persistence and leases.""" + + self._repository = repository or QuantCacheRepository() + self._leases = leases or GLOBAL_QUANT_CACHE_LEASES + + def status(self) -> QuantCacheStatus: + """Return current managed artifact usage.""" + + artifacts = self._repository.list_artifacts() + return QuantCacheStatus( + path=self._repository.relative_display_path(), + usage_bytes=sum(item.manifest.artifact_size_bytes for item in artifacts), + artifact_count=len(artifacts), + active_artifact_count=self._leases.active_count(), + ) + + def enforce_limit( + self, + limit_bytes: int, + protected_paths: frozenset[Path] = frozenset(), + ) -> QuantCacheEvictionResult: + """Evict least-recently-used inactive artifacts until within the limit.""" + + if limit_bytes < 0: + raise ValueError("Quant cache limit must not be negative.") + artifacts = list(self._repository.list_artifacts()) + remaining = sum(item.manifest.artifact_size_bytes for item in artifacts) + removed_count = 0 + removed_bytes = 0 + protected = {path.resolve() for path in protected_paths} + for artifact in sorted(artifacts, key=_last_used): + if remaining <= limit_bytes: + break + if artifact.path.resolve() in protected or self._leases.is_leased( + artifact.path + ): + continue + if not self._repository.remove(artifact): + continue + removed_count += 1 + removed_bytes += artifact.manifest.artifact_size_bytes + remaining -= artifact.manifest.artifact_size_bytes + if removed_count: + LOGGER.info( + "quant cache eviction completed", + extra={ + "removed_artifacts": removed_count, + "removed_bytes": removed_bytes, + "remaining_bytes": remaining, + "limit_bytes": limit_bytes, + }, + ) + return QuantCacheEvictionResult(removed_count, removed_bytes, remaining) + + def clear_inactive(self) -> QuantCacheEvictionResult: + """Remove every inactive managed artifact and preserve loaded models.""" + + return self.enforce_limit(0) + + +def _last_used(artifact: QuantCacheArtifact) -> float: + """Return the manifest timestamp used for deterministic LRU ordering.""" + + try: + return datetime.fromisoformat(artifact.manifest.last_used_at).timestamp() + except ValueError: + return float("-inf") diff --git a/simple_syrup/services/quantized_model_boundaries.py b/simple_syrup/services/quantized_model_boundaries.py new file mode 100644 index 0000000..498d3d9 --- /dev/null +++ b/simple_syrup/services/quantized_model_boundaries.py @@ -0,0 +1,83 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Define collaboration boundaries for reusable quantized model resolution.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +from typing import Protocol + +from ..domain.model_quantization import ModelQuantizationRecipe, QuantizationProfile +from ..domain.quant_cache import SourceCheckpointIdentity +from ..runtime.quant_cache_leases import QuantCacheReservation +from ..runtime.quant_cache_repository import QuantCacheArtifact +from ..runtime.quantization_progress import QuantizationProgressReporter + + +@dataclass(frozen=True) +class ResolvedModelCheckpoint: + """Describe the actual checkpoint selected for one loader execution.""" + + path: Path + profile: QuantizationProfile + cache_artifact: QuantCacheArtifact | None = None + reservation: QuantCacheReservation | None = None + + +class CheckpointQuantizerBoundary(Protocol): + """Convert one source checkpoint into a ComfyUI quantized checkpoint.""" + + def quantize( + self, + *, + source: SourceCheckpointIdentity, + destination_path: Path, + profile: QuantizationProfile, + recipe: ModelQuantizationRecipe, + progress: QuantizationProgressReporter, + progress_base: int, + progress_total: int, + ) -> object: + """Write and validate one derived checkpoint.""" + + +class QuantCacheLimitProvider(Protocol): + """Provide the current global quant cache byte limit.""" + + def limit_bytes(self) -> int: + """Return the validated global byte budget.""" + + +class QuantizationCapabilityBoundary(Protocol): + """Validate requested formats against ComfyUI and the active GPU.""" + + def require_available(self, profile: QuantizationProfile) -> None: + """Reject an unavailable profile.""" + + +class DiffusionModelPathLoaderBoundary(Protocol): + """Resolve and load standalone diffusion checkpoints by path.""" + + def resolve_path(self, diffusion_model: str) -> Path: + """Resolve one workflow model name to its source path.""" + + def load_path(self, model_path: Path, weight_dtype: str) -> object: + """Load one resolved diffusion checkpoint.""" + + +class QuantizedModelResolverBoundary(Protocol): + """Resolve source checkpoints to original or cached derivative paths.""" + + def resolve( + self, + *, + source_model: str, + source_path: Path, + quantization: str, + recipe: ModelQuantizationRecipe, + progress: QuantizationProgressReporter | None = None, + ) -> ResolvedModelCheckpoint: + """Return one protected model checkpoint resolution.""" diff --git a/simple_syrup/services/quantized_model_resolver.py b/simple_syrup/services/quantized_model_resolver.py new file mode 100644 index 0000000..72f85ed --- /dev/null +++ b/simple_syrup/services/quantized_model_resolver.py @@ -0,0 +1,252 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Resolve original checkpoints or globally cached quantized derivatives.""" + +from __future__ import annotations + +import hashlib +import shutil +import time +from pathlib import Path + +from ..domain.model_quantization import ModelQuantizationRecipe, QuantizationProfile +from ..domain.quant_cache import QuantCacheIdentity, SourceCheckpointIdentity +from ..runtime.checkpoint_quantizer import SafetensorsCheckpointQuantizer +from ..runtime.quant_cache_leases import ( + GLOBAL_QUANT_CACHE_LEASES, + QuantCacheLeaseRegistry, +) +from ..runtime.quant_cache_lock import QuantCacheBuildCoordinator +from ..runtime.quant_cache_repository import QuantCacheArtifact, QuantCacheRepository +from ..runtime.quant_cache_settings import SettingsQuantCacheLimitProvider +from ..runtime.quantization_capabilities import QuantizationCapabilityCatalog +from ..runtime.quantization_progress import ( + NullQuantizationProgressReporter, + QuantizationProgressReporter, +) +from ..shared.logging import get_logger +from .quant_cache_service import QuantCacheService +from .quantized_model_boundaries import ( + CheckpointQuantizerBoundary, + QuantCacheLimitProvider, + QuantizationCapabilityBoundary, + ResolvedModelCheckpoint, +) + +LOGGER = get_logger(__name__) +HASH_CHUNK_SIZE = 8 * 1024 * 1024 +MINIMUM_DISK_HEADROOM = 512 * 1024 * 1024 + + +class QuantizedModelResolver: + """Coordinate capability checks, generation, publication, and global LRU.""" + + def __init__( + self, + repository: QuantCacheRepository | None = None, + quantizer: CheckpointQuantizerBoundary | None = None, + capabilities: QuantizationCapabilityBoundary | None = None, + cache_service: QuantCacheService | None = None, + limit_provider: QuantCacheLimitProvider | None = None, + coordinator: QuantCacheBuildCoordinator | None = None, + leases: QuantCacheLeaseRegistry | None = None, + ) -> None: + """Create a resolver with injectable architecture boundaries.""" + + self._repository = repository or QuantCacheRepository() + self._leases = leases or GLOBAL_QUANT_CACHE_LEASES + self._quantizer = quantizer or SafetensorsCheckpointQuantizer() + self._capabilities = capabilities or QuantizationCapabilityCatalog() + self._cache_service = cache_service or QuantCacheService( + self._repository, self._leases + ) + self._limit_provider = limit_provider or SettingsQuantCacheLimitProvider() + self._coordinator = coordinator or QuantCacheBuildCoordinator( + self._repository.lock_directory + ) + + def resolve( + self, + *, + source_model: str, + source_path: Path, + quantization: str, + recipe: ModelQuantizationRecipe, + progress: QuantizationProgressReporter | None = None, + ) -> ResolvedModelCheckpoint: + """Return the original path or create and return its cached derivative.""" + + profile = recipe.profile_from_selection(quantization) + if profile.is_original: + return ResolvedModelCheckpoint(source_path, profile) + self._capabilities.require_available(profile) + + resolved_source = source_path.resolve() + source_stat = resolved_source.stat() + cached = self._repository.find_current( + source_model=source_model, + source_path=resolved_source, + source_size_bytes=source_stat.st_size, + source_modified_ns=source_stat.st_mtime_ns, + profile=profile, + recipe=recipe, + ) + if cached is not None: + LOGGER.debug( + "using cached quantized checkpoint", + extra={ + "source_model": source_model, + "quantization_profile": profile.profile_id, + }, + ) + return self._cached_resolution(cached, profile) + + reporter = progress or NullQuantizationProgressReporter() + total_work = max(source_stat.st_size * 2, 1) + reporter.start( + f"Creating {profile.label} cache for {source_model}", + total_work, + ) + source_sha256 = _sha256_with_progress( + resolved_source, + reporter, + total_work, + ) + identity = QuantCacheIdentity( + source=SourceCheckpointIdentity( + display_name=source_model, + path=resolved_source, + size_bytes=source_stat.st_size, + modified_ns=source_stat.st_mtime_ns, + sha256=source_sha256, + ), + profile=profile, + model_family=recipe.model_family, + recipe_version=recipe.version, + ) + + with self._coordinator.acquire(identity.stable_key): + exact_cached = self._repository.find_identity(identity) + if exact_cached is not None: + reporter.finish() + return self._cached_resolution(exact_cached, profile) + self._require_disk_space(identity) + started = time.monotonic() + LOGGER.info( + "creating %s cache for %s", + profile.label, + source_model, + extra={ + "source_model": source_model, + "quantization_profile": profile.profile_id, + "model_family": recipe.model_family, + "recipe_version": recipe.version, + }, + ) + build_directory = self._repository.create_build_directory(identity) + try: + destination = build_directory / self._repository.artifact_filename( + identity + ) + result = self._quantizer.quantize( + source=identity.source, + destination_path=destination, + profile=profile, + recipe=recipe, + progress=reporter, + progress_base=source_stat.st_size, + progress_total=total_work, + ) + artifact = self._repository.commit(identity, build_directory) + except Exception: + self._repository.discard_build(build_directory) + LOGGER.exception( + "quantized checkpoint generation failed", + extra={ + "source_model": source_model, + "quantization_profile": profile.profile_id, + }, + ) + raise + + reporter.finish() + limit_bytes = self._limit_provider.limit_bytes() + eviction = self._cache_service.enforce_limit( + limit_bytes, + protected_paths=frozenset({artifact.path}), + ) + if eviction.remaining_bytes > limit_bytes: + LOGGER.warning( + "quant cache remains above configured limit because the current " + "or loaded artifacts are protected", + extra={ + "remaining_bytes": eviction.remaining_bytes, + "limit_bytes": limit_bytes, + }, + ) + output_size = getattr(result, "output_size_bytes", artifact.path.stat().st_size) + LOGGER.info( + "created %s cache for %s in %.2f seconds", + profile.label, + source_model, + time.monotonic() - started, + extra={ + "source_model": source_model, + "quantization_profile": profile.profile_id, + "duration_seconds": round(time.monotonic() - started, 2), + "artifact_size_bytes": output_size, + }, + ) + return self._cached_resolution(artifact, profile) + + def _cached_resolution( + self, + artifact: QuantCacheArtifact, + profile: QuantizationProfile, + ) -> ResolvedModelCheckpoint: + """Reserve a cache hit until its consumer establishes a model lease.""" + + return ResolvedModelCheckpoint( + artifact.path, + profile, + artifact, + self._leases.reserve(artifact.path), + ) + + def _require_disk_space(self, identity: QuantCacheIdentity) -> None: + """Fail before conversion when the cache drive lacks safe working space.""" + + ratio = ( + 0.55 + if any(item.value == "nvfp4" for item in identity.profile.required_formats) + else 0.75 + ) + required = int(identity.source.size_bytes * ratio) + MINIMUM_DISK_HEADROOM + self._repository.ensure_root() + available = shutil.disk_usage(self._repository.root).free + if available < required: + raise OSError( + f"Not enough free disk space to create the " + f"{identity.profile.label} cache. Approximately " + f"{required / 1024**3:.1f} GiB is required, but only " + f"{available / 1024**3:.1f} GiB is available on the cache drive." + ) + + +def _sha256_with_progress( + path: Path, + progress: QuantizationProgressReporter, + total_work: int, +) -> str: + """Hash a source file incrementally while updating node progress.""" + + digest = hashlib.sha256() + processed = 0 + with path.open("rb") as source: + while chunk := source.read(HASH_CHUNK_SIZE): + digest.update(chunk) + processed += len(chunk) + progress.advance(processed, total_work) + return digest.hexdigest() diff --git a/tests/test_anima_diffusion_model_service.py b/tests/test_anima_diffusion_model_service.py new file mode 100644 index 0000000..af0a423 --- /dev/null +++ b/tests/test_anima_diffusion_model_service.py @@ -0,0 +1,143 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for Anima's thin integration with reusable quantized model loading.""" + +from __future__ import annotations + +import gc +from dataclasses import dataclass +from pathlib import Path + +from simple_syrup.domain.anima_quantization import NVFP4_MIXED_PROFILE +from simple_syrup.domain.model_quantization import ( + ModelQuantizationRecipe, +) +from simple_syrup.domain.quant_cache import QuantCacheManifest +from simple_syrup.runtime.quant_cache_leases import QuantCacheLeaseRegistry +from simple_syrup.runtime.quant_cache_repository import QuantCacheArtifact +from simple_syrup.runtime.quantization_progress import QuantizationProgressReporter +from simple_syrup.services.anima_diffusion_model_service import ( + AnimaDiffusionModelService, +) +from simple_syrup.services.quantized_model_boundaries import ( + ResolvedModelCheckpoint, +) + + +class FakeUnderlyingModel: + """Weak-referenceable owner shared by a fake Comfy model patcher.""" + + +class FakeModelPatcher: + """Expose Comfy's underlying-model attribute used for cache leases.""" + + def __init__(self) -> None: + """Create one underlying model owner.""" + + self.model = FakeUnderlyingModel() + + +class FakeDiffusionLoader: + """Resolve the source and record actual checkpoint loading.""" + + def __init__(self, source_path: Path) -> None: + """Create a fake for one source path.""" + + self.source_path = source_path + self.load_calls: list[tuple[Path, str]] = [] + + def resolve_path(self, diffusion_model: str) -> Path: + """Return the configured authoritative source path.""" + + del diffusion_model + return self.source_path + + def load_path(self, model_path: Path, weight_dtype: str) -> object: + """Record the selected runtime path and dtype.""" + + self.load_calls.append((model_path, weight_dtype)) + return FakeModelPatcher() + + +@dataclass +class FakeResolver: + """Return a predetermined protected checkpoint resolution.""" + + resolved: ResolvedModelCheckpoint + requested_source_model: str | None = None + + def resolve( + self, + *, + source_model: str, + source_path: Path, + quantization: str, + recipe: ModelQuantizationRecipe, + progress: QuantizationProgressReporter | None = None, + ) -> ResolvedModelCheckpoint: + """Record source identity while returning the configured result.""" + + del source_path, quantization, recipe, progress + self.requested_source_model = source_model + return self.resolved + + +def test_anima_loads_cached_derivative_with_source_provenance_and_lease( + tmp_path: Path, +) -> None: + """Runtime cache paths never replace the workflow's authoritative model name.""" + + source_path = tmp_path / "models" / "diffusion_models" / "anima.safetensors" + artifact_path = tmp_path / "models" / "SyrupQuants" / "Anima" / "cached.safetensors" + artifact_path.parent.mkdir(parents=True) + artifact_path.write_bytes(b"quant") + manifest_path = artifact_path.with_name("manifest.json") + manifest_path.write_text("{}", encoding="utf-8") + manifest = QuantCacheManifest( + source_model="Anima/anima.safetensors", + source_path=str(source_path), + source_sha256="a" * 64, + source_size_bytes=10, + source_modified_ns=1, + profile_id="nvfp4-mixed", + profile_label="NVFP4 (Mixed)", + profile_version=3, + quantization_formats=("float8_e4m3fn", "nvfp4"), + model_family="Anima", + recipe_version=2, + artifact_file=artifact_path.name, + artifact_size_bytes=5, + created_at="2026-01-01T00:00:00+00:00", + last_used_at="2026-01-01T00:00:00+00:00", + ) + artifact = QuantCacheArtifact(artifact_path, manifest_path, manifest) + leases = QuantCacheLeaseRegistry() + resolver = FakeResolver( + ResolvedModelCheckpoint( + artifact_path, + NVFP4_MIXED_PROFILE, + artifact, + leases.reserve(artifact_path), + ) + ) + loader = FakeDiffusionLoader(source_path) + service = AnimaDiffusionModelService(loader, resolver, leases) + + loaded = service.load( + diffusion_model="Anima/anima.safetensors", + diffusion_weight_dtype="fp8_e4m3fn_fast", + quantization="nvfp4-mixed", + ) + + assert loader.load_calls == [(artifact_path, "default")] + assert resolver.requested_source_model == "Anima/anima.safetensors" + provenance = loaded.simple_syrup_model_provenance # type: ignore[attr-defined] + assert provenance["source_model"] == "Anima/anima.safetensors" + assert provenance["quantization_profile"] == "nvfp4-mixed" + assert provenance["quantization_profile_version"] == 3 + assert leases.is_leased(artifact_path) + del loaded + gc.collect() + assert not leases.is_leased(artifact_path) diff --git a/tests/test_anima_loader.py b/tests/test_anima_loader.py index 1d7ad09..06078c4 100644 --- a/tests/test_anima_loader.py +++ b/tests/test_anima_loader.py @@ -17,11 +17,7 @@ from types import ModuleType, TracebackType import pytest import torch -import simple_syrup.runtime.anima_loader as anima_loader_module -from simple_syrup.runtime.anima_loader import ( - AUTO_CHOICE, - AnimaLoaderService, -) +import simple_syrup.services.anima_loader_service as anima_loader_module from simple_syrup.runtime.auto_model_artifact import AutoModelArtifact from simple_syrup.runtime.auto_model_cache import AutoModelCache from simple_syrup.runtime.auto_model_resolver import ( @@ -34,6 +30,15 @@ from simple_syrup.runtime.model_downloads import ( ProgressReporter, ) from simple_syrup.runtime.vae_loader import vae_choices +from simple_syrup.services.anima_loader_service import ( + AUTO_CHOICE, + AnimaLoaderService, +) + + +@dataclass +class FakeLoadedModel: + """Attribute-bearing stand-in for a ComfyUI ModelPatcher.""" @dataclass @@ -45,6 +50,7 @@ class FakeComfyState: vae_paths: list[str] = field(default_factory=list) progress_totals: list[int] = field(default_factory=list) progress_updates: list[list[tuple[int, int | None]]] = field(default_factory=list) + model: object = field(default_factory=FakeLoadedModel) class FakeStreamingResponse: @@ -154,6 +160,7 @@ def test_loader_maps_diffusion_weight_dtype( service.load_models( "anima.safetensors", + "Original", "fp8_e4m3fn_fast", "manual_clip.safetensors", "default", @@ -183,6 +190,7 @@ def test_loader_maps_clip_cpu_device( service.load_models( "anima.safetensors", + "Original", "default", "manual_clip.safetensors", "cpu", @@ -216,6 +224,7 @@ def test_loader_uses_auto_resolver_for_auto_choices( progress = RecordingProgress() service.load_models( "anima.safetensors", + "Original", "default", AUTO_CHOICE, "default", @@ -286,6 +295,7 @@ def test_anima_auto_downloads_emit_comfy_node_progress_end_to_end( service.load_models( "anima.safetensors", + "Original", "default", AUTO_CHOICE, "default", @@ -339,13 +349,14 @@ def test_loader_returns_model_clip_and_vae( result = service.load_models( "anima.safetensors", + "Original", "default", "manual_clip.safetensors", "default", "manual_vae.safetensors", ) - assert result[0] == "model" + assert result[0] is comfy_state.model assert result[1] == "clip" assert result[2] is not None assert comfy_state.vae_paths == [ @@ -403,11 +414,11 @@ def _install_fake_comfy(monkeypatch: pytest.MonkeyPatch) -> FakeComfyState: def load_diffusion_model( path: str, model_options: dict[str, object], - ) -> str: + ) -> object: """Record diffusion model calls.""" state.diffusion_calls.append((path, model_options)) - return "model" + return state.model def load_clip( ckpt_paths: list[str], diff --git a/tests/test_anima_quantization_workflow.py b/tests/test_anima_quantization_workflow.py new file mode 100644 index 0000000..c6b31db --- /dev/null +++ b/tests/test_anima_quantization_workflow.py @@ -0,0 +1,271 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Integration proof for Simple Load Anima on-demand quantization and reuse.""" + +from __future__ import annotations + +import gc +from pathlib import Path +from types import ModuleType +from typing import Any, cast + +import pytest +import torch +from safetensors import safe_open +from safetensors.torch import save_file + +from simple_syrup.domain.model_quantization import ( + ModelQuantizationRecipe, + QuantizationProfile, +) +from simple_syrup.domain.quant_cache import SourceCheckpointIdentity +from simple_syrup.nodes.simple_load_anima import SimpleLoadAnima +from simple_syrup.runtime.checkpoint_quantizer import ( + CheckpointQuantizationResult, + SafetensorsCheckpointQuantizer, +) +from simple_syrup.runtime.diffusion_model_loader import DiffusionModelLoader +from simple_syrup.runtime.quant_cache_leases import QuantCacheLeaseRegistry +from simple_syrup.runtime.quant_cache_repository import QuantCacheRepository +from simple_syrup.runtime.quantization_capabilities import ( + QuantizationCapabilityCatalog, +) +from simple_syrup.runtime.quantization_progress import QuantizationProgressReporter +from simple_syrup.services.anima_diffusion_model_service import ( + AnimaDiffusionModelService, +) +from simple_syrup.services.anima_loader_service import AnimaLoaderService +from simple_syrup.services.quantized_model_resolver import QuantizedModelResolver + + +class WorkflowFolderPaths(ModuleType): + """Resolve a complete tiny workflow fixture below one models directory.""" + + def __init__(self, models_dir: Path) -> None: + """Create folder mappings for diffusion, text encoder, VAE, and embeddings.""" + + super().__init__("folder_paths") + self.models_dir = str(models_dir) + + def get_full_path_or_raise(self, folder_name: str, filename: str) -> str: + """Return one conventional model path.""" + + return str(Path(self.models_dir) / folder_name / filename) + + def get_folder_paths(self, folder_name: str) -> list[str]: + """Return one conventional model directory.""" + + return [str(Path(self.models_dir) / folder_name)] + + +class UnusedAutoResolver: + """Fail if manual workflow choices unexpectedly trigger auto resolution.""" + + def resolve(self, artifact: object, progress: object = None) -> object: + """Reject unexpected automatic resolution.""" + + del artifact, progress + raise AssertionError("Manual workflow fixtures must not auto-resolve models.") + + +class FakeUnderlyingModel: + """Weak-referenceable owner for cache lease verification.""" + + +class FakeModelPatcher: + """Represent the ComfyUI model returned from a resolved checkpoint path.""" + + def __init__(self, loaded_path: str) -> None: + """Retain the runtime path while exposing an underlying model owner.""" + + self.loaded_path = loaded_path + self.model = FakeUnderlyingModel() + + +class FakeVAE: + """Accept the same construction and validation calls as ComfyUI's VAE.""" + + def __init__(self, sd: object, metadata: object = None) -> None: + """Accept decoded VAE state.""" + + del sd, metadata + + def throw_exception_if_invalid(self) -> None: + """Accept the tiny workflow fixture.""" + + +class FixedLimitProvider: + """Keep the tiny integration cache beneath a generous test limit.""" + + def limit_bytes(self) -> int: + """Return one GiB.""" + + return 1024**3 + + +class CountingQuantizer: + """Count calls while delegating real checkpoint conversion to ComfyUI.""" + + def __init__(self) -> None: + """Create the real quantizer and an empty call counter.""" + + self.calls = 0 + self._quantizer = SafetensorsCheckpointQuantizer() + + def quantize( + self, + *, + source: SourceCheckpointIdentity, + destination_path: Path, + profile: QuantizationProfile, + recipe: ModelQuantizationRecipe, + progress: QuantizationProgressReporter, + progress_base: int, + progress_total: int, + ) -> CheckpointQuantizationResult: + """Delegate one real conversion and increment the call count.""" + + self.calls += 1 + return self._quantizer.quantize( + source=source, + destination_path=destination_path, + profile=profile, + recipe=recipe, + progress=progress, + progress_base=progress_base, + progress_total=progress_total, + ) + + +def test_simple_load_anima_generates_and_reuses_native_nvfp4_cache( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A node execution generates once, reports progress, and then loads its cache.""" + + import comfy.model_management as model_management + import comfy.sd as comfy_sd + import comfy.utils as comfy_utils + + device = model_management.get_torch_device() + if not model_management.supports_nvfp4_compute(device): + pytest.skip("NVFP4 compute is unavailable on this GPU.") + + models_dir = tmp_path / "models" + diffusion_path = models_dir / "diffusion_models" / "Anima" / "anima.safetensors" + diffusion_path.parent.mkdir(parents=True) + save_file( + { + "net.blocks.2.attn.q_proj.weight": torch.randn( + 16, 16, dtype=torch.bfloat16 + ), + "net.blocks.0.attn.q_proj.weight": torch.randn( + 16, 16, dtype=torch.bfloat16 + ), + }, + str(diffusion_path), + ) + folder_paths = WorkflowFolderPaths(models_dir) + loaded_paths: list[str] = [] + progress_updates: list[tuple[int, int | None]] = [] + + def load_diffusion_model( + path: str, model_options: dict[str, object] + ) -> FakeModelPatcher: + """Record the derived checkpoint loaded by the node.""" + + assert model_options == {} + loaded_paths.append(path) + return FakeModelPatcher(path) + + class FakeProgressBar: + """Record node progress published during the first conversion.""" + + def __init__(self, total: int) -> None: + """Retain the expected total through later updates.""" + + self.total = total + + def update_absolute(self, value: int, total: int | None = None) -> None: + """Record one absolute progress update.""" + + progress_updates.append((value, total)) + + monkeypatch.setattr(comfy_sd, "load_diffusion_model", load_diffusion_model) + monkeypatch.setattr(comfy_sd, "load_clip", lambda **kwargs: "clip") + monkeypatch.setattr(comfy_sd, "VAE", FakeVAE) + monkeypatch.setattr( + comfy_utils, + "load_torch_file", + lambda path, return_metadata=False: ({}, {}) if return_metadata else {}, + ) + monkeypatch.setattr(comfy_utils, "ProgressBar", FakeProgressBar) + + cache_repository = QuantCacheRepository(models_dir / "SyrupQuants") + leases = QuantCacheLeaseRegistry() + quantizer = CountingQuantizer() + resolver = QuantizedModelResolver( + repository=cache_repository, + quantizer=quantizer, + capabilities=QuantizationCapabilityCatalog(), + limit_provider=FixedLimitProvider(), + leases=leases, + ) + diffusion_service = AnimaDiffusionModelService( + DiffusionModelLoader(folder_paths), + resolver, + leases, + ) + service = AnimaLoaderService( + resolver=UnusedAutoResolver(), # type: ignore[arg-type] + diffusion_service=diffusion_service, + folder_paths_module=folder_paths, + ) + original_service = SimpleLoadAnima._service + SimpleLoadAnima._service = service + try: + first = SimpleLoadAnima().load_models( + diffusion_model="Anima/anima.safetensors", + quantization="nvfp4-mixed", + diffusion_weight_dtype="fp8_e4m3fn_fast", + text_encoder="manual_clip.safetensors", + text_encoder_device="default", + vae="manual_vae.safetensors", + ) + first_progress_count = len(progress_updates) + second = SimpleLoadAnima().load_models( + diffusion_model="Anima/anima.safetensors", + quantization="nvfp4-mixed", + diffusion_weight_dtype="default", + text_encoder="manual_clip.safetensors", + text_encoder_device="default", + vae="manual_vae.safetensors", + ) + finally: + SimpleLoadAnima._service = original_service + + assert quantizer.calls == 1 + assert len(loaded_paths) == 2 + assert loaded_paths[0] == loaded_paths[1] + derived_path = Path(loaded_paths[0]) + assert derived_path.is_relative_to(models_dir / "SyrupQuants") + assert first_progress_count > 1 + assert len(progress_updates) == first_progress_count + assert first[1] == "clip" + assert isinstance(first[2], FakeVAE) + assert second[1] == "clip" + assert isinstance(second[2], FakeVAE) + first_model = cast(Any, first[0]) + assert first_model.simple_syrup_model_provenance["source_model"] == ( + "Anima/anima.safetensors" + ) + with safe_open(str(derived_path), framework="pt", device="cpu") as checkpoint: + assert checkpoint.metadata()["simple_syrup.source_model"] == ( + "Anima/anima.safetensors" + ) + assert "net.blocks.2.attn.q_proj.comfy_quant" in checkpoint.keys() + assert (models_dir / "SyrupQuants" / "README.txt").is_file() + del first, second + gc.collect() diff --git a/tests/test_checkpoint_quantizer.py b/tests/test_checkpoint_quantizer.py new file mode 100644 index 0000000..342b682 --- /dev/null +++ b/tests/test_checkpoint_quantizer.py @@ -0,0 +1,431 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for streaming ComfyUI checkpoint quantization.""" + +from __future__ import annotations + +import hashlib +import importlib +import sys +from dataclasses import dataclass, field +from pathlib import Path +from types import ModuleType +from typing import Any + +import pytest +import torch +from safetensors import safe_open +from safetensors.torch import load_file, save_file + +from simple_syrup.domain.anima_quantization import ( + MXFP8_PROFILE, + NVFP4_MIXED_PROFILE, + AnimaQuantizationRecipe, +) +from simple_syrup.domain.model_quantization import QuantizationFormat +from simple_syrup.domain.quant_cache import SourceCheckpointIdentity +from simple_syrup.runtime.checkpoint_quantizer import SafetensorsCheckpointQuantizer +from simple_syrup.runtime.comfy_safetensors_dtypes import ( + ComfySafetensorsDtypeRegistry, +) + + +@dataclass +class RecordingProgress: + """Record absolute quantization progress updates.""" + + updates: list[tuple[int, int]] = field(default_factory=list) + + def start(self, label: str, total: int) -> None: + """Record no start because the resolver owns operation setup.""" + + del label, total + + def advance(self, current: int, total: int) -> None: + """Record one absolute update.""" + + self.updates.append((current, total)) + + def finish(self) -> None: + """Record no finish because the resolver owns completion.""" + + +class FakeQuantizedTensor: + """Serialize deterministic fake quant tensors through Comfy's boundary.""" + + def __init__(self, source: torch.Tensor) -> None: + """Retain only the source shape needed by the fake serializer.""" + + self._shape = source.shape + + @classmethod + def from_float( + cls, + tensor: torch.Tensor, + layout_name: str, + **kwargs: object, + ) -> FakeQuantizedTensor: + """Accept the same call shape as ComfyUI's QuantizedTensor.""" + + assert layout_name == "FakeLayout" + assert kwargs == {"scale": "recalculate"} + return cls(tensor) + + def state_dict(self, prefix: str) -> dict[str, torch.Tensor]: + """Return representative quantized weight and scale tensors.""" + + return { + prefix: torch.zeros(self._shape, dtype=torch.uint8), + f"{prefix}_scale": torch.ones((), dtype=torch.float32), + } + + +class RecordingDtypeRegistry(ComfySafetensorsDtypeRegistry): + """Record header and marker validation performed before publication.""" + + def __init__(self, *, supports_markers: bool = True) -> None: + """Create an empty validation record.""" + + self._supports_markers = supports_markers + self.validated_paths: list[Path] = [] + self.validated_formats: list[QuantizationFormat] = [] + + def validate_checkpoint_header(self, checkpoint_path: Path) -> None: + """Record the generated checkpoint whose header was validated.""" + + self.validated_paths.append(checkpoint_path) + + def supports_format(self, quantization_format: QuantizationFormat) -> bool: + """Record and accept one generated Comfy marker format.""" + + self.validated_formats.append(quantization_format) + return self._supports_markers + + +def test_quantizer_streams_eligible_tensors_and_preserves_anima_policy( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Only eligible matrices receive Comfy layer markers and scale tensors.""" + + source_path = tmp_path / "anima.safetensors" + destination = tmp_path / "quantized.safetensors" + save_file( + { + "net.blocks.0.attn.q_proj.weight": torch.randn(4, 4, dtype=torch.bfloat16), + "net.blocks.2.attn.q_proj.weight": torch.randn(4, 4, dtype=torch.bfloat16), + "net.blocks.2.attn.q_proj.bias": torch.randn(4, dtype=torch.bfloat16), + }, + str(source_path), + metadata={"source": "test"}, + ) + quant_ops = ModuleType("comfy.quant_ops") + quant_ops.QUANT_ALGOS = { # type: ignore[attr-defined] + "float8_e4m3fn": {"comfy_tensor_layout": "FakeLayout"}, + "nvfp4": {"comfy_tensor_layout": "FakeLayout"}, + } + quant_ops.QuantizedTensor = FakeQuantizedTensor # type: ignore[attr-defined] + model_management = ModuleType("comfy.model_management") + model_management.get_torch_device = lambda: torch.device("cpu") # type: ignore[attr-defined] + monkeypatch.setitem(sys.modules, "comfy.quant_ops", quant_ops) + monkeypatch.setitem(sys.modules, "comfy.model_management", model_management) + progress = RecordingProgress() + source = _identity(source_path) + dtype_registry = RecordingDtypeRegistry() + + result = SafetensorsCheckpointQuantizer(dtype_registry).quantize( + source=source, + destination_path=destination, + profile=NVFP4_MIXED_PROFILE, + recipe=AnimaQuantizationRecipe(), + progress=progress, + progress_base=source.size_bytes, + progress_total=source.size_bytes * 2, + ) + + assert result.quantized_tensor_count == 1 + assert result.preserved_tensor_count == 2 + assert progress.updates + assert dtype_registry.validated_paths == [destination] + assert dtype_registry.validated_formats == [QuantizationFormat.NVFP4] + with safe_open(str(destination), framework="pt", device="cpu") as checkpoint: + assert set(checkpoint.keys()) == { + "net.blocks.0.attn.q_proj.weight", + "net.blocks.2.attn.q_proj.bias", + "net.blocks.2.attn.q_proj.comfy_quant", + "net.blocks.2.attn.q_proj.weight", + "net.blocks.2.attn.q_proj.weight_scale", + } + assert ( + checkpoint.get_tensor("net.blocks.0.attn.q_proj.weight").dtype + is torch.bfloat16 + ) + assert checkpoint.metadata() == { + "source": "test", + "simple_syrup.derived_model": "true", + "simple_syrup.model_family": "Anima", + "simple_syrup.profile_label": "NVFP4 (Mixed)", + "simple_syrup.profile_version": "3", + "simple_syrup.quantization_formats": "float8_e4m3fn,nvfp4", + "simple_syrup.quantization_profile": "nvfp4-mixed", + "simple_syrup.recipe_version": "2", + "simple_syrup.source_model": "Anima/anima.safetensors", + "simple_syrup.source_sha256": source.sha256, + } + + +def test_quantizer_rejects_a_quantized_source( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Derived checkpoints cannot silently become lossy quantization sources.""" + + source_path = tmp_path / "already-quantized.safetensors" + save_file( + { + "net.blocks.2.linear.weight": torch.ones(4, 4), + "net.blocks.2.linear.comfy_quant": torch.tensor( + list(b'{"format":"nvfp4"}'), dtype=torch.uint8 + ), + }, + str(source_path), + ) + quant_ops = ModuleType("comfy.quant_ops") + quant_ops.QUANT_ALGOS = { # type: ignore[attr-defined] + "float8_e4m3fn": {"comfy_tensor_layout": "FakeLayout"}, + "nvfp4": {"comfy_tensor_layout": "FakeLayout"}, + } + quant_ops.QuantizedTensor = FakeQuantizedTensor # type: ignore[attr-defined] + model_management = ModuleType("comfy.model_management") + model_management.get_torch_device = lambda: torch.device("cpu") # type: ignore[attr-defined] + monkeypatch.setitem(sys.modules, "comfy.quant_ops", quant_ops) + monkeypatch.setitem(sys.modules, "comfy.model_management", model_management) + + with pytest.raises(ValueError, match="already contains ComfyUI quantization"): + SafetensorsCheckpointQuantizer().quantize( + source=_identity(source_path), + destination_path=tmp_path / "output.safetensors", + profile=NVFP4_MIXED_PROFILE, + recipe=AnimaQuantizationRecipe(), + progress=RecordingProgress(), + progress_base=0, + progress_total=1, + ) + + +def test_validation_rejects_inconsistent_profile_metadata(tmp_path: Path) -> None: + """A generated artifact cannot publish under mismatched profile metadata.""" + + destination = tmp_path / "mismatched-profile.safetensors" + save_file( + { + "net.blocks.2.attn.q_proj.weight": torch.zeros(4, 4, dtype=torch.uint8), + "net.blocks.2.attn.q_proj.comfy_quant": torch.tensor( + list(b'{"format":"nvfp4"}'), dtype=torch.uint8 + ), + }, + str(destination), + metadata={ + "simple_syrup.quantization_profile": "wrong-profile", + "simple_syrup.profile_version": "3", + "simple_syrup.quantization_formats": "float8_e4m3fn,nvfp4", + }, + ) + + with pytest.raises(ValueError, match="inconsistent profile metadata"): + SafetensorsCheckpointQuantizer(RecordingDtypeRegistry())._validate_output( + destination, + NVFP4_MIXED_PROFILE, + 1, + ) + + +def test_validation_rejects_marker_unsupported_by_active_loader( + tmp_path: Path, +) -> None: + """Publication fails when the active loader cannot read a marker format.""" + + destination = tmp_path / "unsupported-marker.safetensors" + save_file( + { + "net.blocks.2.attn.q_proj.weight": torch.zeros(4, 4, dtype=torch.uint8), + "net.blocks.2.attn.q_proj.comfy_quant": torch.tensor( + list(b'{"format":"nvfp4"}'), dtype=torch.uint8 + ), + }, + str(destination), + metadata={ + "simple_syrup.quantization_profile": "nvfp4-mixed", + "simple_syrup.profile_version": "3", + "simple_syrup.quantization_formats": "float8_e4m3fn,nvfp4", + }, + ) + + with pytest.raises(ValueError, match="unsupported by the active Comfy loader"): + SafetensorsCheckpointQuantizer( + RecordingDtypeRegistry(supports_markers=False) + )._validate_output( + destination, + NVFP4_MIXED_PROFILE, + 1, + ) + + +def test_native_nvfp4_generation_uses_comfy_scale_layout(tmp_path: Path) -> None: + """The installed ComfyUI runtime produces its native NVFP4 checkpoint shape.""" + + try: + import comfy.model_management as model_management + except ImportError: + pytest.skip("ComfyUI runtime is unavailable.") + if not model_management.supports_nvfp4_compute(model_management.get_torch_device()): + pytest.skip("NVFP4 compute is unavailable on this GPU.") + + source_path = tmp_path / "native-anima.safetensors" + destination = tmp_path / "native-nvfp4.safetensors" + save_file( + {"net.blocks.2.attn.q_proj.weight": torch.randn(16, 16, dtype=torch.bfloat16)}, + str(source_path), + ) + source = _identity(source_path) + + SafetensorsCheckpointQuantizer().quantize( + source=source, + destination_path=destination, + profile=NVFP4_MIXED_PROFILE, + recipe=AnimaQuantizationRecipe(), + progress=RecordingProgress(), + progress_base=source.size_bytes, + progress_total=source.size_bytes * 2, + ) + + with safe_open(str(destination), framework="pt", device="cpu") as checkpoint: + keys = set(checkpoint.keys()) + assert "net.blocks.2.attn.q_proj.comfy_quant" in keys + assert "net.blocks.2.attn.q_proj.weight_scale" in keys + assert "net.blocks.2.attn.q_proj.weight_scale_2" in keys + assert ( + checkpoint.get_tensor("net.blocks.2.attn.q_proj.weight").dtype + is torch.uint8 + ) + + import comfy.ops + import comfy.quant_ops + + state_dict = { + key.removeprefix("net.blocks.2.attn.q_proj."): value + for key, value in load_file(str(destination)).items() + } + operations = comfy.ops.mixed_precision_ops( + {"mixed_ops": True}, compute_dtype=torch.bfloat16 + ) + layer = operations.Linear( + 16, + 16, + bias=False, + device=torch.device("cpu"), + dtype=torch.bfloat16, + ) + + layer.load_state_dict(state_dict, strict=True) + + assert isinstance(layer.weight, comfy.quant_ops.QuantizedTensor) + assert layer.quant_format == "nvfp4" + + +def test_native_mxfp8_generation_reloads_through_aimdo(tmp_path: Path) -> None: + """E8M0 scale tensors reload through Comfy's active AIMDO mmap path.""" + + import comfy.model_management as model_management + import comfy.utils as comfy_utils + + aimdo_control: Any = importlib.import_module("comfy_aimdo.control") + aimdo_model_mmap: Any = importlib.import_module("comfy_aimdo.model_mmap") + + device = model_management.get_torch_device() + if not model_management.supports_mxfp8_compute(device): + pytest.skip("MXFP8 compute is unavailable on this GPU.") + source_path = tmp_path / "native-mxfp8-source.safetensors" + destination = tmp_path / "native-mxfp8.safetensors" + save_file( + {"net.blocks.14.attn.q_proj.weight": torch.randn(32, 32, dtype=torch.bfloat16)}, + str(source_path), + ) + source = _identity(source_path) + + SafetensorsCheckpointQuantizer().quantize( + source=source, + destination_path=destination, + profile=MXFP8_PROFILE, + recipe=AnimaQuantizationRecipe(), + progress=RecordingProgress(), + progress_base=source.size_bytes, + progress_total=source.size_bytes * 2, + ) + if not aimdo_control.init(): + pytest.skip("comfy-aimdo is unavailable.") + importlib.reload(aimdo_model_mmap) + state_dict, metadata = comfy_utils.load_safetensors(str(destination)) + + assert metadata["simple_syrup.quantization_profile"] == "mxfp8" + assert any(tensor.dtype is torch.float8_e8m0fnu for tensor in state_dict.values()) + + +def test_native_recommended_profile_serializes_both_formats(tmp_path: Path) -> None: + """Recommended profile emits native NVFP4 and FP8 markers in one checkpoint.""" + + import comfy.model_management as model_management + + device = model_management.get_torch_device() + if not model_management.supports_nvfp4_compute(device): + pytest.skip("NVFP4 compute is unavailable on this GPU.") + source_path = tmp_path / "native-mixed-source.safetensors" + destination = tmp_path / "native-mixed.safetensors" + save_file( + { + "net.blocks.14.attn.q_proj.weight": torch.randn( + 32, 32, dtype=torch.bfloat16 + ), + "net.blocks.14.attn.v_proj.weight": torch.randn( + 32, 32, dtype=torch.bfloat16 + ), + }, + str(source_path), + ) + source = _identity(source_path) + + SafetensorsCheckpointQuantizer().quantize( + source=source, + destination_path=destination, + profile=NVFP4_MIXED_PROFILE, + recipe=AnimaQuantizationRecipe(), + progress=RecordingProgress(), + progress_base=source.size_bytes, + progress_total=source.size_bytes * 2, + ) + + with safe_open(str(destination), framework="pt", device="cpu") as checkpoint: + q_marker = bytes( + checkpoint.get_tensor("net.blocks.14.attn.q_proj.comfy_quant").tolist() + ) + v_marker = bytes( + checkpoint.get_tensor("net.blocks.14.attn.v_proj.comfy_quant").tolist() + ) + assert b'"format":"nvfp4"' in q_marker + assert b'"format":"float8_e4m3fn"' in v_marker + + +def _identity(path: Path) -> SourceCheckpointIdentity: + """Create a complete source identity for a tiny test checkpoint.""" + + content = path.read_bytes() + stat = path.stat() + return SourceCheckpointIdentity( + display_name="Anima/anima.safetensors", + path=path.resolve(), + size_bytes=stat.st_size, + modified_ns=stat.st_mtime_ns, + sha256=hashlib.sha256(content).hexdigest(), + ) diff --git a/tests/test_comfy_safetensors_dtypes.py b/tests/test_comfy_safetensors_dtypes.py new file mode 100644 index 0000000..119906c --- /dev/null +++ b/tests/test_comfy_safetensors_dtypes.py @@ -0,0 +1,60 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for guarded Comfy safetensors dtype compatibility.""" + +from __future__ import annotations + +from types import ModuleType + +import torch + +from simple_syrup.domain.model_quantization import QuantizationFormat +from simple_syrup.runtime.comfy_safetensors_dtypes import ( + ComfySafetensorsDtypeRegistry, +) + + +def _comfy_types(types: dict[str, object]) -> ModuleType: + """Create a fake Comfy utils module with a private dtype registry.""" + + module = ModuleType("comfy.utils") + module._TYPES = types # type: ignore[attr-defined] + return module + + +def test_registry_adds_e8m0_when_host_supports_it() -> None: + """SimpleSyrup fills the missing mapping without replacing host entries.""" + + types: dict[str, object] = {"F8_E4M3": torch.float8_e4m3fn} + registry = ComfySafetensorsDtypeRegistry(_comfy_types(types)) + + assert registry.register_extension_dtypes() + assert types["F8_E8M0"] == torch.float8_e8m0fnu + assert registry.supports_format(QuantizationFormat.MXFP8) + + +def test_registry_accepts_identical_existing_mapping() -> None: + """A host-provided identical mapping is an idempotent success.""" + + types: dict[str, object] = { + "F8_E4M3": torch.float8_e4m3fn, + "F8_E8M0": torch.float8_e8m0fnu, + } + + assert ComfySafetensorsDtypeRegistry( + _comfy_types(types) + ).register_extension_dtypes() + + +def test_registry_rejects_conflict_without_overwriting_it() -> None: + """A conflicting host mapping fails closed and remains untouched.""" + + conflict = object() + types = {"F8_E4M3": torch.float8_e4m3fn, "F8_E8M0": conflict} + registry = ComfySafetensorsDtypeRegistry(_comfy_types(types)) + + assert not registry.register_extension_dtypes() + assert not registry.supports_format(QuantizationFormat.MXFP8) + assert types["F8_E8M0"] is conflict diff --git a/tests/test_model_quantization.py b/tests/test_model_quantization.py new file mode 100644 index 0000000..cfde3ce --- /dev/null +++ b/tests/test_model_quantization.py @@ -0,0 +1,107 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for reusable and Anima-specific quantization profiles.""" + +from __future__ import annotations + +import pytest + +from simple_syrup.domain.anima_quantization import ( + FP8_E4M3_PROFILE, + FP8_E5M2_PROFILE, + MXFP8_PROFILE, + NVFP4_MIXED_PROFILE, + ORIGINAL_PROFILE, + AnimaQuantizationRecipe, +) +from simple_syrup.domain.model_quantization import ( + QuantizationFormat, + QuantizationProfile, + TensorDescriptor, +) + + +def test_anima_profiles_parse_labels_and_stable_ids() -> None: + """Workflow labels and persisted identifiers resolve consistently.""" + + recipe = AnimaQuantizationRecipe() + + assert recipe.profile_from_selection("Original") == ORIGINAL_PROFILE + assert recipe.profile_from_selection("nvfp4-mixed") == NVFP4_MIXED_PROFILE + assert QuantizationFormat.NVFP4.label == "NVFP4" + with pytest.raises(ValueError, match="quantization profile must be one of"): + recipe.profile_from_selection("unknown") + + +@pytest.mark.parametrize( + "name", + [ + "net.blocks.0.attn.q_proj.weight", + "net.blocks.1.attn.q_proj.weight", + "net.blocks.27.attn.q_proj.weight", + "net.blocks.14.adaln_modulation.1.weight", + "net.final_layer.linear.weight", + "net.llm_adapter.proj.weight", + "net.t_embedder.1.weight", + "net.x_embedder.proj.weight", + "some_other_model.blocks.14.attn.q_proj.weight", + ], +) +@pytest.mark.parametrize( + "profile", + [FP8_E4M3_PROFILE, MXFP8_PROFILE, NVFP4_MIXED_PROFILE], +) +def test_anima_profiles_preserve_quality_sensitive_weights( + name: str, + profile: QuantizationProfile, +) -> None: + """Every profile keeps known sensitive and out-of-envelope weights.""" + + recipe = AnimaQuantizationRecipe() + selected = recipe.profile_from_selection(profile.profile_id) + + assert recipe.policy_for(TensorDescriptor(name, (16, 16), "BF16"), selected) is None + + +def test_anima_mixed_profile_assigns_projection_specific_formats() -> None: + """Recommended mixed precision follows the intended attention/MLP split.""" + + recipe = AnimaQuantizationRecipe() + + def assigned(name: str) -> QuantizationFormat | None: + return recipe.policy_for( + TensorDescriptor(f"net.blocks.14.{name}.weight", (16, 16), "BF16"), + NVFP4_MIXED_PROFILE, + ) + + assert assigned("attn.q_proj") is QuantizationFormat.NVFP4 + assert assigned("attn.k_proj") is QuantizationFormat.NVFP4 + assert assigned("attn.output_proj") is QuantizationFormat.NVFP4 + assert assigned("attn.v_proj") is QuantizationFormat.FP8_E4M3 + assert assigned("mlp.fc1") is QuantizationFormat.FP8_E4M3 + assert assigned("unmatched") is None + + +@pytest.mark.parametrize( + ("profile", "expected"), + [ + (FP8_E4M3_PROFILE, QuantizationFormat.FP8_E4M3), + (FP8_E5M2_PROFILE, QuantizationFormat.FP8_E5M2), + (MXFP8_PROFILE, QuantizationFormat.MXFP8), + ], +) +def test_uniform_profiles_quantize_only_eligible_middle_block_matrices( + profile: QuantizationProfile, + expected: QuantizationFormat, +) -> None: + """Uniform profiles share the safety envelope while selecting their format.""" + + descriptor = TensorDescriptor( + "diffusion_model.blocks.14.self_attn.q_proj.weight", + (16, 16), + "BF16", + ) + + assert AnimaQuantizationRecipe().policy_for(descriptor, profile) is expected diff --git a/tests/test_persisted_widget_order_contract.py b/tests/test_persisted_widget_order_contract.py index 05a52c4..fd7e6a8 100644 --- a/tests/test_persisted_widget_order_contract.py +++ b/tests/test_persisted_widget_order_contract.py @@ -219,6 +219,7 @@ _PERSISTED_WIDGET_PREFIXES: tuple[tuple[type[_ClassicNode], tuple[str, ...]], .. SimpleLoadAnima, ( "diffusion_model", + "quantization", "diffusion_weight_dtype", "text_encoder", "text_encoder_device", diff --git a/tests/test_quant_cache.py b/tests/test_quant_cache.py new file mode 100644 index 0000000..5db348b --- /dev/null +++ b/tests/test_quant_cache.py @@ -0,0 +1,247 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for readable quant cache persistence, leases, and global LRU policy.""" + +from __future__ import annotations + +import hashlib +import json +from dataclasses import replace +from pathlib import Path + +from simple_syrup.domain.anima_quantization import ( + MXFP8_PROFILE, + NVFP4_MIXED_PROFILE, + AnimaQuantizationRecipe, +) +from simple_syrup.domain.quant_cache import ( + QuantCacheIdentity, + SourceCheckpointIdentity, +) +from simple_syrup.runtime.quant_cache_leases import QuantCacheLeaseRegistry +from simple_syrup.runtime.quant_cache_repository import ( + MANIFEST_FILENAME, + README_FILENAME, + QuantCacheArtifact, + QuantCacheRepository, +) +from simple_syrup.services.quant_cache_service import QuantCacheService + + +def test_repository_creates_a_readable_global_cache_layout(tmp_path: Path) -> None: + """Users browsing models/SyrupQuants can identify ownership and artifacts.""" + + repository = QuantCacheRepository(tmp_path / "models" / "SyrupQuants") + identity = _identity(tmp_path, "Anima/base model.safetensors", b"source-a") + + artifact = _commit(repository, identity, b"quantized") + + assert ( + (repository.root / README_FILENAME) + .read_text(encoding="utf-8") + .startswith("SimpleSyrup Quantized Model Cache") + ) + assert artifact.path.relative_to(repository.root).as_posix() == ( + "Anima/base_model/nvfp4-mixed/" + f"{identity.source.sha256[:12]}-profile-3-recipe-2/" + "base_model--nvfp4-mixed.safetensors" + ) + payload = json.loads(artifact.manifest_path.read_text(encoding="utf-8")) + assert payload["source_model"] == "Anima/base model.safetensors" + assert payload["profile_id"] == "nvfp4-mixed" + assert payload["profile_version"] == 3 + assert payload["quantization_formats"] == ["float8_e4m3fn", "nvfp4"] + assert repository.relative_display_path() == "models/SyrupQuants" + + +def test_repository_finds_only_an_unchanged_source_signature(tmp_path: Path) -> None: + """Fast cache hits invalidate when the authoritative source file changes.""" + + repository = QuantCacheRepository(tmp_path / "SyrupQuants") + identity = _identity(tmp_path, "anima.safetensors", b"source") + artifact = _commit(repository, identity, b"quant") + + found = repository.find_current( + source_model=identity.source.display_name, + source_path=identity.source.path, + source_size_bytes=identity.source.size_bytes, + source_modified_ns=identity.source.modified_ns, + profile=identity.profile, + recipe=AnimaQuantizationRecipe(), + ) + changed = repository.find_current( + source_model=identity.source.display_name, + source_path=identity.source.path, + source_size_bytes=identity.source.size_bytes + 1, + source_modified_ns=identity.source.modified_ns, + profile=identity.profile, + recipe=AnimaQuantizationRecipe(), + ) + + assert found is not None + assert found.path == artifact.path + assert changed is None + + +def test_profiles_have_isolated_readable_cache_identities(tmp_path: Path) -> None: + """Mixed and MXFP8 recipes never share a path or cache artifact.""" + + repository = QuantCacheRepository(tmp_path / "SyrupQuants") + mixed = _identity(tmp_path, "anima.safetensors", b"source") + mxfp8 = replace(mixed, profile=MXFP8_PROFILE) + + assert mixed.stable_key != mxfp8.stable_key + assert repository.artifact_directory(mixed) != repository.artifact_directory(mxfp8) + + +def test_global_lru_preserves_reserved_artifacts_and_clears_inactive( + tmp_path: Path, +) -> None: + """Cache clearing never deletes an artifact reserved for active model work.""" + + repository = QuantCacheRepository(tmp_path / "SyrupQuants") + old = _commit( + repository, + _identity(tmp_path, "old.safetensors", b"old"), + b"old-quant", + ) + new = _commit( + repository, + _identity(tmp_path, "new.safetensors", b"new"), + b"new-quantized", + ) + _set_last_used(old, "2026-01-01T00:00:00+00:00") + _set_last_used(new, "2026-02-01T00:00:00+00:00") + leases = QuantCacheLeaseRegistry() + reservation = leases.reserve(old.path) + service = QuantCacheService(repository, leases) + + result = service.clear_inactive() + + assert result.removed_artifacts == 1 + assert old.path.is_file() + assert not new.path.exists() + assert service.status().active_artifact_count == 1 + reservation.release() + second = service.clear_inactive() + assert second.removed_artifacts == 1 + assert not old.path.exists() + + +def test_lru_evicts_oldest_artifact_until_under_budget(tmp_path: Path) -> None: + """The global byte budget removes least-recently-used artifacts first.""" + + repository = QuantCacheRepository(tmp_path / "SyrupQuants") + old = _commit( + repository, + _identity(tmp_path, "old.safetensors", b"old-source"), + b"12345", + ) + new = _commit( + repository, + _identity(tmp_path, "new.safetensors", b"new-source"), + b"1234567", + ) + _set_last_used(old, "2026-01-01T00:00:00+00:00") + _set_last_used(new, "2026-02-01T00:00:00+00:00") + + result = QuantCacheService(repository).enforce_limit(7) + + assert result.removed_artifacts == 1 + assert not old.path.exists() + assert new.path.is_file() + assert result.remaining_bytes == 7 + + +def test_v1_artifact_is_never_reused_but_remains_clearable(tmp_path: Path) -> None: + """Legacy soft quants participate in cleanup without matching v2 profiles.""" + + repository = QuantCacheRepository(tmp_path / "SyrupQuants") + identity = _identity(tmp_path, "anima.safetensors", b"source") + legacy_directory = repository.root / "Anima" / "anima" / "nvfp4" / "legacy" + legacy_directory.mkdir(parents=True) + artifact_path = legacy_directory / "anima--nvfp4.safetensors" + artifact_path.write_bytes(b"legacy") + legacy_payload = { + "schema_version": 1, + "managed_by": "SimpleSyrup", + "source_model": identity.source.display_name, + "source_path": str(identity.source.path), + "source_sha256": identity.source.sha256, + "source_size_bytes": identity.source.size_bytes, + "source_modified_ns": identity.source.modified_ns, + "quantization_format": "nvfp4", + "model_family": "Anima", + "recipe_version": 1, + "artifact_file": artifact_path.name, + "artifact_size_bytes": artifact_path.stat().st_size, + "created_at": "2026-01-01T00:00:00+00:00", + "last_used_at": "2026-01-01T00:00:00+00:00", + } + (legacy_directory / MANIFEST_FILENAME).write_text( + json.dumps(legacy_payload), encoding="utf-8" + ) + + assert ( + repository.find_current( + source_model=identity.source.display_name, + source_path=identity.source.path, + source_size_bytes=identity.source.size_bytes, + source_modified_ns=identity.source.modified_ns, + profile=identity.profile, + recipe=AnimaQuantizationRecipe(), + ) + is None + ) + assert len(repository.list_artifacts()) == 1 + assert QuantCacheService(repository).clear_inactive().removed_artifacts == 1 + assert not artifact_path.exists() + + +def _identity( + tmp_path: Path, + display_name: str, + content: bytes, +) -> QuantCacheIdentity: + """Create an identity backed by an authoritative test source file.""" + + source_path = tmp_path / f"source-{hashlib.sha256(content).hexdigest()[:8]}.bin" + source_path.write_bytes(content) + stat = source_path.stat() + return QuantCacheIdentity( + source=SourceCheckpointIdentity( + display_name=display_name, + path=source_path.resolve(), + size_bytes=stat.st_size, + modified_ns=stat.st_mtime_ns, + sha256=hashlib.sha256(content).hexdigest(), + ), + profile=NVFP4_MIXED_PROFILE, + model_family="Anima", + recipe_version=2, + ) + + +def _commit( + repository: QuantCacheRepository, + identity: QuantCacheIdentity, + content: bytes, +) -> QuantCacheArtifact: + """Publish one tiny managed artifact through the production repository.""" + + build = repository.create_build_directory(identity) + (build / repository.artifact_filename(identity)).write_bytes(content) + return repository.commit(identity, build) + + +def _set_last_used(artifact: QuantCacheArtifact, timestamp: str) -> None: + """Set deterministic LRU order in one human-readable manifest.""" + + manifest = replace(artifact.manifest, last_used_at=timestamp) + artifact.manifest_path.write_text( + json.dumps(manifest.to_payload(), indent=2, sort_keys=True) + "\n", + encoding="utf-8", + ) + assert artifact.manifest_path.name == MANIFEST_FILENAME diff --git a/tests/test_quant_cache_routes.py b/tests/test_quant_cache_routes.py new file mode 100644 index 0000000..7718203 --- /dev/null +++ b/tests/test_quant_cache_routes.py @@ -0,0 +1,160 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for global quant cache settings-menu backend routes.""" + +from __future__ import annotations + +import asyncio +import json +from collections.abc import Callable +from typing import cast + +from aiohttp import web + +from simple_syrup.runtime.quant_cache_routes import ( + QUANT_CACHE_ROUTE, + Handler, + QuantCachePromptServerProtocol, + register_quant_cache_routes, +) +from simple_syrup.services.quant_cache_service import ( + QuantCacheEvictionResult, + QuantCacheStatus, +) + + +class FakeRoutes: + """Record GET and DELETE cache handlers.""" + + def __init__(self) -> None: + """Create empty route maps.""" + + self.get_handlers: dict[str, Handler] = {} + self.post_handlers: dict[str, Handler] = {} + self.delete_handlers: dict[str, Handler] = {} + + def get(self, path: str) -> Callable[[Handler], Handler]: + """Record a GET handler decorator.""" + + def decorator(handler: Handler) -> Handler: + self.get_handlers[path] = handler + return handler + + return decorator + + def delete(self, path: str) -> Callable[[Handler], Handler]: + """Record a DELETE handler decorator.""" + + def decorator(handler: Handler) -> Handler: + self.delete_handlers[path] = handler + return handler + + return decorator + + def post(self, path: str) -> Callable[[Handler], Handler]: + """Record a POST handler decorator.""" + + def decorator(handler: Handler) -> Handler: + self.post_handlers[path] = handler + return handler + + return decorator + + +class FakePromptServer: + """Expose the fake cache route table.""" + + def __init__(self) -> None: + """Create one route table.""" + + self.routes = FakeRoutes() + + +class FakeCacheService: + """Return deterministic status and clear results.""" + + def __init__(self) -> None: + """Create a service with one visible artifact.""" + + self.cleared = False + + def status(self) -> QuantCacheStatus: + """Return current fake cache state.""" + + return QuantCacheStatus( + path="models/SyrupQuants", + usage_bytes=0 if self.cleared else 1024, + artifact_count=0 if self.cleared else 1, + active_artifact_count=0, + ) + + def clear_inactive(self) -> QuantCacheEvictionResult: + """Record clearing and return its result.""" + + self.cleared = True + return QuantCacheEvictionResult(1, 1024, 0) + + def enforce_limit(self, limit_bytes: int) -> QuantCacheEvictionResult: + """Record enforcement while leaving the tiny artifact within budget.""" + + assert limit_bytes == 20 * 1024**3 + return QuantCacheEvictionResult(0, 0, 1024) + + +class FakeLimitProvider: + """Return the configured 20 GiB byte limit.""" + + def limit_bytes(self) -> int: + """Return the deterministic limit.""" + + return 20 * 1024**3 + + +def test_routes_report_and_clear_the_global_cache() -> None: + """Settings UI endpoints expose location, limit, usage, and clear results.""" + + server = FakePromptServer() + cache_service = FakeCacheService() + registered = register_quant_cache_routes( + cache_service=cache_service, + limit_provider=FakeLimitProvider(), + prompt_server=cast(QuantCachePromptServerProtocol, server), + ) + + assert registered is True + get_response = asyncio.run(server.routes.get_handlers[QUANT_CACHE_ROUTE](object())) + post_response = asyncio.run( + server.routes.post_handlers[QUANT_CACHE_ROUTE](object()) + ) + delete_response = asyncio.run( + server.routes.delete_handlers[QUANT_CACHE_ROUTE](object()) + ) + + assert _payload(get_response) == { + "path": "models/SyrupQuants", + "usage_bytes": 1024, + "limit_bytes": 20 * 1024**3, + "artifact_count": 1, + "active_artifact_count": 0, + } + assert _payload(delete_response) == { + "path": "models/SyrupQuants", + "usage_bytes": 0, + "limit_bytes": 20 * 1024**3, + "artifact_count": 0, + "active_artifact_count": 0, + "removed_artifacts": 1, + "removed_bytes": 1024, + } + assert _payload(post_response)["removed_artifacts"] == 0 + + +def _payload(response: web.Response) -> dict[str, object]: + """Decode a JSON response created by aiohttp helpers.""" + + assert response.text is not None + payload = json.loads(response.text) + assert isinstance(payload, dict) + return cast(dict[str, object], payload) diff --git a/tests/test_quantization_capabilities.py b/tests/test_quantization_capabilities.py new file mode 100644 index 0000000..0ff431d --- /dev/null +++ b/tests/test_quantization_capabilities.py @@ -0,0 +1,129 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for GPU-aware ComfyUI quantization capability discovery.""" + +from __future__ import annotations + +from types import ModuleType + +import pytest + +from simple_syrup.domain.anima_quantization import ( + NVFP4_MIXED_PROFILE, + AnimaQuantizationRecipe, +) +from simple_syrup.domain.model_quantization import QuantizationFormat +from simple_syrup.runtime.comfy_safetensors_dtypes import ( + ComfySafetensorsDtypeRegistry, +) +from simple_syrup.runtime.quantization_capabilities import ( + QuantizationCapabilityCatalog, +) + + +class FakeQuantOps(ModuleType): + """Expose an injectable subset of ComfyUI's quantization registry.""" + + def __init__(self, formats: tuple[str, ...]) -> None: + """Register the requested fake formats.""" + + super().__init__("comfy.quant_ops") + self.QUANT_ALGOS = {name: object() for name in formats} + + +class FakeModelManagement(ModuleType): + """Expose deterministic device capability decisions.""" + + def __init__(self, *, fp8: bool, nvfp4: bool, mxfp8: bool) -> None: + """Create a fake with explicit format support.""" + + super().__init__("comfy.model_management") + self._fp8 = fp8 + self._nvfp4 = nvfp4 + self._mxfp8 = mxfp8 + + def get_torch_device(self) -> str: + """Return a stable fake device.""" + + return "cuda:0" + + def supports_fp8_compute(self, device: object) -> bool: + """Return configured FP8 support.""" + + del device + return self._fp8 + + def supports_nvfp4_compute(self, device: object) -> bool: + """Return configured NVFP4 support.""" + + del device + return self._nvfp4 + + def supports_mxfp8_compute(self, device: object) -> bool: + """Return configured MXFP8 support.""" + + del device + return self._mxfp8 + + +class FakeDtypeRegistry(ComfySafetensorsDtypeRegistry): + """Expose deterministic loader compatibility decisions.""" + + def __init__( + self, unsupported: frozenset[QuantizationFormat] = frozenset() + ) -> None: + """Store formats that the fake loader cannot deserialize.""" + + self._unsupported = unsupported + + def supports_format(self, quantization_format: QuantizationFormat) -> bool: + """Return whether the fake loader supports one format.""" + + return quantization_format not in self._unsupported + + +def test_catalog_intersects_gpu_support_with_comfy_registry() -> None: + """Dropdown choices require both device support and a registered layout.""" + + catalog = QuantizationCapabilityCatalog( + FakeQuantOps(("float8_e4m3fn", "float8_e5m2", "nvfp4")), + FakeModelManagement(fp8=True, nvfp4=True, mxfp8=True), + FakeDtypeRegistry(), + ) + + assert catalog.available_formats() == frozenset( + { + QuantizationFormat.FP8_E4M3, + QuantizationFormat.FP8_E5M2, + QuantizationFormat.NVFP4, + } + ) + + +def test_catalog_always_keeps_original_and_rejects_unavailable_format() -> None: + """Users can always load the source model even without quant GPU support.""" + + catalog = QuantizationCapabilityCatalog( + FakeQuantOps(("nvfp4",)), + FakeModelManagement(fp8=False, nvfp4=False, mxfp8=False), + FakeDtypeRegistry(), + ) + + profiles = AnimaQuantizationRecipe().profiles + assert catalog.selection_labels(profiles) == ["Original"] + with pytest.raises(ValueError, match="unavailable on the current"): + catalog.require_available(NVFP4_MIXED_PROFILE) + + +def test_catalog_hides_profile_when_loader_cannot_read_required_dtype() -> None: + """Compute support does not advertise an unloadable MXFP8 checkpoint.""" + + catalog = QuantizationCapabilityCatalog( + FakeQuantOps(("mxfp8",)), + FakeModelManagement(fp8=False, nvfp4=False, mxfp8=True), + FakeDtypeRegistry(frozenset({QuantizationFormat.MXFP8})), + ) + + assert catalog.selection_labels(AnimaQuantizationRecipe().profiles) == ["Original"] diff --git a/tests/test_quantized_model_resolver.py b/tests/test_quantized_model_resolver.py new file mode 100644 index 0000000..9003146 --- /dev/null +++ b/tests/test_quantized_model_resolver.py @@ -0,0 +1,360 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for reusable quantized checkpoint resolution and generation.""" + +from __future__ import annotations + +import logging +import threading +import time +from dataclasses import dataclass, field +from pathlib import Path + +import pytest + +from simple_syrup.domain.anima_quantization import ( + NVFP4_MIXED_PROFILE, + AnimaQuantizationRecipe, +) +from simple_syrup.domain.model_quantization import ( + ModelQuantizationRecipe, + QuantizationFormat, + QuantizationProfile, + TensorDescriptor, +) +from simple_syrup.domain.quant_cache import SourceCheckpointIdentity +from simple_syrup.runtime.quant_cache_leases import QuantCacheLeaseRegistry +from simple_syrup.runtime.quant_cache_repository import QuantCacheRepository +from simple_syrup.services.quantized_model_resolver import QuantizedModelResolver + + +class RecordingCapabilities: + """Accept and record requested test formats.""" + + def __init__(self) -> None: + """Create an empty request list.""" + + self.requests: list[QuantizationProfile] = [] + + def require_available(self, profile: QuantizationProfile) -> None: + """Record an accepted profile.""" + + self.requests.append(profile) + + +class FixedLimitProvider: + """Return a deterministic generous cache budget.""" + + def limit_bytes(self) -> int: + """Return one GiB for tiny test artifacts.""" + + return 1024**3 + + +@dataclass(frozen=True) +class FakeQuantizationResult: + """Expose output size through the quantizer result boundary.""" + + output_size_bytes: int + + +class RecordingQuantizer: + """Write a tiny artifact and record conversion calls.""" + + def __init__(self, delay: float = 0.0, fail: bool = False) -> None: + """Create a fake with optional concurrency delay or failure.""" + + self.delay = delay + self.fail = fail + self.calls: list[SourceCheckpointIdentity] = [] + self._lock = threading.Lock() + + def quantize( + self, + *, + source: SourceCheckpointIdentity, + destination_path: Path, + profile: QuantizationProfile, + recipe: ModelQuantizationRecipe, + progress: object, + progress_base: int, + progress_total: int, + ) -> FakeQuantizationResult: + """Write deterministic bytes through the production build directory.""" + + del profile, recipe, progress, progress_base, progress_total + with self._lock: + self.calls.append(source) + if self.delay: + time.sleep(self.delay) + if self.fail: + raise RuntimeError("conversion failed") + destination_path.write_bytes(b"quantized") + return FakeQuantizationResult(destination_path.stat().st_size) + + +@dataclass +class RecordingProgress: + """Record resolver-owned progress lifecycle and absolute updates.""" + + starts: list[tuple[str, int]] = field(default_factory=list) + updates: list[tuple[int, int]] = field(default_factory=list) + finishes: int = 0 + + def start(self, label: str, total: int) -> None: + """Record operation start.""" + + self.starts.append((label, total)) + + def advance(self, current: int, total: int) -> None: + """Record one absolute update.""" + + self.updates.append((current, total)) + + def finish(self) -> None: + """Record operation completion.""" + + self.finishes += 1 + + +@dataclass(frozen=True) +class AlternateLoaderRecipe: + """Represent a future loader's independently versioned quantization policy.""" + + model_family: str = "FutureLoader" + version: int = 3 + + @property + def profiles(self) -> tuple[QuantizationProfile, ...]: + """Return the illustrative loader's supported profiles.""" + + return (NVFP4_MIXED_PROFILE,) + + def profile_from_selection(self, selection: str) -> QuantizationProfile: + """Parse the illustrative loader's sole profile.""" + + if selection in (NVFP4_MIXED_PROFILE.label, NVFP4_MIXED_PROFILE.profile_id): + return NVFP4_MIXED_PROFILE + raise ValueError("unsupported profile") + + def policy_for( + self, + tensor: TensorDescriptor, + profile: QuantizationProfile, + ) -> QuantizationFormat | None: + """Select matrices for the illustrative future loader.""" + + del profile + return QuantizationFormat.NVFP4 if len(tensor.shape) == 2 else None + + +def test_resolver_generates_once_then_uses_the_global_cache( + tmp_path: Path, + caplog: pytest.LogCaptureFixture, +) -> None: + """A cache miss converts once and later unchanged requests are cache hits.""" + + source = tmp_path / "anima.safetensors" + source.write_bytes(b"authoritative-source") + repository = QuantCacheRepository(tmp_path / "models" / "SyrupQuants") + quantizer = RecordingQuantizer() + capabilities = RecordingCapabilities() + leases = QuantCacheLeaseRegistry() + resolver = QuantizedModelResolver( + repository=repository, + quantizer=quantizer, + capabilities=capabilities, + limit_provider=FixedLimitProvider(), + leases=leases, + ) + first_progress = RecordingProgress() + + with caplog.at_level(logging.INFO): + first = resolver.resolve( + source_model="Anima/anima.safetensors", + source_path=source, + quantization="nvfp4-mixed", + recipe=AnimaQuantizationRecipe(), + progress=first_progress, + ) + assert first.reservation is not None + first.reservation.release() + second_progress = RecordingProgress() + with caplog.at_level(logging.INFO): + second = resolver.resolve( + source_model="Anima/anima.safetensors", + source_path=source, + quantization="nvfp4-mixed", + recipe=AnimaQuantizationRecipe(), + progress=second_progress, + ) + assert second.reservation is not None + second.reservation.release() + + assert first.path == second.path + assert len(quantizer.calls) == 1 + assert capabilities.requests == [ + NVFP4_MIXED_PROFILE, + NVFP4_MIXED_PROFILE, + ] + assert len(first_progress.starts) == 1 + assert first_progress.finishes == 1 + assert second_progress.starts == [] + assert (repository.root / "README.txt").is_file() + resolver_info = [ + record.message + for record in caplog.records + if record.name.endswith("quantized_model_resolver") + and record.levelno == logging.INFO + ] + assert resolver_info[0] == ( + "creating NVFP4 (Mixed) cache for Anima/anima.safetensors" + ) + assert resolver_info[1].startswith( + "created NVFP4 (Mixed) cache for Anima/anima.safetensors in " + ) + assert len(resolver_info) == 2 + + +def test_original_resolution_does_not_read_or_create_the_cache(tmp_path: Path) -> None: + """Original selection returns its path without touching a missing source file.""" + + repository = QuantCacheRepository(tmp_path / "SyrupQuants") + resolver = QuantizedModelResolver( + repository=repository, + quantizer=RecordingQuantizer(), + capabilities=RecordingCapabilities(), + limit_provider=FixedLimitProvider(), + ) + missing_source = tmp_path / "not-read.safetensors" + + resolved = resolver.resolve( + source_model="not-read.safetensors", + source_path=missing_source, + quantization="Original", + recipe=AnimaQuantizationRecipe(), + ) + + assert resolved.path == missing_source + assert resolved.cache_artifact is None + assert not repository.root.exists() + + +def test_model_family_and_recipe_version_are_part_of_cache_identity( + tmp_path: Path, +) -> None: + """Future loaders can reuse infrastructure without sharing incompatible quants.""" + + source = tmp_path / "shared.safetensors" + source.write_bytes(b"shared-authoritative-source") + repository = QuantCacheRepository(tmp_path / "SyrupQuants") + quantizer = RecordingQuantizer() + resolver = QuantizedModelResolver( + repository=repository, + quantizer=quantizer, + capabilities=RecordingCapabilities(), + limit_provider=FixedLimitProvider(), + ) + + anima = resolver.resolve( + source_model="shared.safetensors", + source_path=source, + quantization="nvfp4-mixed", + recipe=AnimaQuantizationRecipe(), + progress=RecordingProgress(), + ) + future = resolver.resolve( + source_model="shared.safetensors", + source_path=source, + quantization="nvfp4-mixed", + recipe=AlternateLoaderRecipe(), + progress=RecordingProgress(), + ) + assert anima.reservation is not None + assert future.reservation is not None + anima.reservation.release() + future.reservation.release() + + assert anima.path != future.path + assert "Anima" in anima.path.parts + assert "FutureLoader" in future.path.parts + assert len(quantizer.calls) == 2 + + +def test_failed_generation_removes_partial_build_directory(tmp_path: Path) -> None: + """Conversion failures leave no valid-looking partial cache artifacts.""" + + source = tmp_path / "anima.safetensors" + source.write_bytes(b"source") + repository = QuantCacheRepository(tmp_path / "SyrupQuants") + resolver = QuantizedModelResolver( + repository=repository, + quantizer=RecordingQuantizer(fail=True), + capabilities=RecordingCapabilities(), + limit_provider=FixedLimitProvider(), + ) + + with pytest.raises(RuntimeError, match="conversion failed"): + resolver.resolve( + source_model="anima.safetensors", + source_path=source, + quantization="nvfp4-mixed", + recipe=AnimaQuantizationRecipe(), + progress=RecordingProgress(), + ) + + building_root = repository.root / ".building" + assert not building_root.exists() or list(building_root.iterdir()) == [] + assert repository.list_artifacts() == () + + +def test_concurrent_requests_share_one_generated_artifact(tmp_path: Path) -> None: + """Thread and file locks prevent duplicate work for the same cache identity.""" + + source = tmp_path / "anima.safetensors" + source.write_bytes(b"same-source") + repository = QuantCacheRepository(tmp_path / "SyrupQuants") + quantizer = RecordingQuantizer(delay=0.2) + leases = QuantCacheLeaseRegistry() + resolver = QuantizedModelResolver( + repository=repository, + quantizer=quantizer, + capabilities=RecordingCapabilities(), + limit_provider=FixedLimitProvider(), + leases=leases, + ) + results: list[Path] = [] + failures: list[BaseException] = [] + results_lock = threading.Lock() + + def resolve_once() -> None: + """Resolve and release one temporary cache reservation.""" + + try: + resolved = resolver.resolve( + source_model="anima.safetensors", + source_path=source, + quantization="nvfp4-mixed", + recipe=AnimaQuantizationRecipe(), + progress=RecordingProgress(), + ) + assert resolved.reservation is not None + resolved.reservation.release() + with results_lock: + results.append(resolved.path) + except BaseException as error: + with results_lock: + failures.append(error) + + threads = [threading.Thread(target=resolve_once) for _ in range(2)] + for thread in threads: + thread.start() + for thread in threads: + thread.join(timeout=5) + + assert failures == [] + assert len(results) == 2 + assert results[0] == results[1] + assert len(quantizer.calls) == 1 diff --git a/tests/test_settings.py b/tests/test_settings.py index bef5775..676ba4f 100644 --- a/tests/test_settings.py +++ b/tests/test_settings.py @@ -15,14 +15,28 @@ import pytest from simple_syrup.runtime.settings import ( SimpleSyrupSettings, SimpleSyrupSettingsError, - SimpleSyrupSettingsRepository, ) +from simple_syrup.runtime.settings_repository import SimpleSyrupSettingsRepository def test_default_settings_show_downloadable_models() -> None: """Default settings favor low-friction model discovery.""" assert SimpleSyrupSettings().show_downloadable_models is True + assert SimpleSyrupSettings().quant_cache_limit_gib == 20 + + +@pytest.mark.parametrize("value", [0, 2049, 1.5, True, "20"]) +def test_settings_reject_invalid_quant_cache_limits(value: object) -> None: + """The global cache limit is a bounded whole number of GiB.""" + + with pytest.raises(SimpleSyrupSettingsError, match="quant_cache_limit_gib"): + SimpleSyrupSettings.from_payload( + { + "show_downloadable_models": True, + "quant_cache_limit_gib": value, + } + ) def test_missing_settings_file_returns_defaults(tmp_path: Path) -> None: @@ -83,6 +97,7 @@ def test_saving_settings_writes_validated_schema(tmp_path: Path) -> None: "cached_models": [], "default_model": "", }, + "quant_cache_limit_gib": 20, "show_downloadable_models": False, } diff --git a/tests/test_settings_routes.py b/tests/test_settings_routes.py index 1d152ad..c921eea 100644 --- a/tests/test_settings_routes.py +++ b/tests/test_settings_routes.py @@ -19,8 +19,8 @@ import simple_syrup.runtime.settings_routes as settings_routes from simple_syrup.runtime.settings import ( ExternalLLMSettings, SimpleSyrupSettings, - SimpleSyrupSettingsRepository, ) +from simple_syrup.runtime.settings_repository import SimpleSyrupSettingsRepository from simple_syrup.runtime.settings_routes import ( SETTINGS_ROUTE, Handler, @@ -121,6 +121,7 @@ def test_get_settings_returns_current_settings(tmp_path: Path) -> None: "cached_models": [], "default_model": "", }, + "quant_cache_limit_gib": 20, "show_downloadable_models": False, } @@ -145,11 +146,34 @@ def test_post_settings_validates_and_persists_payload(tmp_path: Path) -> None: "cached_models": [], "default_model": "", }, + "quant_cache_limit_gib": 20, "show_downloadable_models": False, } assert repository.load().show_downloadable_models is False +def test_post_settings_persists_global_quant_cache_limit(tmp_path: Path) -> None: + """The settings route persists the user-selected global cache budget.""" + + prompt_server = FakePromptServer() + repository = SimpleSyrupSettingsRepository(tmp_path / "settings.json") + register_fake_routes(repository, prompt_server) + + response = asyncio.run( + prompt_server.routes.post_handlers[SETTINGS_ROUTE]( + FakeRequest( + { + "show_downloadable_models": True, + "quant_cache_limit_gib": 35, + } + ) + ) + ) + + assert response.status == 200 + assert repository.load().quant_cache_limit_gib == 35 + + def test_post_settings_preserves_external_llm_when_payload_omits_it( tmp_path: Path, ) -> None: @@ -183,6 +207,7 @@ def test_post_settings_preserves_external_llm_when_payload_omits_it( "cached_models": ["model-a"], "default_model": "model-a", }, + "quant_cache_limit_gib": 20, "show_downloadable_models": False, } loaded = repository.load() @@ -190,6 +215,31 @@ def test_post_settings_preserves_external_llm_when_payload_omits_it( assert loaded.external_llm.base_url == "https://provider.example/v1" +def test_post_settings_preserves_quant_limit_when_payload_omits_it( + tmp_path: Path, +) -> None: + """Older or focused setting updates do not reset the global cache budget.""" + + prompt_server = FakePromptServer() + repository = SimpleSyrupSettingsRepository(tmp_path / "settings.json") + repository.save( + SimpleSyrupSettings( + show_downloadable_models=True, + quant_cache_limit_gib=42, + ) + ) + register_fake_routes(repository, prompt_server) + + response = asyncio.run( + prompt_server.routes.post_handlers[SETTINGS_ROUTE]( + FakeRequest({"show_downloadable_models": False}) + ) + ) + + assert response.status == 200 + assert repository.load().quant_cache_limit_gib == 42 + + def test_post_settings_rejects_non_boolean_payload(tmp_path: Path) -> None: """POST rejects malformed setting values.""" diff --git a/tests/test_simple_load_anima_node.py b/tests/test_simple_load_anima_node.py index 59b211c..afb8292 100644 --- a/tests/test_simple_load_anima_node.py +++ b/tests/test_simple_load_anima_node.py @@ -15,6 +15,9 @@ import pytest from simple_syrup.nodes.simple_load_anima import SimpleLoadAnima from simple_syrup.nodes_v3.legacy_node_wrappers import SimpleLoadAnimaV3 from simple_syrup.runtime.model_downloads import ComfyProgressReporter +from simple_syrup.runtime.quantization_progress import ( + ComfyQuantizationProgressReporter, +) class FakeFolderPaths(ModuleType): @@ -37,6 +40,16 @@ class FakeFolderPaths(ModuleType): return self.files[folder_name] +class FakeQuantizationCapabilities: + """Return deterministic quantization labels for node schema tests.""" + + def selection_labels(self, profiles: object) -> list[str]: + """Return representative original and GPU-backed choices.""" + + del profiles + return ["Original", "NVFP4 (Mixed)"] + + def test_simple_load_anima_contract() -> None: """Simple Load Anima exposes MODEL, CLIP, and VAE sockets.""" @@ -52,12 +65,18 @@ def test_simple_load_anima_declares_expected_inputs( """Input declarations mirror the combined ComfyUI loader controls.""" monkeypatch.setitem(sys.modules, "folder_paths", FakeFolderPaths()) + monkeypatch.setattr( + SimpleLoadAnima, + "_quantization_capabilities", + FakeQuantizationCapabilities(), + ) input_types: dict[str, dict[str, tuple[Any, ...]]] = SimpleLoadAnima.INPUT_TYPES() required = input_types["required"] assert list(required) == [ "diffusion_model", + "quantization", "diffusion_weight_dtype", "text_encoder", "text_encoder_device", @@ -65,6 +84,9 @@ def test_simple_load_anima_declares_expected_inputs( ] assert required["diffusion_model"][0] == ["anima.safetensors"] assert "advanced" not in required["diffusion_model"][1] + assert required["quantization"][0] == ["Original", "NVFP4 (Mixed)"] + assert required["quantization"][1]["default"] == "Original" + assert required["quantization"][1]["advanced"] is True assert required["diffusion_weight_dtype"][0] == [ "default", "fp8_e4m3fn", @@ -88,11 +110,17 @@ def test_simple_load_anima_v3_keeps_only_diffusion_model_primary( """The v3 schema marks every control after diffusion selection as advanced.""" monkeypatch.setitem(sys.modules, "folder_paths", FakeFolderPaths()) + monkeypatch.setattr( + SimpleLoadAnima, + "_quantization_capabilities", + FakeQuantizationCapabilities(), + ) schema = SimpleLoadAnimaV3.define_schema() assert [input_item.id for input_item in schema.inputs] == [ "diffusion_model", + "quantization", "diffusion_weight_dtype", "text_encoder", "text_encoder_device", @@ -125,6 +153,7 @@ def test_simple_load_anima_delegates_to_service() -> None: try: result = SimpleLoadAnima().load_models( diffusion_model="anima.safetensors", + quantization="Original", diffusion_weight_dtype="default", text_encoder="auto", text_encoder_device="default", @@ -137,4 +166,9 @@ def test_simple_load_anima_delegates_to_service() -> None: assert fake_service.kwargs is not None assert fake_service.kwargs["text_encoder"] == "auto" assert fake_service.kwargs["vae"] == "auto" + assert fake_service.kwargs["quantization"] == "Original" assert isinstance(fake_service.kwargs["progress"], ComfyProgressReporter) + assert isinstance( + fake_service.kwargs["quantization_progress"], + ComfyQuantizationProgressReporter, + ) diff --git a/web/dist/simple-syrup.js b/web/dist/simple-syrup.js index 5dceb46..932e7bb 100644 --- a/web/dist/simple-syrup.js +++ b/web/dist/simple-syrup.js @@ -3,6 +3,7 @@ import { app } from "../../../scripts/app.js"; // web/src/api.ts var SETTINGS_ROUTE = "/simple-syrup/settings"; +var QUANT_CACHE_ROUTE = "/simple-syrup/quant-cache"; var EXTERNAL_LLM_SETTINGS_ROUTE = "/simple-syrup/external-llm/settings"; var EXTERNAL_LLM_API_KEY_ROUTE = "/simple-syrup/external-llm/api-key"; var EXTERNAL_LLM_MODELS_REFRESH_ROUTE = "/simple-syrup/external-llm/models/refresh"; @@ -69,9 +70,54 @@ function parseSettings(payload) { ); } return { - show_downloadable_models: payload.show_downloadable_models + show_downloadable_models: payload.show_downloadable_models, + quant_cache_limit_gib: payload.quant_cache_limit_gib }; } +async function getQuantCacheStatus(fetchImpl = fetch) { + const response = await fetchImpl(QUANT_CACHE_ROUTE); + if (!response.ok) { + throw new Error( + await backendErrorMessage( + response, + `Could not load quant cache status. Backend returned ${String(response.status)}.` + ) + ); + } + return parseQuantCacheStatus(await response.json()); +} +async function clearQuantCache(fetchImpl = fetch) { + const response = await fetchImpl(QUANT_CACHE_ROUTE, { method: "DELETE" }); + if (!response.ok) { + throw new Error( + await backendErrorMessage( + response, + `Could not clear quant cache. Backend returned ${String(response.status)}.` + ) + ); + } + return parseQuantCacheStatus(await response.json()); +} +async function enforceQuantCacheLimit(fetchImpl = fetch) { + const response = await fetchImpl(QUANT_CACHE_ROUTE, { method: "POST" }); + if (!response.ok) { + throw new Error( + await backendErrorMessage( + response, + `Could not enforce quant cache limit. Backend returned ${String(response.status)}.` + ) + ); + } + return parseQuantCacheStatus(await response.json()); +} +function parseQuantCacheStatus(payload) { + if (!isQuantCacheStatusPayload(payload)) { + throw new Error( + "SimpleSyrup quant cache status is invalid. Expected path, byte usage, limit, and artifact counts." + ); + } + return { ...payload }; +} async function getExternalLLMSettings(fetchImpl = fetch) { const response = await fetchImpl(EXTERNAL_LLM_SETTINGS_ROUTE); if (!response.ok) { @@ -144,7 +190,17 @@ function parseExternalLLMSettings(payload) { }; } function isSettingsPayload(payload) { - return typeof payload === "object" && payload !== null && typeof payload.show_downloadable_models === "boolean"; + return typeof payload === "object" && payload !== null && typeof payload.show_downloadable_models === "boolean" && Number.isInteger( + payload.quant_cache_limit_gib + ) && Number(payload.quant_cache_limit_gib) > 0; +} +function isQuantCacheStatusPayload(payload) { + if (typeof payload !== "object" || payload === null) return false; + const candidate = payload; + return typeof candidate.path === "string" && isNonNegativeInteger(candidate.usage_bytes) && isNonNegativeInteger(candidate.limit_bytes) && isNonNegativeInteger(candidate.artifact_count) && isNonNegativeInteger(candidate.active_artifact_count) && (candidate.removed_artifacts === void 0 || isNonNegativeInteger(candidate.removed_artifacts)) && (candidate.removed_bytes === void 0 || isNonNegativeInteger(candidate.removed_bytes)); +} +function isNonNegativeInteger(value) { + return typeof value === "number" && Number.isInteger(value) && value >= 0; } function isExternalLLMSettingsPayload(payload) { return typeof payload === "object" && payload !== null && typeof payload.base_url === "string" && Array.isArray(payload.cached_models) && payload.cached_models?.every( @@ -173,124 +229,41 @@ async function backendErrorMessage(response, fallback) { return fallback; } -// web/src/settings.ts +// web/src/downloadableModelsSetting.ts var SIMPLE_SYRUP_SETTING_ID = "SimpleSyrup.ShowDownloadableModels"; var SIMPLE_SYRUP_SETTING_LABEL = "SimpleSyrup: Show downloadable models in loader dropdowns"; var SIMPLE_SYRUP_SETTING_DESCRIPTION = "Show known downloadable SAM, GroundingDINO, and ViTMatte models even when they are not installed locally."; -var EXTERNAL_LLM_ENDPOINT_SETTING_ID = "SimpleSyrup.ExternalLLM.Endpoint"; -var EXTERNAL_LLM_ENDPOINT_SETTING_LABEL = "SimpleSyrup: External LLM endpoint"; -var EXTERNAL_LLM_ENDPOINT_SETTING_DESCRIPTION = "OpenAI-compatible endpoint base URL used by SimpleSyrup prompt nodes."; -var EXTERNAL_LLM_API_KEY_SETTING_ID = "SimpleSyrup.ExternalLLM.ApiKey"; -var EXTERNAL_LLM_API_KEY_SETTING_LABEL = "SimpleSyrup: External LLM API key"; -var EXTERNAL_LLM_API_KEY_SETTING_DESCRIPTION = "Stores the API key for the configured external LLM endpoint in OS credential storage."; -var DEFAULT_SETTINGS = { - show_downloadable_models: true -}; -async function registerSimpleSyrupSettings(app2, api = { - getSettings, - saveSettings, - getExternalLLMSettings, - saveExternalLLMSettings, - saveExternalLLMApiKey -}, logger = console) { - let initialSettings = DEFAULT_SETTINGS; - let externalLLMSettings = { - base_url: "", - cached_models: [], - default_model: "", - has_api_key: false - }; - try { - initialSettings = await api.getSettings(); - } catch (error) { - logger.warn( - "Could not load SimpleSyrup settings. Using the default setting until the backend is available.", - error - ); - } - try { - externalLLMSettings = await api.getExternalLLMSettings(); - } catch (error) { - logger.warn( - "Could not load SimpleSyrup external LLM settings. Using empty endpoint settings until the backend is available.", - error - ); - } - let savedSettings = initialSettings; - let savedExternalLLMSettings = externalLLMSettings; - installSimpleSyrupSettingsStyle(); +function registerDownloadableModelsSetting(app2, context, logger) { const setting = app2.ui.settings.addSetting({ id: SIMPLE_SYRUP_SETTING_ID, name: SIMPLE_SYRUP_SETTING_LABEL, type: "boolean", - defaultValue: initialSettings.show_downloadable_models, + defaultValue: context.getSettings().show_downloadable_models, tooltip: SIMPLE_SYRUP_SETTING_DESCRIPTION, onChange: async (value) => { + const previous = context.getSettings(); try { - const saved = await api.saveSettings({ + const saved = await context.saveSettings({ + ...previous, show_downloadable_models: value }); - savedSettings = saved; + context.setSettings(saved); setting.value = saved.show_downloadable_models; } catch (error) { logger.warn( "Could not save SimpleSyrup settings. The backend rejected the setting update.", error ); - setting.value = savedSettings.show_downloadable_models; + setting.value = previous.show_downloadable_models; } } }); - setting.value = initialSettings.show_downloadable_models; - app2.ui.settings.addSetting({ - id: EXTERNAL_LLM_ENDPOINT_SETTING_ID, - name: EXTERNAL_LLM_ENDPOINT_SETTING_LABEL, - sortOrder: 320, - type: () => createExternalLLMEndpointControl({ - api, - logger, - refreshModelChoices: () => refreshExternalLLMModelChoices(app2, logger), - getSettings: () => savedExternalLLMSettings, - setSettings: (settings) => { - savedExternalLLMSettings = settings; - } - }), - defaultValue: externalLLMSettings.base_url, - tooltip: EXTERNAL_LLM_ENDPOINT_SETTING_DESCRIPTION - }); - app2.ui.settings.addSetting({ - id: EXTERNAL_LLM_API_KEY_SETTING_ID, - name: EXTERNAL_LLM_API_KEY_SETTING_LABEL, - sortOrder: 319, - type: () => createExternalLLMApiKeyControl({ - api, - logger, - refreshModelChoices: () => refreshExternalLLMModelChoices(app2, logger), - getSettings: () => savedExternalLLMSettings, - setSettings: (settings) => { - savedExternalLLMSettings = settings; - } - }), - defaultValue: "", - tooltip: EXTERNAL_LLM_API_KEY_SETTING_DESCRIPTION - }); -} -function endpointShouldBeSaved(value) { - const endpoint = value.trim(); - if (!endpoint) { - return true; - } - try { - const parsed = new URL(endpoint); - return (parsed.protocol === "http:" || parsed.protocol === "https:") && parsed.hostname.length > 0; - } catch { - return false; - } + setting.value = context.getSettings().show_downloadable_models; } + +// web/src/settingsUi.ts function installSimpleSyrupSettingsStyle() { - if (document.getElementById("simple-syrup-settings-style")) { - return; - } + if (document.getElementById("simple-syrup-settings-style")) return; const style = document.createElement("style"); style.id = "simple-syrup-settings-style"; style.textContent = ` @@ -346,6 +319,93 @@ function installSimpleSyrupSettingsStyle() { `; document.head.appendChild(style); } +function createElement(tagName, className) { + const element = document.createElement(tagName); + element.className = className; + return element; +} +function setPending(element, pending) { + element.dataset.pending = pending ? "true" : "false"; + for (const control of Array.from(element.querySelectorAll("input, button"))) { + if (control instanceof HTMLInputElement || control instanceof HTMLButtonElement) { + control.disabled = pending; + } + } +} + +// web/src/externalLlmSettings.ts +var EXTERNAL_LLM_ENDPOINT_SETTING_ID = "SimpleSyrup.ExternalLLM.Endpoint"; +var EXTERNAL_LLM_ENDPOINT_SETTING_LABEL = "SimpleSyrup: External LLM endpoint"; +var EXTERNAL_LLM_ENDPOINT_SETTING_DESCRIPTION = "OpenAI-compatible endpoint base URL used by SimpleSyrup prompt nodes."; +var EXTERNAL_LLM_API_KEY_SETTING_ID = "SimpleSyrup.ExternalLLM.ApiKey"; +var EXTERNAL_LLM_API_KEY_SETTING_LABEL = "SimpleSyrup: External LLM API key"; +var EXTERNAL_LLM_API_KEY_SETTING_DESCRIPTION = "Stores the API key for the configured external LLM endpoint in OS credential storage."; +async function registerExternalLLMSettings(app2, api = { + getExternalLLMSettings, + saveExternalLLMSettings, + saveExternalLLMApiKey +}, logger = console) { + let externalLLMSettings = { + base_url: "", + cached_models: [], + default_model: "", + has_api_key: false + }; + try { + externalLLMSettings = await api.getExternalLLMSettings(); + } catch (error) { + logger.warn( + "Could not load SimpleSyrup external LLM settings. Using empty endpoint settings until the backend is available.", + error + ); + } + let savedExternalLLMSettings = externalLLMSettings; + installSimpleSyrupSettingsStyle(); + app2.ui.settings.addSetting({ + id: EXTERNAL_LLM_ENDPOINT_SETTING_ID, + name: EXTERNAL_LLM_ENDPOINT_SETTING_LABEL, + sortOrder: 320, + type: () => createExternalLLMEndpointControl({ + api, + logger, + refreshModelChoices: () => refreshExternalLLMModelChoices(app2, logger), + getSettings: () => savedExternalLLMSettings, + setSettings: (settings) => { + savedExternalLLMSettings = settings; + } + }), + defaultValue: externalLLMSettings.base_url, + tooltip: EXTERNAL_LLM_ENDPOINT_SETTING_DESCRIPTION + }); + app2.ui.settings.addSetting({ + id: EXTERNAL_LLM_API_KEY_SETTING_ID, + name: EXTERNAL_LLM_API_KEY_SETTING_LABEL, + sortOrder: 319, + type: () => createExternalLLMApiKeyControl({ + api, + logger, + refreshModelChoices: () => refreshExternalLLMModelChoices(app2, logger), + getSettings: () => savedExternalLLMSettings, + setSettings: (settings) => { + savedExternalLLMSettings = settings; + } + }), + defaultValue: "", + tooltip: EXTERNAL_LLM_API_KEY_SETTING_DESCRIPTION + }); +} +function endpointShouldBeSaved(value) { + const endpoint = value.trim(); + if (!endpoint) { + return true; + } + try { + const parsed = new URL(endpoint); + return (parsed.protocol === "http:" || parsed.protocol === "https:") && parsed.hostname.length > 0; + } catch { + return false; + } +} function createExternalLLMEndpointControl(context) { const wrapper = createElement("div", "simple-syrup-settings-row"); const input = createElement("input", "simple-syrup-settings-input"); @@ -489,19 +549,6 @@ function openExternalLLMApiKeyDialog(options) { document.body.appendChild(overlay); input.focus(); } -function createElement(tagName, className) { - const element = document.createElement(tagName); - element.className = className; - return element; -} -function setPending(element, pending) { - element.dataset.pending = pending ? "true" : "false"; - for (const control of Array.from(element.querySelectorAll("input, button"))) { - if (control instanceof HTMLInputElement || control instanceof HTMLButtonElement) { - control.disabled = pending; - } - } -} function errorMessage(error, fallback) { if (error instanceof Error && error.message) { return error.message; @@ -526,6 +573,144 @@ async function refreshExternalLLMModelChoices(app2, logger) { } } +// web/src/quantCacheSetting.ts +var QUANT_CACHE_SETTING_ID = "SimpleSyrup.QuantCache"; +var QUANT_CACHE_SETTING_LABEL = "SimpleSyrup: Quantized model cache"; +var QUANT_CACHE_SETTING_DESCRIPTION = "Sets the global models/SyrupQuants cache limit in GiB for every SimpleSyrup loader; least-recently-used inactive copies are removed automatically."; +async function registerQuantCacheSetting(app2, settings, api, logger) { + installSimpleSyrupSettingsStyle(); + let initialStatus = null; + try { + initialStatus = await api.getQuantCacheStatus(); + } catch (error) { + logger.warn("Could not load SimpleSyrup quant cache status.", error); + } + app2.ui.settings.addSetting({ + id: QUANT_CACHE_SETTING_ID, + name: QUANT_CACHE_SETTING_LABEL, + sortOrder: 321, + type: () => createQuantCacheControl({ settings, api, logger, initialStatus }), + defaultValue: settings.getSettings().quant_cache_limit_gib, + tooltip: QUANT_CACHE_SETTING_DESCRIPTION + }); +} +function createQuantCacheControl(context) { + const wrapper = createElement("div", "simple-syrup-settings-row"); + const limitInput = createElement("input", "simple-syrup-settings-input"); + limitInput.type = "number"; + limitInput.min = "1"; + limitInput.max = "2048"; + limitInput.step = "1"; + limitInput.value = String(context.settings.getSettings().quant_cache_limit_gib); + limitInput.setAttribute("aria-label", "Quant cache limit in GiB"); + const unitLabel = createElement("span", "simple-syrup-settings-status"); + unitLabel.textContent = "GiB limit"; + const saveButton = createElement("button", "simple-syrup-settings-button"); + saveButton.type = "button"; + saveButton.textContent = "Save Limit"; + const clearButton = createElement("button", "simple-syrup-settings-button"); + clearButton.type = "button"; + clearButton.textContent = "Clear Inactive"; + const status = createElement("span", "simple-syrup-settings-status"); + const renderStatus = (cacheStatus) => { + status.textContent = cacheStatus ? `${formatGiB(cacheStatus.usage_bytes)} GiB used in ${cacheStatus.path} (${String(cacheStatus.artifact_count)} cached, ${String(cacheStatus.active_artifact_count)} active).` : "Cache status unavailable. Generated models are stored in models/SyrupQuants."; + }; + renderStatus(context.initialStatus); + saveButton.addEventListener("click", () => { + void saveLimit(); + }); + limitInput.addEventListener("keydown", (event) => { + if (event.key === "Enter") { + event.preventDefault(); + void saveLimit(); + } + }); + clearButton.addEventListener("click", () => { + void clearInactive(); + }); + const saveLimit = async () => { + const limit = Number(limitInput.value); + if (!Number.isInteger(limit) || limit < 1 || limit > 2048) { + status.textContent = "Enter a whole-number cache limit from 1 to 2048 GiB."; + return; + } + const previous = context.settings.getSettings(); + setPending(wrapper, true); + try { + const saved = await context.settings.saveSettings({ + ...previous, + quant_cache_limit_gib: limit + }); + context.settings.setSettings(saved); + limitInput.value = String(saved.quant_cache_limit_gib); + const refreshed = await context.api.enforceQuantCacheLimit(); + renderStatus(refreshed); + } catch (error) { + context.logger.warn("Could not save SimpleSyrup quant cache limit.", error); + limitInput.value = String(previous.quant_cache_limit_gib); + status.textContent = "Cache limit was not saved."; + } finally { + setPending(wrapper, false); + } + }; + const clearInactive = async () => { + setPending(wrapper, true); + try { + const cleared = await context.api.clearQuantCache(); + renderStatus(cleared); + } catch (error) { + context.logger.warn("Could not clear SimpleSyrup quant cache.", error); + status.textContent = "Inactive cached models were not cleared."; + } finally { + setPending(wrapper, false); + } + }; + wrapper.append(limitInput, unitLabel, saveButton, clearButton, status); + return wrapper; +} +function formatGiB(bytes) { + return (bytes / 1024 ** 3).toFixed(2); +} + +// web/src/settingsRegistration.ts +var DEFAULT_SETTINGS = { + show_downloadable_models: true, + quant_cache_limit_gib: 20 +}; +async function registerSimpleSyrupSettings(app2, api = defaultApi(), logger = console) { + let savedSettings = DEFAULT_SETTINGS; + try { + savedSettings = await api.getSettings(); + } catch (error) { + logger.warn( + "Could not load SimpleSyrup settings. Using defaults until the backend is available.", + error + ); + } + const settingsContext = { + getSettings: () => savedSettings, + saveSettings: (settings) => api.saveSettings(settings), + setSettings: (settings) => { + savedSettings = settings; + } + }; + registerDownloadableModelsSetting(app2, settingsContext, logger); + await registerQuantCacheSetting(app2, settingsContext, api, logger); + await registerExternalLLMSettings(app2, api, logger); +} +function defaultApi() { + return { + getSettings, + saveSettings, + getQuantCacheStatus, + enforceQuantCacheLimit, + clearQuantCache, + getExternalLLMSettings, + saveExternalLLMSettings, + saveExternalLLMApiKey + }; +} + // web/src/refresh.ts var REFRESH_WRAPPED = /* @__PURE__ */ Symbol.for("SimpleSyrup.ExternalLLM.RefreshWrapped"); function registerExternalLLMRefreshHook(app2, api = { refreshExternalLLMModels }, logger = console) { diff --git a/web/src/api.ts b/web/src/api.ts index 22c5094..9275199 100644 --- a/web/src/api.ts +++ b/web/src/api.ts @@ -6,6 +6,17 @@ import type { ComfyImageResult, ComfyNodeExecutionOutput } from "./types"; export interface SimpleSyrupSettings { show_downloadable_models: boolean; + quant_cache_limit_gib: number; +} + +export interface QuantCacheStatus { + path: string; + usage_bytes: number; + limit_bytes: number; + artifact_count: number; + active_artifact_count: number; + removed_artifacts?: number; + removed_bytes?: number; } export interface ExternalLLMSettings { @@ -30,6 +41,7 @@ export type FetchLike = ( ) => Promise; const SETTINGS_ROUTE = "/simple-syrup/settings"; +const QUANT_CACHE_ROUTE = "/simple-syrup/quant-cache"; const EXTERNAL_LLM_SETTINGS_ROUTE = "/simple-syrup/external-llm/settings"; const EXTERNAL_LLM_API_KEY_ROUTE = "/simple-syrup/external-llm/api-key"; const EXTERNAL_LLM_MODELS_REFRESH_ROUTE = @@ -113,10 +125,65 @@ export function parseSettings(payload: unknown): SimpleSyrupSettings { ); } return { - show_downloadable_models: payload.show_downloadable_models + show_downloadable_models: payload.show_downloadable_models, + quant_cache_limit_gib: payload.quant_cache_limit_gib }; } +export async function getQuantCacheStatus( + fetchImpl: FetchLike = fetch +): Promise { + const response = await fetchImpl(QUANT_CACHE_ROUTE); + if (!response.ok) { + throw new Error( + await backendErrorMessage( + response, + `Could not load quant cache status. Backend returned ${String(response.status)}.` + ) + ); + } + return parseQuantCacheStatus(await response.json()); +} + +export async function clearQuantCache( + fetchImpl: FetchLike = fetch +): Promise { + const response = await fetchImpl(QUANT_CACHE_ROUTE, { method: "DELETE" }); + if (!response.ok) { + throw new Error( + await backendErrorMessage( + response, + `Could not clear quant cache. Backend returned ${String(response.status)}.` + ) + ); + } + return parseQuantCacheStatus(await response.json()); +} + +export async function enforceQuantCacheLimit( + fetchImpl: FetchLike = fetch +): Promise { + const response = await fetchImpl(QUANT_CACHE_ROUTE, { method: "POST" }); + if (!response.ok) { + throw new Error( + await backendErrorMessage( + response, + `Could not enforce quant cache limit. Backend returned ${String(response.status)}.` + ) + ); + } + return parseQuantCacheStatus(await response.json()); +} + +export function parseQuantCacheStatus(payload: unknown): QuantCacheStatus { + if (!isQuantCacheStatusPayload(payload)) { + throw new Error( + "SimpleSyrup quant cache status is invalid. Expected path, byte usage, limit, and artifact counts." + ); + } + return { ...payload }; +} + export async function getExternalLLMSettings( fetchImpl: FetchLike = fetch ): Promise { @@ -227,10 +294,36 @@ function isSettingsPayload(payload: unknown): payload is SimpleSyrupSettings { typeof payload === "object" && payload !== null && typeof (payload as Partial).show_downloadable_models === - "boolean" + "boolean" && + Number.isInteger( + (payload as Partial).quant_cache_limit_gib + ) && + Number((payload as Partial).quant_cache_limit_gib) > 0 ); } +function isQuantCacheStatusPayload( + payload: unknown +): payload is QuantCacheStatus { + if (typeof payload !== "object" || payload === null) return false; + const candidate = payload as Partial; + return ( + typeof candidate.path === "string" && + isNonNegativeInteger(candidate.usage_bytes) && + isNonNegativeInteger(candidate.limit_bytes) && + isNonNegativeInteger(candidate.artifact_count) && + isNonNegativeInteger(candidate.active_artifact_count) && + (candidate.removed_artifacts === undefined || + isNonNegativeInteger(candidate.removed_artifacts)) && + (candidate.removed_bytes === undefined || + isNonNegativeInteger(candidate.removed_bytes)) + ); +} + +function isNonNegativeInteger(value: unknown): value is number { + return typeof value === "number" && Number.isInteger(value) && value >= 0; +} + function isExternalLLMSettingsPayload( payload: unknown ): payload is ExternalLLMSettings { diff --git a/web/src/downloadableModelsSetting.ts b/web/src/downloadableModelsSetting.ts new file mode 100644 index 0000000..de8b082 --- /dev/null +++ b/web/src/downloadableModelsSetting.ts @@ -0,0 +1,50 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +import type { SimpleSyrupSettings } from "./api"; +import type { ComfyApp, Logger } from "./types"; + +export const SIMPLE_SYRUP_SETTING_ID = "SimpleSyrup.ShowDownloadableModels"; +export const SIMPLE_SYRUP_SETTING_LABEL = + "SimpleSyrup: Show downloadable models in loader dropdowns"; +export const SIMPLE_SYRUP_SETTING_DESCRIPTION = + "Show known downloadable SAM, GroundingDINO, and ViTMatte models even when they are not installed locally."; + +export interface GeneralSettingsContext { + getSettings(): SimpleSyrupSettings; + saveSettings(settings: SimpleSyrupSettings): Promise; + setSettings(settings: SimpleSyrupSettings): void; +} + +export function registerDownloadableModelsSetting( + app: ComfyApp, + context: GeneralSettingsContext, + logger: Logger +): void { + const setting = app.ui.settings.addSetting({ + id: SIMPLE_SYRUP_SETTING_ID, + name: SIMPLE_SYRUP_SETTING_LABEL, + type: "boolean", + defaultValue: context.getSettings().show_downloadable_models, + tooltip: SIMPLE_SYRUP_SETTING_DESCRIPTION, + onChange: async (value: boolean) => { + const previous = context.getSettings(); + try { + const saved = await context.saveSettings({ + ...previous, + show_downloadable_models: value + }); + context.setSettings(saved); + setting.value = saved.show_downloadable_models; + } catch (error) { + logger.warn( + "Could not save SimpleSyrup settings. The backend rejected the setting update.", + error + ); + setting.value = previous.show_downloadable_models; + } + } + }); + setting.value = context.getSettings().show_downloadable_models; +} diff --git a/web/src/settings.ts b/web/src/externalLlmSettings.ts similarity index 70% rename from web/src/settings.ts rename to web/src/externalLlmSettings.ts index 91d429a..7d2738a 100644 --- a/web/src/settings.ts +++ b/web/src/externalLlmSettings.ts @@ -2,21 +2,21 @@ // Copyright (C) 2026 Artificial Sweetener and contributors // SPDX-License-Identifier: AGPL-3.0-or-later +// Owns the external LLM endpoint and credential settings presentation. + import { getExternalLLMSettings, - getSettings, saveExternalLLMApiKey, - saveExternalLLMSettings, - saveSettings + saveExternalLLMSettings } from "./api"; -import type { ExternalLLMSettings, SimpleSyrupSettings } from "./api"; +import type { ExternalLLMSettings } from "./api"; +import { + createElement, + installSimpleSyrupSettingsStyle, + setPending +} from "./settingsUi"; import type { ComfyApp, Logger } from "./types"; -export const SIMPLE_SYRUP_SETTING_ID = "SimpleSyrup.ShowDownloadableModels"; -export const SIMPLE_SYRUP_SETTING_LABEL = - "SimpleSyrup: Show downloadable models in loader dropdowns"; -export const SIMPLE_SYRUP_SETTING_DESCRIPTION = - "Show known downloadable SAM, GroundingDINO, and ViTMatte models even when they are not installed locally."; export const EXTERNAL_LLM_ENDPOINT_SETTING_ID = "SimpleSyrup.ExternalLLM.Endpoint"; export const EXTERNAL_LLM_ENDPOINT_SETTING_LABEL = @@ -30,13 +30,7 @@ export const EXTERNAL_LLM_API_KEY_SETTING_LABEL = export const EXTERNAL_LLM_API_KEY_SETTING_DESCRIPTION = "Stores the API key for the configured external LLM endpoint in OS credential storage."; -const DEFAULT_SETTINGS: SimpleSyrupSettings = { - show_downloadable_models: true -}; - -export interface SimpleSyrupSettingsApi { - getSettings(): Promise; - saveSettings(settings: SimpleSyrupSettings): Promise; +export interface ExternalLLMSettingsApi { getExternalLLMSettings(): Promise; saveExternalLLMSettings( settings: Pick @@ -44,18 +38,15 @@ export interface SimpleSyrupSettingsApi { saveExternalLLMApiKey(settings: { api_key: string }): Promise; } -export async function registerSimpleSyrupSettings( +export async function registerExternalLLMSettings( app: ComfyApp, - api: SimpleSyrupSettingsApi = { - getSettings, - saveSettings, + api: ExternalLLMSettingsApi = { getExternalLLMSettings, saveExternalLLMSettings, saveExternalLLMApiKey }, logger: Logger = console ): Promise { - let initialSettings = DEFAULT_SETTINGS; let externalLLMSettings: ExternalLLMSettings = { base_url: "", cached_models: [], @@ -63,14 +54,6 @@ export async function registerSimpleSyrupSettings( has_api_key: false }; - try { - initialSettings = await api.getSettings(); - } catch (error) { - logger.warn( - "Could not load SimpleSyrup settings. Using the default setting until the backend is available.", - error - ); - } try { externalLLMSettings = await api.getExternalLLMSettings(); } catch (error) { @@ -79,35 +62,9 @@ export async function registerSimpleSyrupSettings( error ); } - let savedSettings = initialSettings; let savedExternalLLMSettings = externalLLMSettings; installSimpleSyrupSettingsStyle(); - const setting = app.ui.settings.addSetting({ - id: SIMPLE_SYRUP_SETTING_ID, - name: SIMPLE_SYRUP_SETTING_LABEL, - type: "boolean", - defaultValue: initialSettings.show_downloadable_models, - tooltip: SIMPLE_SYRUP_SETTING_DESCRIPTION, - onChange: async (value: boolean) => { - try { - const saved = await api.saveSettings({ - show_downloadable_models: value - }); - savedSettings = saved; - setting.value = saved.show_downloadable_models; - } catch (error) { - logger.warn( - "Could not save SimpleSyrup settings. The backend rejected the setting update.", - error - ); - setting.value = savedSettings.show_downloadable_models; - } - } - }); - - setting.value = initialSettings.show_downloadable_models; - app.ui.settings.addSetting({ id: EXTERNAL_LLM_ENDPOINT_SETTING_ID, name: EXTERNAL_LLM_ENDPOINT_SETTING_LABEL, @@ -163,69 +120,8 @@ function endpointShouldBeSaved(value: string): boolean { } } -function installSimpleSyrupSettingsStyle(): void { - if (document.getElementById("simple-syrup-settings-style")) { - return; - } - - const style = document.createElement("style"); - style.id = "simple-syrup-settings-style"; - style.textContent = ` - .simple-syrup-settings-row { - display: flex; - align-items: center; - gap: 0.5rem; - min-width: min(38rem, 100%); - } - .simple-syrup-settings-row[data-pending="true"] { - opacity: 0.75; - } - .simple-syrup-settings-input { - min-width: 16rem; - flex: 1 1 auto; - } - .simple-syrup-settings-button { - flex: 0 0 auto; - white-space: nowrap; - } - .simple-syrup-settings-status { - color: var(--fg-color); - opacity: 0.8; - white-space: normal; - overflow-wrap: anywhere; - } - .simple-syrup-dialog-backdrop { - position: fixed; - inset: 0; - z-index: 2147483647; - display: flex; - align-items: center; - justify-content: center; - background: rgb(0 0 0 / 45%); - } - .simple-syrup-dialog { - display: grid; - gap: 0.75rem; - min-width: min(28rem, calc(100vw - 2rem)); - padding: 1rem; - background: var(--comfy-menu-bg); - color: var(--fg-color); - } - .simple-syrup-dialog-title { - margin: 0; - font-size: 1rem; - } - .simple-syrup-dialog-actions { - display: flex; - justify-content: flex-end; - gap: 0.5rem; - } - `; - document.head.appendChild(style); -} - interface ExternalLLMControlContext { - api: SimpleSyrupSettingsApi; + api: ExternalLLMSettingsApi; logger: Logger; refreshModelChoices(): Promise; getSettings(): ExternalLLMSettings; @@ -403,27 +299,6 @@ function openExternalLLMApiKeyDialog(options: { input.focus(); } -function createElement( - tagName: TTag, - className: string -): HTMLElementTagNameMap[TTag] { - const element = document.createElement(tagName); - element.className = className; - return element; -} - -function setPending(element: HTMLElement, pending: boolean): void { - element.dataset.pending = pending ? "true" : "false"; - for (const control of Array.from(element.querySelectorAll("input, button"))) { - if ( - control instanceof HTMLInputElement || - control instanceof HTMLButtonElement - ) { - control.disabled = pending; - } - } -} - function errorMessage(error: unknown, fallback: string): string { if (error instanceof Error && error.message) { return error.message; diff --git a/web/src/main.ts b/web/src/main.ts index fb0000c..d2aae75 100644 --- a/web/src/main.ts +++ b/web/src/main.ts @@ -5,7 +5,7 @@ // @ts-expect-error ComfyUI serves this host module outside the extension package. import { app } from "../../../scripts/app.js"; -import { registerSimpleSyrupSettings } from "./settings"; +import { registerSimpleSyrupSettings } from "./settingsRegistration"; import { registerExternalLLMRefreshHook } from "./refresh"; import { registerMaskBatchUpload } from "./maskBatchUpload"; import { registerImageListUpload } from "./imageListUpload"; diff --git a/web/src/quantCacheSetting.ts b/web/src/quantCacheSetting.ts new file mode 100644 index 0000000..ed7809b --- /dev/null +++ b/web/src/quantCacheSetting.ts @@ -0,0 +1,141 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +import type { QuantCacheStatus } from "./api"; +import type { ComfyApp, Logger } from "./types"; +import type { GeneralSettingsContext } from "./downloadableModelsSetting"; +import { + createElement, + installSimpleSyrupSettingsStyle, + setPending +} from "./settingsUi"; + +export const QUANT_CACHE_SETTING_ID = "SimpleSyrup.QuantCache"; +export const QUANT_CACHE_SETTING_LABEL = "SimpleSyrup: Quantized model cache"; +export const QUANT_CACHE_SETTING_DESCRIPTION = + "Sets the global models/SyrupQuants cache limit in GiB for every SimpleSyrup loader; least-recently-used inactive copies are removed automatically."; + +export interface QuantCacheSettingsApi { + getQuantCacheStatus(): Promise; + enforceQuantCacheLimit(): Promise; + clearQuantCache(): Promise; +} + +interface QuantCacheControlContext { + settings: GeneralSettingsContext; + api: QuantCacheSettingsApi; + logger: Logger; + initialStatus: QuantCacheStatus | null; +} + +export async function registerQuantCacheSetting( + app: ComfyApp, + settings: GeneralSettingsContext, + api: QuantCacheSettingsApi, + logger: Logger +): Promise { + installSimpleSyrupSettingsStyle(); + let initialStatus: QuantCacheStatus | null = null; + try { + initialStatus = await api.getQuantCacheStatus(); + } catch (error) { + logger.warn("Could not load SimpleSyrup quant cache status.", error); + } + app.ui.settings.addSetting({ + id: QUANT_CACHE_SETTING_ID, + name: QUANT_CACHE_SETTING_LABEL, + sortOrder: 321, + type: () => + createQuantCacheControl({ settings, api, logger, initialStatus }), + defaultValue: settings.getSettings().quant_cache_limit_gib, + tooltip: QUANT_CACHE_SETTING_DESCRIPTION + }); +} + +function createQuantCacheControl(context: QuantCacheControlContext): HTMLElement { + const wrapper = createElement("div", "simple-syrup-settings-row"); + const limitInput = createElement("input", "simple-syrup-settings-input"); + limitInput.type = "number"; + limitInput.min = "1"; + limitInput.max = "2048"; + limitInput.step = "1"; + limitInput.value = String(context.settings.getSettings().quant_cache_limit_gib); + limitInput.setAttribute("aria-label", "Quant cache limit in GiB"); + const unitLabel = createElement("span", "simple-syrup-settings-status"); + unitLabel.textContent = "GiB limit"; + + const saveButton = createElement("button", "simple-syrup-settings-button"); + saveButton.type = "button"; + saveButton.textContent = "Save Limit"; + const clearButton = createElement("button", "simple-syrup-settings-button"); + clearButton.type = "button"; + clearButton.textContent = "Clear Inactive"; + const status = createElement("span", "simple-syrup-settings-status"); + + const renderStatus = (cacheStatus: QuantCacheStatus | null): void => { + status.textContent = cacheStatus + ? `${formatGiB(cacheStatus.usage_bytes)} GiB used in ${cacheStatus.path} (${String(cacheStatus.artifact_count)} cached, ${String(cacheStatus.active_artifact_count)} active).` + : "Cache status unavailable. Generated models are stored in models/SyrupQuants."; + }; + renderStatus(context.initialStatus); + + saveButton.addEventListener("click", () => { + void saveLimit(); + }); + limitInput.addEventListener("keydown", (event) => { + if (event.key === "Enter") { + event.preventDefault(); + void saveLimit(); + } + }); + clearButton.addEventListener("click", () => { + void clearInactive(); + }); + + const saveLimit = async (): Promise => { + const limit = Number(limitInput.value); + if (!Number.isInteger(limit) || limit < 1 || limit > 2048) { + status.textContent = "Enter a whole-number cache limit from 1 to 2048 GiB."; + return; + } + const previous = context.settings.getSettings(); + setPending(wrapper, true); + try { + const saved = await context.settings.saveSettings({ + ...previous, + quant_cache_limit_gib: limit + }); + context.settings.setSettings(saved); + limitInput.value = String(saved.quant_cache_limit_gib); + const refreshed = await context.api.enforceQuantCacheLimit(); + renderStatus(refreshed); + } catch (error) { + context.logger.warn("Could not save SimpleSyrup quant cache limit.", error); + limitInput.value = String(previous.quant_cache_limit_gib); + status.textContent = "Cache limit was not saved."; + } finally { + setPending(wrapper, false); + } + }; + + const clearInactive = async (): Promise => { + setPending(wrapper, true); + try { + const cleared = await context.api.clearQuantCache(); + renderStatus(cleared); + } catch (error) { + context.logger.warn("Could not clear SimpleSyrup quant cache.", error); + status.textContent = "Inactive cached models were not cleared."; + } finally { + setPending(wrapper, false); + } + }; + + wrapper.append(limitInput, unitLabel, saveButton, clearButton, status); + return wrapper; +} + +function formatGiB(bytes: number): string { + return (bytes / 1024 ** 3).toFixed(2); +} diff --git a/web/src/settingsRegistration.ts b/web/src/settingsRegistration.ts new file mode 100644 index 0000000..93a07c3 --- /dev/null +++ b/web/src/settingsRegistration.ts @@ -0,0 +1,86 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +import { + clearQuantCache, + enforceQuantCacheLimit, + getExternalLLMSettings, + getQuantCacheStatus, + getSettings, + saveExternalLLMApiKey, + saveExternalLLMSettings, + saveSettings +} from "./api"; +import type { + ExternalLLMSettings, + QuantCacheStatus, + SimpleSyrupSettings +} from "./api"; +import { + registerDownloadableModelsSetting, + type GeneralSettingsContext +} from "./downloadableModelsSetting"; +import { + registerExternalLLMSettings, + type ExternalLLMSettingsApi +} from "./externalLlmSettings"; +import { + registerQuantCacheSetting, + type QuantCacheSettingsApi +} from "./quantCacheSetting"; +import type { ComfyApp, Logger } from "./types"; + +const DEFAULT_SETTINGS: SimpleSyrupSettings = { + show_downloadable_models: true, + quant_cache_limit_gib: 20 +}; + +export interface SimpleSyrupSettingsApi + extends ExternalLLMSettingsApi, + QuantCacheSettingsApi { + getSettings(): Promise; + saveSettings(settings: SimpleSyrupSettings): Promise; +} + +export async function registerSimpleSyrupSettings( + app: ComfyApp, + api: SimpleSyrupSettingsApi = defaultApi(), + logger: Logger = console +): Promise { + let savedSettings = DEFAULT_SETTINGS; + try { + savedSettings = await api.getSettings(); + } catch (error) { + logger.warn( + "Could not load SimpleSyrup settings. Using defaults until the backend is available.", + error + ); + } + + const settingsContext: GeneralSettingsContext = { + getSettings: () => savedSettings, + saveSettings: (settings) => api.saveSettings(settings), + setSettings: (settings) => { + savedSettings = settings; + } + }; + registerDownloadableModelsSetting(app, settingsContext, logger); + await registerQuantCacheSetting(app, settingsContext, api, logger); + await registerExternalLLMSettings(app, api, logger); +} + +function defaultApi(): SimpleSyrupSettingsApi { + return { + getSettings, + saveSettings, + getQuantCacheStatus, + enforceQuantCacheLimit, + clearQuantCache, + getExternalLLMSettings, + saveExternalLLMSettings, + saveExternalLLMApiKey + }; +} + +export type { ExternalLLMSettings, QuantCacheStatus, SimpleSyrupSettings }; diff --git a/web/src/settingsUi.ts b/web/src/settingsUi.ts new file mode 100644 index 0000000..389e681 --- /dev/null +++ b/web/src/settingsUi.ts @@ -0,0 +1,83 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +export function installSimpleSyrupSettingsStyle(): void { + if (document.getElementById("simple-syrup-settings-style")) return; + + const style = document.createElement("style"); + style.id = "simple-syrup-settings-style"; + style.textContent = ` + .simple-syrup-settings-row { + display: flex; + align-items: center; + gap: 0.5rem; + min-width: min(38rem, 100%); + } + .simple-syrup-settings-row[data-pending="true"] { + opacity: 0.75; + } + .simple-syrup-settings-input { + min-width: 16rem; + flex: 1 1 auto; + } + .simple-syrup-settings-button { + flex: 0 0 auto; + white-space: nowrap; + } + .simple-syrup-settings-status { + color: var(--fg-color); + opacity: 0.8; + white-space: normal; + overflow-wrap: anywhere; + } + .simple-syrup-dialog-backdrop { + position: fixed; + inset: 0; + z-index: 2147483647; + display: flex; + align-items: center; + justify-content: center; + background: rgb(0 0 0 / 45%); + } + .simple-syrup-dialog { + display: grid; + gap: 0.75rem; + min-width: min(28rem, calc(100vw - 2rem)); + padding: 1rem; + background: var(--comfy-menu-bg); + color: var(--fg-color); + } + .simple-syrup-dialog-title { + margin: 0; + font-size: 1rem; + } + .simple-syrup-dialog-actions { + display: flex; + justify-content: flex-end; + gap: 0.5rem; + } + `; + document.head.appendChild(style); +} + +export function createElement( + tagName: TTag, + className: string +): HTMLElementTagNameMap[TTag] { + const element = document.createElement(tagName); + element.className = className; + return element; +} + +export function setPending(element: HTMLElement, pending: boolean): void { + element.dataset.pending = pending ? "true" : "false"; + for (const control of Array.from(element.querySelectorAll("input, button"))) { + if ( + control instanceof HTMLInputElement || + control instanceof HTMLButtonElement + ) { + control.disabled = pending; + } + } +} diff --git a/web/tests/api.test.ts b/web/tests/api.test.ts index c5053d7..06ea864 100644 --- a/web/tests/api.test.ts +++ b/web/tests/api.test.ts @@ -5,12 +5,16 @@ import { describe, expect, it, vi } from "vitest"; import { + clearQuantCache, deleteExternalLLMApiKey, getExternalLLMSettings, getMaskBatchPreview, + getQuantCacheStatus, getSettings, + enforceQuantCacheLimit, parseExternalLLMSettings, parseMaskBatchPreview, + parseQuantCacheStatus, parseSettings, refreshExternalLLMModels, saveExternalLLMApiKey, @@ -23,29 +27,45 @@ import { createJsonResponse } from "./testUtils"; describe("settings API", () => { it("loads SimpleSyrup settings from the backend route", async () => { const fetchImpl = vi.fn().mockResolvedValue( - createJsonResponse({ show_downloadable_models: false }) + createJsonResponse({ + show_downloadable_models: false, + quant_cache_limit_gib: 20 + }) ); await expect(getSettings(fetchImpl)).resolves.toEqual({ - show_downloadable_models: false + show_downloadable_models: false, + quant_cache_limit_gib: 20 }); expect(fetchImpl).toHaveBeenCalledWith("/simple-syrup/settings"); }); it("saves SimpleSyrup settings to the backend route", async () => { const fetchImpl = vi.fn().mockResolvedValue( - createJsonResponse({ show_downloadable_models: true }) + createJsonResponse({ + show_downloadable_models: true, + quant_cache_limit_gib: 30 + }) ); await expect( - saveSettings({ show_downloadable_models: true }, fetchImpl) - ).resolves.toEqual({ show_downloadable_models: true }); + saveSettings( + { show_downloadable_models: true, quant_cache_limit_gib: 30 }, + fetchImpl + ) + ).resolves.toEqual({ + show_downloadable_models: true, + quant_cache_limit_gib: 30 + }); expect(fetchImpl).toHaveBeenCalledWith( "/simple-syrup/settings", expect.objectContaining({ method: "POST", headers: { "Content-Type": "application/json" }, - body: JSON.stringify({ show_downloadable_models: true }) + body: JSON.stringify({ + show_downloadable_models: true, + quant_cache_limit_gib: 30 + }) }) ); }); @@ -64,7 +84,10 @@ describe("settings API", () => { .mockResolvedValue(createJsonResponse({ error: "nope" }, { status: 400 })); await expect( - saveSettings({ show_downloadable_models: false }, fetchImpl) + saveSettings( + { show_downloadable_models: false, quant_cache_limit_gib: 20 }, + fetchImpl + ) ).rejects.toThrow("nope"); }); @@ -73,6 +96,36 @@ describe("settings API", () => { "SimpleSyrup settings payload is invalid" ); }); + + it("loads and clears the global quant cache through its dedicated route", async () => { + const status = { + path: "models/SyrupQuants", + usage_bytes: 1024, + limit_bytes: 20 * 1024 ** 3, + artifact_count: 1, + active_artifact_count: 0 + }; + const fetchImpl = vi + .fn() + .mockImplementation(() => Promise.resolve(createJsonResponse(status))); + + await expect(getQuantCacheStatus(fetchImpl)).resolves.toEqual(status); + await expect(enforceQuantCacheLimit(fetchImpl)).resolves.toEqual(status); + await expect(clearQuantCache(fetchImpl)).resolves.toEqual(status); + expect(fetchImpl).toHaveBeenNthCalledWith(1, "/simple-syrup/quant-cache"); + expect(fetchImpl).toHaveBeenNthCalledWith(2, "/simple-syrup/quant-cache", { + method: "POST" + }); + expect(fetchImpl).toHaveBeenNthCalledWith(3, "/simple-syrup/quant-cache", { + method: "DELETE" + }); + }); + + it("rejects malformed quant cache status payloads", () => { + expect(() => + parseQuantCacheStatus({ path: "models/SyrupQuants" }) + ).toThrow("quant cache status is invalid"); + }); }); describe("mask batch preview API", () => { diff --git a/web/tests/settings.test.ts b/web/tests/settings.test.ts index 56036ed..2c3df08 100644 --- a/web/tests/settings.test.ts +++ b/web/tests/settings.test.ts @@ -6,12 +6,17 @@ import { describe, expect, it, vi } from "vitest"; import { SIMPLE_SYRUP_SETTING_ID, - SIMPLE_SYRUP_SETTING_LABEL, + SIMPLE_SYRUP_SETTING_LABEL +} from "../src/downloadableModelsSetting"; +import { EXTERNAL_LLM_API_KEY_SETTING_ID, - EXTERNAL_LLM_ENDPOINT_SETTING_ID, + EXTERNAL_LLM_ENDPOINT_SETTING_ID +} from "../src/externalLlmSettings"; +import { registerSimpleSyrupSettings -} from "../src/settings"; -import type { SimpleSyrupSettingsApi } from "../src/settings"; +} from "../src/settingsRegistration"; +import type { SimpleSyrupSettingsApi } from "../src/settingsRegistration"; +import { QUANT_CACHE_SETTING_ID } from "../src/quantCacheSetting"; import { createFakeComfyApp } from "./testUtils"; describe("Comfy settings registration", () => { @@ -21,7 +26,7 @@ describe("Comfy settings registration", () => { await registerSimpleSyrupSettings(app, api); - expect(app.ui.settings.definitions).toHaveLength(3); + expect(app.ui.settings.definitions).toHaveLength(4); expect(app.ui.settings.definitions[0]).toMatchObject({ id: SIMPLE_SYRUP_SETTING_ID, name: SIMPLE_SYRUP_SETTING_LABEL, @@ -30,25 +35,39 @@ describe("Comfy settings registration", () => { }); expect(app.ui.settings.settings[0]?.value).toBe(false); expect(app.ui.settings.definitions[1]).toMatchObject({ - id: EXTERNAL_LLM_ENDPOINT_SETTING_ID, - sortOrder: 320 + id: QUANT_CACHE_SETTING_ID, + sortOrder: 321 }); expect(typeof app.ui.settings.definitions[1]?.type).toBe("function"); expect(app.ui.settings.definitions[2]).toMatchObject({ + id: EXTERNAL_LLM_ENDPOINT_SETTING_ID, + sortOrder: 320 + }); + expect(typeof app.ui.settings.definitions[2]?.type).toBe("function"); + expect(app.ui.settings.definitions[3]).toMatchObject({ id: EXTERNAL_LLM_API_KEY_SETTING_ID, sortOrder: 319 }); - expect(typeof app.ui.settings.definitions[2]?.type).toBe("function"); + expect(typeof app.ui.settings.definitions[3]?.type).toBe("function"); }); it("saves setting changes to the backend", async () => { const app = createFakeComfyApp(); const saveSettings = vi .fn() - .mockResolvedValue({ show_downloadable_models: true }); + .mockResolvedValue({ + show_downloadable_models: true, + quant_cache_limit_gib: 20 + }); const api: SimpleSyrupSettingsApi = { - getSettings: vi.fn().mockResolvedValue({ show_downloadable_models: false }), + getSettings: vi.fn().mockResolvedValue({ + show_downloadable_models: false, + quant_cache_limit_gib: 20 + }), saveSettings, + getQuantCacheStatus: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), + enforceQuantCacheLimit: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), + clearQuantCache: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), getExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), saveExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), saveExternalLLMApiKey: vi.fn().mockResolvedValue(defaultExternalLLMSettings()) @@ -58,7 +77,8 @@ describe("Comfy settings registration", () => { await app.ui.settings.definitions[0]?.onChange?.(true); expect(saveSettings).toHaveBeenCalledWith({ - show_downloadable_models: true + show_downloadable_models: true, + quant_cache_limit_gib: 20 }); expect(app.ui.settings.settings[0]?.value).toBe(true); }); @@ -68,7 +88,13 @@ describe("Comfy settings registration", () => { const logger = { warn: vi.fn() }; const api: SimpleSyrupSettingsApi = { getSettings: vi.fn().mockRejectedValue(new Error("offline")), - saveSettings: vi.fn().mockResolvedValue({ show_downloadable_models: true }), + saveSettings: vi.fn().mockResolvedValue({ + show_downloadable_models: true, + quant_cache_limit_gib: 20 + }), + getQuantCacheStatus: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), + enforceQuantCacheLimit: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), + clearQuantCache: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), getExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), saveExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), saveExternalLLMApiKey: vi.fn().mockResolvedValue(defaultExternalLLMSettings()) @@ -88,11 +114,20 @@ describe("Comfy settings registration", () => { const logger = { warn: vi.fn() }; const saveSettings = vi .fn() - .mockResolvedValueOnce({ show_downloadable_models: true }) + .mockResolvedValueOnce({ + show_downloadable_models: true, + quant_cache_limit_gib: 20 + }) .mockRejectedValueOnce(new Error("rejected")); const api: SimpleSyrupSettingsApi = { - getSettings: vi.fn().mockResolvedValue({ show_downloadable_models: false }), + getSettings: vi.fn().mockResolvedValue({ + show_downloadable_models: false, + quant_cache_limit_gib: 20 + }), saveSettings, + getQuantCacheStatus: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), + enforceQuantCacheLimit: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), + clearQuantCache: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), getExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), saveExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), saveExternalLLMApiKey: vi.fn().mockResolvedValue(defaultExternalLLMSettings()) @@ -109,6 +144,81 @@ describe("Comfy settings registration", () => { expect(app.ui.settings.settings[0]?.value).toBe(true); }); + it("shows global quant cache usage and saves its GiB limit", async () => { + const app = createFakeComfyApp(); + const api = fakeSettingsApi(true); + const saveSettings = vi + .fn() + .mockResolvedValue({ + show_downloadable_models: true, + quant_cache_limit_gib: 30 + }); + api.saveSettings = saveSettings; + + await registerSimpleSyrupSettings(app, api); + const control = renderSetting(getDefinition(app, 1)); + const input = requiredInput(control, "input[type=number]"); + const saveButton = requiredButton(control, "button"); + + expect(input.value).toBe("20"); + expect(control.textContent).toContain("models/SyrupQuants"); + input.value = "30"; + saveButton.click(); + await flushPromises(); + + expect(saveSettings).toHaveBeenCalledWith({ + show_downloadable_models: true, + quant_cache_limit_gib: 30 + }); + expect(input.value).toBe("30"); + }); + + it("clears inactive quant artifacts and refreshes cache status", async () => { + const app = createFakeComfyApp(); + const api = fakeSettingsApi(true); + const clearQuantCache = vi.fn().mockResolvedValue({ + ...defaultQuantCacheStatus(), + removed_artifacts: 2, + removed_bytes: 1024 + }); + api.clearQuantCache = clearQuantCache; + + await registerSimpleSyrupSettings(app, api); + const control = renderSetting(getDefinition(app, 1)); + const buttons = control.querySelectorAll("button"); + const clearButton = buttons[1]; + if (!(clearButton instanceof HTMLButtonElement)) { + throw new Error("Expected quant cache clear button."); + } + clearButton.click(); + await flushPromises(); + + expect(clearQuantCache).toHaveBeenCalledOnce(); + expect(control.textContent).toContain("0.00 GiB used"); + }); + + it("restores the previous quant limit when backend saving fails", async () => { + const app = createFakeComfyApp(); + const logger = { warn: vi.fn() }; + const api = fakeSettingsApi(true); + api.saveSettings = vi.fn().mockRejectedValue(new Error("rejected")); + + await registerSimpleSyrupSettings(app, api, logger); + const control = renderSetting(getDefinition(app, 1)); + const input = requiredInput(control, "input[type=number]"); + const saveButton = requiredButton(control, "button"); + input.value = "30"; + saveButton.click(); + await flushPromises(); + + expect(input.value).toBe("20"); + expect(control.textContent).toContain("was not saved"); + expect(logger.warn).toHaveBeenCalledWith( + expect.stringContaining("quant cache limit"), + expect.any(Error) + ); + }); + it("saves external LLM endpoint changes to the backend", async () => { const app = createFakeComfyApp(); const refreshComboInNodes = vi.fn().mockResolvedValue(undefined); @@ -122,7 +232,7 @@ describe("Comfy settings registration", () => { api.saveExternalLLMSettings = saveExternalLLMSettings; await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 1)); + const control = renderSetting(getDefinition(app, 2)); const input = requiredInput(control, "input"); const button = requiredButton(control, "button"); @@ -146,7 +256,7 @@ describe("Comfy settings registration", () => { api.saveExternalLLMSettings = saveExternalLLMSettings; await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 1)); + const control = renderSetting(getDefinition(app, 2)); const input = requiredInput(control, "input"); const button = requiredButton(control, "button"); @@ -174,7 +284,7 @@ describe("Comfy settings registration", () => { api.saveExternalLLMApiKey = saveExternalLLMApiKey; await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 2)); + const control = renderSetting(getDefinition(app, 3)); const addButton = requiredButton(control, "button"); expect(addButton.textContent).toBe("Add API Key"); @@ -209,7 +319,7 @@ describe("Comfy settings registration", () => { api.saveExternalLLMSettings = saveExternalLLMSettings; await registerSimpleSyrupSettings(app, api, logger); - const control = renderSetting(getDefinition(app, 1)); + const control = renderSetting(getDefinition(app, 2)); const input = requiredInput(control, "input"); const button = requiredButton(control, "button"); @@ -233,7 +343,7 @@ describe("Comfy settings registration", () => { api.saveExternalLLMApiKey = saveExternalLLMApiKey; await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 2)); + const control = renderSetting(getDefinition(app, 3)); const addButton = requiredButton(control, "button"); addButton.click(); @@ -254,7 +364,7 @@ describe("Comfy settings registration", () => { .mockRejectedValue(new Error("Configure an external LLM endpoint first.")); await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 2)); + const control = renderSetting(getDefinition(app, 3)); const addButton = requiredButton(control, "button"); addButton.click(); @@ -281,7 +391,7 @@ describe("Comfy settings registration", () => { .mockResolvedValue({ ...defaultExternalLLMSettings(), has_api_key: true }); await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 2)); + const control = renderSetting(getDefinition(app, 3)); const button = control.querySelector("button"); expect(button?.textContent).toBe("Replace API Key"); @@ -294,9 +404,13 @@ function fakeSettingsApi( ): SimpleSyrupSettingsApi { return { getSettings: vi.fn().mockResolvedValue({ - show_downloadable_models: showDownloadableModels + show_downloadable_models: showDownloadableModels, + quant_cache_limit_gib: 20 }), saveSettings: vi.fn().mockImplementation((settings) => Promise.resolve(settings)), + getQuantCacheStatus: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), + enforceQuantCacheLimit: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), + clearQuantCache: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), getExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), saveExternalLLMSettings: vi .fn() @@ -307,6 +421,16 @@ function fakeSettingsApi( }; } +function defaultQuantCacheStatus() { + return { + path: "models/SyrupQuants", + usage_bytes: 0, + limit_bytes: 20 * 1024 ** 3, + artifact_count: 0, + active_artifact_count: 0 + }; +} + function defaultExternalLLMSettings() { return { base_url: "",