Coverage for tests/test_wrapper_flux.py: 100%
78 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1#!/usr/bin/env python3
3import sys
4import pytest
6from typing import Any
8from unittest.mock import patch
9from unittest.mock import MagicMock
10from tests.torch_mock import TorchMock
11from tests.diffusers_mock import DiffusersMock
13from PIL import Image
15mock_torch = TorchMock()
16mock_diffusers = DiffusersMock()
18sys.path.append("wrapper")
19sys.path.append("wrapper/flux")
21mock_modules = {
22 'torch': mock_torch,
23 'nvidia_smi': MagicMock(),
24 'imageio': MagicMock(),
25 'cv2': MagicMock(),
26 'xfuser': MagicMock(),
27 'xfuser.config': MagicMock(),
28 'xfuser.core': MagicMock(),
29 'xfuser.core.distributed': MagicMock(),
30 'xfuser.model_executor': MagicMock(),
31 'xfuser.model_executor.models': MagicMock(),
32 'xfuser.model_executor.models.transformers.transformer_flux': MagicMock(),
33 'xfuser.model_executor.layers': MagicMock(),
34 'xfuser.model_executor.layers.attention_processor': MagicMock(),
35}
36mock_modules.update(mock_torch.get_sub_modules())
37mock_modules.update(mock_diffusers.get_sub_modules())
39with patch.dict(sys.modules, mock_modules):
40 from flux.wrapper_flux import FluxGeneration
41 _flux_module = sys.modules['flux.wrapper_flux']
44@pytest.mark.asyncio
45async def test_wrapper_flux() -> None:
46 model = FluxGeneration()
47 assert model is not None
48 assert model.model_name == "flux"
49 assert model.status == "initializing"
51 with pytest.raises(ValueError, match="Model not initialized. Current status: initializing."):
52 await model.generate(64, 48, "test prompt")
54 model.init()
55 assert model.status == "ok"
57 health = model.get_health()
58 assert health is not None
59 timestamps = model.get_timestamps()
60 assert timestamps is not None
62 with pytest.raises(ValueError, match="Missing JSON body"):
63 await model.get_rest_args(None)
64 with pytest.raises(ValueError, match="Missing 'prompt' parameter"):
65 await model.get_rest_args({})
66 await model.get_rest_args({
67 "job_id": "unittest",
68 "prompt": "Test prompt",
69 "width": 80,
70 "height": 60,
71 "seed": 7,
72 })
74 await model.warmup()
76 image = await model.generate(
77 width=1024,
78 height=1024,
79 prompt="Test prompt")
80 assert image is not None
81 assert isinstance(image, Image.Image)
82 assert image.size == (1024, 1024)
84 # 48x48 not supported for 2 GPUs (latent shape 9, odd).
85 model.world_size = 2
86 with pytest.raises(ValueError, match="48x48 not supported for 2 GPUs"):
87 await model.generate(
88 width=48,
89 height=48,
90 prompt="Test prompt")
92 del model
95@pytest.mark.asyncio
96async def test_additional_coverage() -> None:
97 """Cover seed path, step callbacks, parallelism init, and compile-disabled path."""
98 model = FluxGeneration()
99 model.init()
100 assert model.status == "ok"
102 image = await model.generate(
103 width=256,
104 height=256,
105 prompt="Seed coverage test",
106 seed=42)
107 assert isinstance(image, Image.Image)
109 pipeline_instance = model.pipeline
111 def _pipeline_with_callback(*args: Any, **kwargs: Any) -> Any:
112 n_steps = kwargs.get("num_inference_steps", 2)
113 callback = kwargs.get("callback_on_step_end")
114 if callback:
115 for step in range(n_steps):
116 callback(pipeline_instance, step, 0, {})
117 out = MagicMock()
118 out.images = [Image.new("RGB", (kwargs.get("width", 64), kwargs.get("height", 64)))]
119 return out
121 pipeline_instance.side_effect = _pipeline_with_callback
122 image = await model.generate(
123 width=256,
124 height=256,
125 prompt="Callback coverage test",
126 sampling_steps=2)
127 assert isinstance(image, Image.Image)
129 model.world_size = 2
130 with patch.object(_flux_module, 'parallelize_transformer'):
131 model.init_model_parallelism()
133 model.torch_compile = False
134 model.model_compile()
136 del model
139def test_model_compile_no_pipeline() -> None:
140 """model_compile() with pipeline=None raises ValueError."""
141 model = FluxGeneration()
142 assert model.pipeline is None
143 with pytest.raises(ValueError, match="FLUX pipeline not initialized"):
144 model.model_compile()
145 del model