Spaces:
Running on Zero
Running on Zero
Download tests/test_gpu_boundary.py from BlueSkyXN/ImageGen-Studio: direct link, hf CLI and curl.
- Browser
- Download file 8.51 kB
-
https://huggingface.co/spaces/BlueSkyXN/ImageGen-Studio/resolve/main/tests/test_gpu_boundary.py
- Command line
-
hf download hf://spaces/BlueSkyXN/ImageGen-Studio/tests/test_gpu_boundary.py
-
curl -L -o test_gpu_boundary.py https://huggingface.co/spaces/BlueSkyXN/ImageGen-Studio/resolve/main/tests/test_gpu_boundary.py
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() | |