ImageGen-Studio / tests /test_gpu_boundary.py
BlueSkyXN's picture
Unify image references and guided UI/API workflows
ca9a89c verified
Raw History Blame Contribute Delete
8.51 kB
from __future__ import annotations
import importlib.util
import pickle
import sys
import tempfile
import threading
import types
import unittest
from pathlib import Path
from unittest import mock
from core.task_scheduler import TaskCancelledError
ROOT = Path(__file__).resolve().parents[1]
class _LockedProgress:
def __init__(self, *args, **kwargs):
self.lock = threading.Lock()
self.updates = []
def __call__(self, value, desc=None):
self.updates.append((value, desc))
class _Assembler:
def __init__(self, after_assembly=None):
self.after_assembly = after_assembly
def assemble(self, values):
self.ui_values = values
if self.after_assembly:
self.after_assembly()
return {"1": {"class_type": "TestNode", "inputs": {"seed": values["seed"]}}}
def _module(name, **attributes):
module = types.ModuleType(name)
module.__dict__.update(attributes)
return module
class GpuBoundaryTests(unittest.TestCase):
"""执行真实 pipeline 调用链,只替换下载、推理依赖和 GPU 传输。"""
def setUp(self):
output_dir = self.enterContext(tempfile.TemporaryDirectory())
self.payloads = []
self.assemblers = []
self.after_assembly = None
self.download = mock.Mock()
self.release_models = mock.Mock()
self.execute_workflow = mock.Mock(
return_value=types.SimpleNamespace(shape=(0,))
)
def gpu_decorator(**_options):
def decorate(function):
def run(*args, **kwargs):
# 真正序列化入参,不能像 UI 烟测那样直接跳过生成调用。
args, kwargs = pickle.loads(pickle.dumps((args, kwargs)))
self.payloads.append(kwargs)
return function(*args, **kwargs)
return run
return decorate
def create_assembler(*_args, **_kwargs):
assembler = _Assembler(self.after_assembly)
self.assemblers.append(assembler)
return assembler
prepared_keys = (
"temp_files_to_clean", "active_loras_for_gpu", "active_loras_for_meta",
"active_controlnets", "active_anima_controlnets", "active_diffsynth_controlnets",
"active_ipadapters", "active_flux1_ipadapters", "active_sd3_ipadapters",
"active_styles", "active_reference_latents", "active_hidream_o1_reference",
"active_conditioning",
)
base_name = "core.pipelines.base_pipeline"
pipeline_name = "core.pipelines.sd_image_pipeline"
stubs = {
"gradio": _module("gradio", Error=RuntimeError, Progress=_LockedProgress),
"spaces": _module("spaces", GPU=gpu_decorator),
"torch": _module("torch"),
"imageio": _module("imageio"),
"numpy": _module("numpy"),
"core.model_manager": _module(
"core.model_manager",
model_manager=types.SimpleNamespace(ensure_models_downloaded=self.download),
release_loaded_models=self.release_models,
),
"core.workflow_assembler": _module(
"core.workflow_assembler", WorkflowAssembler=create_assembler
),
"imagegen_utils.app_utils": _module(
"imagegen_utils.app_utils", sanitize_prompt=lambda value: value
),
"core.pipelines.pipeline_input_processor": _module(
"core.pipelines.pipeline_input_processor",
process_pipeline_inputs=lambda *_: {key: [] for key in prepared_keys},
),
"core.pipelines.workflow_executor": _module(
"core.pipelines.workflow_executor",
WorkflowExecutor=types.SimpleNamespace(execute_workflow=self.execute_workflow),
),
}
specs = []
for name in (base_name, pipeline_name):
spec = importlib.util.spec_from_file_location(
name, ROOT / "core" / "pipelines" / f"{name.rsplit('.', 1)[1]}.py"
)
stubs[name] = importlib.util.module_from_spec(spec)
specs.append(spec)
self.enterContext(mock.patch.dict(sys.modules, stubs))
for spec in specs:
spec.loader.exec_module(stubs[spec.name])
self.module = stubs[pipeline_name]
self.module.OUTPUT_DIR = output_dir
self.pipeline = self.module.SdImagePipeline()
def inputs(self, **overrides):
values = {
"task_type": "txt2img",
"model_display_name": "circlestone-labs/Anima-Turbo-v1.0",
"positive_prompt": "测试图片",
"negative_prompt": "",
"seed": 42,
"batch_size": 1,
"num_inference_steps": 4,
"guidance_scale": 1.0,
"sampler": "euler",
"scheduler": "simple",
"width": 64,
"height": 64,
"denoise": 1.0,
}
values.update(overrides)
return values
def test_repeated_ui_calls_serialize_and_preserve_model_release(self):
for release in (True, False):
with self.subTest(release=release):
cancellation = threading.Event()
values = self.inputs(
_cancel_event=cancellation, _release_models_after_run=release
)
progress = _LockedProgress()
self.assertEqual(self.pipeline.run(values, progress), [])
payload = self.payloads[-1]
self.assertEqual(set(payload), {"ui_inputs", "loras_string", "workflow"})
self.assertNotIn("_cancel_event", payload["ui_inputs"])
self.assertEqual(payload["ui_inputs"]["_release_models_after_run"], release)
self.assertIs(values["_cancel_event"], cancellation)
self.assertIs(self.assemblers[-1].ui_values["_cancel_event"], cancellation)
self.assertTrue(progress.updates)
self.assertEqual(len(self.payloads), 2)
self.release_models.assert_called_once_with()
def test_request_without_cancel_event_serializes(self):
self.pipeline.run(self.inputs(), _LockedProgress())
self.assertNotIn("_cancel_event", self.payloads[0]["ui_inputs"])
self.release_models.assert_not_called()
def test_unsupported_stale_reference_does_not_block_img2img(self):
from PIL import Image
self.pipeline.run(self.inputs(task_type="img2img", img2img_image=Image.new("RGB", (64, 64)), qwen_image_edit_data=[Image.new("RGB", (32, 32))], positive_prompt="保留图1"), _LockedProgress())
self.download.assert_called_once()
self.assertEqual(self.payloads[0]["ui_inputs"]["qwen_image_edit_data"], [])
self.assertEqual(self.payloads[0]["ui_inputs"]["positive_prompt"], "保留image 1")
def test_invalid_numbered_reference_rejected_before_download_and_gpu(self):
with self.assertRaisesRegex(RuntimeError, "本次只有 0"):
self.pipeline.run(self.inputs(positive_prompt="img2"), _LockedProgress())
self.download.assert_not_called()
self.assertEqual(self.payloads, [])
def test_cancellation_after_download_stays_on_cpu(self):
cancellation = threading.Event()
self.download.side_effect = lambda *args, **kwargs: cancellation.set()
with self.assertRaises(TaskCancelledError):
self.pipeline.run(self.inputs(_cancel_event=cancellation), _LockedProgress())
self.assertEqual(self.payloads, [])
self.execute_workflow.assert_not_called()
def test_cancellation_after_assembly_stays_on_cpu(self):
cancellation = threading.Event()
self.after_assembly = cancellation.set
with self.assertRaises(TaskCancelledError):
self.pipeline.run(self.inputs(_cancel_event=cancellation), _LockedProgress())
self.assertEqual(self.payloads, [])
self.execute_workflow.assert_not_called()
def test_gpu_failure_still_releases_model_state(self):
self.execute_workflow.side_effect = RuntimeError("sampling failed")
with self.assertRaisesRegex(RuntimeError, "sampling failed"):
self.pipeline.run(
self.inputs(_cancel_event=threading.Event()), _LockedProgress()
)
self.release_models.assert_called_once_with()
if __name__ == "__main__":
unittest.main()