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

287 lines
11 KiB
Python

# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""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