修改架构,增加可视化

This commit is contained in:
2026-07-23 13:07:49 +08:00
parent e7152727ce
commit f17aa7db8a
27 changed files with 1725 additions and 2278 deletions
+3 -3
View File
@@ -1,4 +1,4 @@
[pytest]
testpaths = . # Run tests in the current directory
python_files = test_*.py # Run tests in files that start with "test_"
norecursedirs = .. # Don't run tests in the parent directory
testpaths = .
python_files = test_*.py
norecursedirs = __pycache__
+185 -15
View File
@@ -1,21 +1,191 @@
#!/usr/bin/env python
"""SmartSaveImage 节点回归测试。"""
"""Tests for `SmartSaveImage` package."""
from pathlib import Path
import pytest
from src.SmartSaveImage.nodes import Example
import torch
from PIL import Image
@pytest.fixture
def example_node():
"""Fixture to create an Example node instance."""
return Example()
pytest.importorskip("folder_paths")
def test_example_node_initialization(example_node):
"""Test that the node can be instantiated."""
assert isinstance(example_node, Example)
from src.SmartSaveImage.nodes import NODE_CLASS_MAPPINGS, SmartSaveImage
from src.SmartSaveImage.nodes.smart_save import (
build_subfolder,
expand_template,
extract_context,
)
from src.SmartSaveImage.server_routes import compute_preview
def test_return_types():
"""Test the node's metadata."""
assert Example.RETURN_TYPES == ("IMAGE",)
assert Example.FUNCTION == "test"
assert Example.CATEGORY == "Example"
def test_node_is_registered():
assert NODE_CLASS_MAPPINGS["SmartSaveImage"] is SmartSaveImage
assert SmartSaveImage.OUTPUT_NODE is True
assert SmartSaveImage.RETURN_TYPES == ()
def test_template_expansion_and_sanitizing():
context = {
"model": "model:name",
"seed": "42",
"positive": "portrait / studio",
"width": "1024",
"height": "768",
}
folder = build_subfolder("%model%/%seed%/%prompt%", context)
assert Path(folder).parts == ("model_name", "42", "portrait _ studio")
def test_extracts_model_name_from_unet_loader():
prompt = {
"1": {
"class_type": "UNETLoader",
"inputs": {"unet_name": "Krea/Krea2_fp8.safetensors"},
}
}
context = extract_context(prompt)
assert context["model"] == "Krea2_fp8"
def test_extract_context_ignores_links_and_reads_sampler():
prompt = {
"1": {"class_type": "CheckpointLoaderSimple",
"inputs": {"ckpt_name": "sdxl/dreamShaperXL.safetensors"}},
"2": {"class_type": "LoraLoader",
"inputs": {"lora_name": "add_detail.safetensors", "model": ["1", 0]}},
"3": {"class_type": "KSampler",
"inputs": {"seed": 999, "steps": 30, "cfg": 6.0,
"sampler_name": "euler", "scheduler": "normal", "model": ["2", 0]}},
}
ctx = extract_context(prompt)
assert ctx["model"] == "dreamShaperXL"
assert ctx["lora"] == "add_detail"
assert ctx["seed"] == "999"
assert ctx["sampler"] == "euler"
def test_collision_mode_adds_counter(tmp_path):
original = tmp_path / "image.png"
original.write_bytes(b"existing")
result = SmartSaveImage._unique_path(str(original), overwrite=False, digits=3)
assert Path(result).name == "image_001.png"
assert original.read_bytes() == b"existing"
def test_preview_uses_next_available_filename(tmp_path):
(tmp_path / "image.png").write_bytes(b"existing")
preview = compute_preview({
"root_mode": "custom",
"custom_root": str(tmp_path),
"folder_template": "",
"filename_template": "image",
"file_format": "png",
"collision_mode": "increment",
"counter_digits": 3,
})
assert preview["example_filenames"] == ["image_001.png"]
def test_batch_token_in_filename():
ctx = extract_context({})
assert expand_template("img_%batch%", ctx, 5) == "img_05"
@pytest.mark.parametrize("file_format,pil_format", [
("png", "PNG"),
("jpeg", "JPEG"),
("webp", "WEBP"),
])
def test_saves_real_image_batches(tmp_path, file_format, pil_format):
images = torch.rand((2, 16, 20, 3))
node = SmartSaveImage()
node.save_images(
images=images,
root_mode="custom",
custom_root=str(tmp_path),
folder_template="case_%model%",
filename_template="sample_%batch%",
file_format=file_format,
quality=91,
collision_mode=SmartSaveImage.COLLISION_INCREMENT,
save_mode=SmartSaveImage.MODE_SAVE_ONLY,
manual_model="auto",
embed_workflow=True,
counter_digits=3,
prompt={
"1": {
"class_type": "CheckpointLoaderSimple",
"inputs": {"ckpt_name": "models/demo.safetensors"},
}
},
extra_pnginfo={"workflow": {"nodes": []}},
)
ext = SmartSaveImage._extension(file_format).lstrip(".")
saved = sorted(tmp_path.rglob(f"*.{ext}"))
assert [path.stem for path in saved] == ["sample_00", "sample_01"]
assert {Image.open(path).format for path in saved} == {pil_format}
def test_overwrite_batch_with_zero_digits_keeps_every_image(tmp_path):
images = torch.rand((2, 8, 8, 3))
SmartSaveImage().save_images(
images=images,
root_mode="custom",
custom_root=str(tmp_path),
folder_template="",
filename_template="image",
file_format="png",
quality=95,
collision_mode=SmartSaveImage.COLLISION_OVERWRITE,
save_mode=SmartSaveImage.MODE_SAVE_ONLY,
manual_model="auto",
embed_workflow=False,
counter_digits=0,
)
assert sorted(path.name for path in tmp_path.glob("*.png")) == ["image_0.png", "image_1.png"]
def test_png_compression_defaults_to_comfyui_save_image_level():
config = SmartSaveImage.INPUT_TYPES()["optional"]["png_compression"][1]
assert config["default"] == 4
assert config["min"] == 0
assert config["max"] == 9
def test_png_compression_changes_size_without_changing_pixels(tmp_path):
image = torch.zeros((1, 64, 64, 3))
node = SmartSaveImage()
common = {
"images": image,
"root_mode": "custom",
"custom_root": str(tmp_path),
"folder_template": "",
"file_format": "png",
"quality": 95,
"collision_mode": SmartSaveImage.COLLISION_OVERWRITE,
"save_mode": SmartSaveImage.MODE_SAVE_ONLY,
"embed_workflow": False,
}
node.save_images(filename_template="level_0", png_compression=0, **common)
node.save_images(filename_template="level_4", png_compression=4, **common)
level_0 = tmp_path / "level_0.png"
level_4 = tmp_path / "level_4.png"
with Image.open(level_0) as first, Image.open(level_4) as second:
assert first.tobytes() == second.tobytes()
assert level_4.stat().st_size < level_0.stat().st_size