Files
spawner1145-comfyui-lsnet/tests/test_model_loading.py
T
2026-10-06 00:30:52 +08:00

395 lines
20 KiB
Python

"""Functional tests with real, small DINOv3 networks and local checkpoints."""
import base64
import importlib.util
import io
import json
from pathlib import Path
import sys
import tempfile
import types
import unittest
from unittest.mock import patch
import torch
from PIL import Image
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model_loading import load_model_bundle, find_checkpoint, model_folders
from inference_artist import classify_image, resolved_mode
from kaloscope_dinov3.models.vision_transformer import DinoVisionTransformer
class ReferenceInferenceModel(torch.nn.Module):
"""Independent forward-only fixture for comparing loaded checkpoint outputs."""
def __init__(self, config, num_classes):
super().__init__()
self.backbone = DinoVisionTransformer(**config["kwargs"])
self.backbone.init_weights()
self.feature_dim = 2 * self.backbone.embed_dim
self.head = torch.nn.Linear(self.feature_dim, num_classes)
self.projector = torch.nn.Sequential(
torch.nn.Linear(self.feature_dim, self.feature_dim), torch.nn.GELU(),
torch.nn.Linear(self.feature_dim, config["projection_dim"]),
)
def forward(self, images, projections=False):
tokens = self.backbone.forward_features(images)
features = torch.cat((tokens["x_norm_clstoken"], tokens["x_norm_patchtokens"].mean(1)), dim=1)
result = {"features": features, "logits": self.head(features)}
if projections:
result["projections"] = self.projector(features)
return result
def load_nodes():
fake = types.ModuleType("folder_paths")
fake.models_dir = str(ROOT / "models")
spec = importlib.util.spec_from_file_location("kaloscope_test_nodes", ROOT / "__init__.py")
module = importlib.util.module_from_spec(spec)
original = sys.modules.get("folder_paths")
sys.modules["folder_paths"] = fake
try:
spec.loader.exec_module(module)
finally:
if original is None:
del sys.modules["folder_paths"]
else:
sys.modules["folder_paths"] = original
return module
class ModelLoadingTests(unittest.TestCase):
@classmethod
def setUpClass(cls):
torch.set_num_threads(2)
cls.options = {"name": "custom_vit", "pooling": "cls_mean", "projection_dim": 8,
"kwargs": {"embed_dim": 24, "depth": 2, "num_heads": 3, "patch_size": 16,
"n_storage_tokens": 2,
"pos_embed_rope_dtype": "fp32"}}
torch.manual_seed(1)
cls.original = ReferenceInferenceModel(cls.options, num_classes=3).eval()
cls.nodes = load_nodes()
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.directory = Path(self.temp.name)
self.config = {"model": "custom_vit", "input_size": 32}
self.payload = {"model": self.original.state_dict(), "model_config": self.options,
"classes": ["a", "b", "c"]}
self.save()
def tearDown(self):
self.temp.cleanup()
def save(self):
(self.directory / "config.json").write_text(json.dumps(self.config), encoding="utf-8")
torch.save(self.payload, self.directory / "best.pt")
def bundle(self):
return load_model_bundle(self.directory, device="cpu")
def tensor(self, bundle):
return bundle["transform"](Image.new("RGB", (48, 32), (70, 130, 190))).unsqueeze(0)
def test_head_and_features_match_reference_model(self):
bundle = self.bundle()
batch = self.tensor(bundle)
with torch.inference_mode():
expected = self.original(batch)
torch.testing.assert_close(bundle["model"](batch), expected["logits"])
torch.testing.assert_close(bundle["model"](batch, return_features=True), expected["features"])
self.assertTrue(bundle["has_classifier"])
self.assertEqual(bundle["feature_dim"], 48)
results = classify_image(bundle["model"], batch, "cpu", bundle["class_mapping"], top_k=20)
self.assertEqual(len(results), 3)
self.assertEqual({row["class_name"] for row in results}, {"a", "b", "c"})
self.assertAlmostEqual(sum(row["probability"] for row in results), 1.0, places=6)
def test_no_head_is_feature_only_even_with_metadata_classes(self):
self.payload["model"] = {k: v for k, v in self.payload["model"].items() if not k.startswith("head.")}
self.save()
bundle = self.bundle()
self.assertFalse(bundle["has_classifier"])
self.assertEqual(resolved_mode(bundle["model"], "auto"), "cluster")
with self.assertRaisesRegex(ValueError, "no classification head"):
classify_image(bundle["model"], self.tensor(bundle), "cpu")
with torch.inference_mode():
torch.testing.assert_close(bundle["model"](self.tensor(bundle), return_features=True),
self.original(self.tensor(bundle))["features"])
def test_temporal_projector_is_loaded_and_used(self):
self.payload["model"] = {k: v for k, v in self.payload["model"].items() if not k.startswith("head.")}
self.payload["model"].update(log_temperature=torch.tensor(1.0), bias=torch.tensor(-1.0))
self.save()
bundle = self.bundle()
self.assertEqual(bundle["feature_source"], "projector")
self.assertEqual(bundle["feature_dim"], 8)
with torch.inference_mode():
expected = self.original(self.tensor(bundle), projections=True)["projections"]
torch.testing.assert_close(bundle["model"](self.tensor(bundle), return_features=True), expected)
def test_raw_backbone_and_official_linear_head(self):
self.payload = dict(self.original.backbone.state_dict())
self.payload.update({"linear_head.weight": self.original.head.weight,
"linear_head.bias": self.original.head.bias})
self.config["model"] = {"name": "custom_vit", "kwargs": self.options["kwargs"]}
self.save()
bundle = self.bundle()
self.assertEqual(bundle["model"].pooling, "cls_mean")
with torch.inference_mode():
torch.testing.assert_close(bundle["model"](self.tensor(bundle)), self.original(self.tensor(bundle))["logits"])
self.assertTrue(classify_image(bundle["model"], self.tensor(bundle), "cpu")[0]["class_name"].startswith("Class "))
def test_raw_headless_safetensors(self):
from safetensors.torch import save_file
self.config["model"] = {"name": "custom_vit", "kwargs": self.options["kwargs"]}
self.config["checkpoint"] = "backbone.safetensors"
self.save()
save_file(self.original.backbone.state_dict(), self.directory / "backbone.safetensors")
bundle = self.bundle()
self.assertFalse(bundle["has_classifier"])
self.assertEqual(bundle["feature_dim"], 24)
with torch.inference_mode():
expected = self.original.backbone.forward_features(self.tensor(bundle))["x_norm_clstoken"]
torch.testing.assert_close(bundle["model"](self.tensor(bundle), return_features=True), expected)
def test_wrapped_state_dict_prefixes(self):
self.payload = {"state_dict": {"module._orig_mod." + k: v for k, v in self.payload["model"].items()},
"model_config": self.options}
self.save()
self.assertTrue(self.bundle()["has_classifier"])
def test_checkpoint_transform_path_works_without_finetune_module(self):
from kaloscope_dinov3.preprocessing import lvd_transform
image = Image.new("RGB", (48, 32), "blue")
for prefix in ("dinov3.finetune.data.", "kaloscope_dinov3.finetune.data."):
self.config["data"] = {"custom_transform": prefix + "lvd_transform"}
self.save()
torch.testing.assert_close(self.bundle()["transform"](image), lvd_transform(32)(image))
def test_all_final_feature_outputs_match_backbone(self):
bundle = self.bundle()
model = bundle['model']
batch = self.tensor(bundle)
with torch.inference_mode():
raw = self.original.backbone.forward_features(batch)
cls, patches, storage = raw['x_norm_clstoken'], raw['x_norm_patchtokens'], raw['x_storage_tokens']
pooled = torch.cat((cls, patches.mean(1)), dim=1)
expected = {
'default': pooled, 'backbone': pooled, 'cls': cls, 'mean': patches.mean(1),
'cls_mean': pooled, 'patch_tokens': patches, 'storage_tokens': storage,
'all_tokens': torch.cat((cls.unsqueeze(1), storage, patches), dim=1),
'prenorm': raw['x_prenorm'], 'projector': self.original.projector(pooled),
'patch_map': patches.transpose(1, 2).reshape(1, 24, 2, 2),
}
for kind, tensor in expected.items():
with self.subTest(output=kind):
torch.testing.assert_close(model.extract_tensor(batch, kind), tensor)
# Loading backbone as default must still retain trained projector for selection.
self.assertEqual(bundle['feature_source'], 'backbone')
self.assertIsNotNone(model.projector)
def test_intermediate_features_keep_layer_order_and_shape(self):
model = self.bundle()['model']
batch = self.tensor(self.bundle())
with torch.inference_mode():
native = self.original.backbone.get_intermediate_layers(
batch, n=[0, 1], return_class_token=True, return_extra_tokens=True)
for kind in ('cls', 'mean', 'cls_mean', 'patch_tokens', 'patch_map', 'storage_tokens', 'all_tokens'):
expected = []
for patches, cls, storage in reversed(native):
if kind == 'cls':
tensor = cls
elif kind == 'mean':
tensor = patches.mean(1)
elif kind == 'cls_mean':
tensor = torch.cat((cls, patches.mean(1)), dim=1)
elif kind == 'patch_tokens':
tensor = patches
elif kind == 'patch_map':
tensor = patches.transpose(1, 2).reshape(1, 24, 2, 2)
elif kind == 'storage_tokens':
tensor = storage
else:
tensor = torch.cat((cls.unsqueeze(1), storage, patches), dim=1)
expected.append(tensor)
with self.subTest(output=kind):
actual = model.extract_tensor(batch, 'intermediate_' + kind, '-1,0')
torch.testing.assert_close(actual, torch.stack(expected, dim=1))
unnorm = self.original.backbone.get_intermediate_layers(
batch, n=[0, 1], norm=False, return_class_token=True, return_extra_tokens=True)
expected = torch.stack([torch.cat((cls.unsqueeze(1), storage, patches), dim=1)
for patches, cls, storage in unnorm], dim=1)
torch.testing.assert_close(model.extract_tensor(batch, 'intermediate_prenorm', '0,1'), expected)
torch.testing.assert_close(model.extract_tensor(batch, 'intermediate_all_tokens', '0,1', False), expected)
self.assertEqual(tuple(model.extract_tensor(batch, 'intermediate_patch_tokens').shape), (1, 1, 4, 24))
def test_feature_selection_does_not_change_classifier(self):
bundle = self.bundle()
model = bundle['model']
batch = self.tensor(bundle)
with torch.inference_mode():
before = model(batch)
for kind in ('mean', 'projector', 'patch_tokens', 'cls'):
model.extract_tensor(batch, kind)
torch.testing.assert_close(model(batch), before)
self.assertEqual(model.pooling, 'cls_mean')
self.assertEqual(model.feature_source, 'backbone')
def test_missing_projector_and_invalid_layers_fail_clearly(self):
self.payload['model'] = {k: v for k, v in self.payload['model'].items() if not k.startswith('projector.')}
self.save()
bundle = self.bundle()
batch = self.tensor(bundle)
with self.assertRaisesRegex(ValueError, 'no supported projector'):
bundle['model'].extract_tensor(batch, 'projector')
for layers in ('2', '-3', '0,0', 'one', ''):
with self.subTest(layers=layers), self.assertRaises(ValueError):
bundle['model'].extract_tensor(batch, 'intermediate_cls', layers)
def test_feature_node_returns_selected_tensor_for_image_batch(self):
bundle = self.bundle()
image = torch.full((2, 32, 48, 3), 0.5)
node = self.nodes.KaloscopeExtractFeaturesNode()
for kind, shape in {
'cls': (2, 24), 'mean': (2, 24), 'cls_mean': (2, 48), 'projector': (2, 8),
'patch_tokens': (2, 4, 24), 'patch_map': (2, 24, 2, 2),
'storage_tokens': (2, 2, 24), 'all_tokens': (2, 7, 24), 'prenorm': (2, 7, 24),
'intermediate_patch_tokens': (2, 2, 4, 24),
'intermediate_patch_map': (2, 2, 24, 2, 2),
}.items():
with self.subTest(output=kind):
result = node.extract(image, bundle, kind, '0,1')[0]
self.assertEqual(tuple(result.shape), shape)
self.assertEqual(result.device.type, 'cpu')
self.assertFalse(result.requires_grad)
self.assertTrue(torch.isfinite(result).all())
def test_convnext_feature_outputs_and_variable_stages(self):
from kaloscope_dinov3.models.convnext import ConvNeXt
from model_loading import DinoInferenceModel
backbone = ConvNeXt(depths=[1, 1, 1, 1], dims=[8, 16, 24, 32]).eval()
model = DinoInferenceModel(backbone, 'cls').eval()
images = torch.rand(2, 3, 64, 96)
with torch.inference_mode():
raw = backbone.forward_features(images)
torch.testing.assert_close(model.extract_tensor(images, 'patch_tokens'), raw['x_norm_patchtokens'])
patch_map = model.extract_tensor(images, 'patch_map')
self.assertEqual(tuple(patch_map.shape), (2, 32, 2, 3))
torch.testing.assert_close(patch_map.flatten(2).transpose(1, 2), raw['x_norm_patchtokens'])
intermediate = model.extract_tensor(images, 'intermediate_patch_map', '1')
self.assertEqual(tuple(intermediate.shape), (2, 1, 16, 8, 12))
torch.testing.assert_close(model.extract_tensor(images, 'intermediate_prenorm', '-1')[:, 0], raw['x_prenorm'])
with self.assertRaisesRegex(ValueError, 'no storage/register tokens'):
model.extract_tensor(images, 'storage_tokens')
with self.assertRaisesRegex(ValueError, 'different tensor shapes'):
model.extract_tensor(images, 'intermediate_patch_tokens', '0,1')
def test_lsnet_feature_node_preserves_default_and_rejects_dino_outputs(self):
encoder = torch.nn.Identity()
encoder.forward = lambda batch, return_features: batch.mean((2, 3))
from kaloscope_dinov3.preprocessing import image_transform
bundle = {'model': encoder, 'transform': image_transform(32), 'device': 'cpu'}
node = self.nodes.KaloscopeExtractFeaturesNode()
image = torch.ones(1, 32, 32, 3)
self.assertEqual(tuple(node.extract(image, bundle)[0].shape), (1, 3))
with self.assertRaisesRegex(ValueError, 'requires a DINOv3'):
node.extract(image, bundle, 'patch_tokens')
def test_bad_architecture_is_not_silently_ignored(self):
self.config["model"] = "dinov3_typo"
self.save()
with self.assertRaises(ValueError):
self.bundle()
def test_incomplete_backbone_fails(self):
del self.payload["model"]["backbone.cls_token"]
self.save()
with self.assertRaisesRegex(RuntimeError, "cls_token"):
self.bundle()
def test_class_mapping_must_match_head(self):
(self.directory / "class_mapping.csv").write_text("class_id,class_name\n0,one\n2,two\n", encoding="utf-8")
with self.assertRaisesRegex(ValueError, "mapping IDs"):
self.bundle()
def test_model_discovery_only_uses_kaloscope(self):
for folder in ("lsnet/shared", "kaloscope/shared", "lsnet/other"):
(self.directory / folder).mkdir(parents=True)
folders = model_folders(self.directory)
self.assertEqual(folders, {"shared": self.directory / "kaloscope/shared"})
def test_architecture_must_be_explicit(self):
del self.config['model']
self.save()
with self.assertRaisesRegex(ValueError, 'model architecture in config.json'):
self.bundle()
bundle = load_model_bundle(self.directory, device='cpu', model_name='custom_vit')
self.assertEqual(bundle['model_type'], 'custom_vit')
def test_ambiguous_checkpoint_requires_selection(self):
(self.directory / "best.pt").rename(self.directory / "one.pt")
torch.save(self.payload, self.directory / "two.pth")
with self.assertRaisesRegex(ValueError, "Expected one checkpoint"):
find_checkpoint(self.directory)
self.config["checkpoint"] = "two.pth"
self.save()
self.assertEqual(find_checkpoint(self.directory).name, "two.pth")
def test_kaloscope_nodes_and_model_sockets(self):
bundle = self.bundle()
nodes = self.nodes
self.assertEqual(len(nodes.NODE_CLASS_MAPPINGS), 10)
self.assertEqual(set(nodes.NODE_CLASS_MAPPINGS), set(nodes.NODE_DISPLAY_NAME_MAPPINGS))
self.assertEqual(nodes.KaloscopeModelLoader.RETURN_TYPES, ("KALOSCOPE_MODEL",))
for name, node in nodes.NODE_CLASS_MAPPINGS.items():
schema = node.INPUT_TYPES()
self.assertTrue(name.startswith('Kaloscope'))
if "model" in schema.get("required", {}):
self.assertEqual(schema["required"]["model"][0], 'KALOSCOPE_MODEL')
self.assertTrue(node.__name__.startswith("Kaloscope"))
self.assertIn(node.CATEGORY, ("Kaloscope", "Kaloscope/Analysis"))
with patch.object(nodes, "model_folders", return_value={"sample": self.directory}):
loaded = nodes.KaloscopeModelLoader().load("sample", "cpu")[0]
self.assertTrue(loaded["has_classifier"])
image = torch.full((2, 32, 48, 3), 0.5)
features = nodes.KaloscopeExtractFeaturesNode().extract(image, bundle)[0]
self.assertEqual(tuple(features.shape), (2, 48))
tags, predictions = nodes.KaloscopeArtistInferenceNode().process(image, bundle, 5, 0.0)
self.assertEqual(set(tags.split(",")), {"a", "b", "c"})
self.assertEqual(len(json.loads(predictions)), 3)
result, visualization = nodes.KaloscopeClusteringNode().cluster(
"kmeans", 2, 0.5, 2, True, "pca", 5, group_1=torch.randn(4, 48))
self.assertEqual(len(json.loads(result)["labels"]), 4)
self.assertEqual(visualization.shape[-1], 3)
def test_backend_auto_features_and_api(self):
from backend_lsnet.inference import process_image_from_pil
from backend_lsnet.api import Api
from fastapi import FastAPI
from fastapi.testclient import TestClient
self.payload["model"] = {k: v for k, v in self.payload["model"].items() if not k.startswith("head.")}
self.save()
image = Image.new("RGB", (48, 32), "blue")
result = process_image_from_pil(image, checkpoint=str(self.directory / "best.pt"), device="cpu")
self.assertEqual(len(result["features"]), 48)
buffer = io.BytesIO()
image.save(buffer, format="PNG")
app = FastAPI()
api = Api(app)
with patch("backend_lsnet.api.get_available_checkpoints", return_value=["best.pt"]), \
patch("backend_lsnet.api.get_checkpoint_path", return_value=str(self.directory / "best.pt")), \
patch("backend_lsnet.api.get_class_csv", return_value=None), TestClient(app) as client:
response = client.post('/kaloscope/v1/infer', json={"input_image": base64.b64encode(buffer.getvalue()).decode(), "device": "cpu"})
self.assertEqual(response.status_code, 200, response.text)
self.assertEqual(len(response.json()["results"]["features"]), 48)
self.assertTrue(all(route.path.startswith('/kaloscope/v1/') for route in app.routes
if hasattr(route, 'endpoint') and route.endpoint.__module__ == 'backend_lsnet.api'))
api.executor.shutdown()
if __name__ == "__main__":
unittest.main()