Coverage for tests/test_wrapper_fluxkrea.py: 100%
80 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
13mock_torch = TorchMock()
14mock_diffusers = DiffusersMock()
16sys.path.append("fluxkrea")
17sys.path.append("flux")
19mock_modules = {
20 'nvidia_smi': MagicMock(),
21 'imageio': MagicMock(),
22 'cv2': MagicMock(),
23 'torch': mock_torch,
24 'xfuser': MagicMock(),
25 'xfuser.config': MagicMock(),
26 'xfuser.core': MagicMock(),
27 'xfuser.core.distributed': MagicMock(),
28 'xfuser.model_executor': MagicMock(),
29 'xfuser.model_executor.models': MagicMock(),
30 'xfuser.model_executor.models.transformers.transformer_flux': MagicMock(),
31 'xfuser.model_executor.layers': MagicMock(),
32 'xfuser.model_executor.layers.attention_processor': MagicMock(),
33}
34mock_modules.update(mock_torch.get_sub_modules())
35mock_modules.update(mock_diffusers.get_sub_modules())
37with patch.dict(sys.modules, mock_modules):
38 from fluxkrea.wrapper_fluxkrea import FluxKreaGeneration
39 _fluxkrea_module = sys.modules['fluxkrea.wrapper_fluxkrea']
42@pytest.mark.asyncio
43async def test_wrapper_fluxkrea() -> None:
44 model = FluxKreaGeneration()
45 assert model is not None
46 assert model.model_name == "fluxkrea"
47 assert model.status == "initializing"
49 with pytest.raises(ValueError, match="Model not initialized"):
50 await model.generate(
51 width=128,
52 height=80,
53 prompt="test prompt")
55 model.init()
56 assert model.status == "ok"
58 # Mock pipeline return object
59 mock_output = MagicMock()
60 mock_output.images = ["image"]
61 model.pipeline = MagicMock(return_value=mock_output)
62 model.pipeline.vae_scale_factor = 8
64 health = model.get_health()
65 assert health is not None
66 timestamps = model.get_timestamps()
67 assert timestamps is not None
69 with pytest.raises(ValueError):
70 await model.get_rest_args(None)
71 with pytest.raises(ValueError):
72 await model.get_rest_args({})
73 await model.get_rest_args({
74 "job_id": "unittest",
75 "prompt": "Test prompt",
76 "width": 80,
77 "height": 60,
78 })
80 await model.warmup()
82 # TODO this should fail
83 await model.generate(
84 width=17,
85 height=13,
86 prompt="Test prompt")
88 image = await model.generate(
89 width=256,
90 height=256,
91 prompt="Test prompt")
92 assert image is not None
94 image = await model.generate(
95 width=256,
96 height=160,
97 prompt="Test prompt")
98 assert image is not None
100 # 15x17 not supported for 2 GPUs.
101 model.world_size = 2
102 # TODO
103 # with pytest.raises(ValueError, msg="15x17 not supported for 2 GPUs."):
104 image = await model.generate(
105 width=15,
106 height=17,
107 prompt="Test prompt")
109 del model
112@pytest.mark.asyncio
113async def test_additional_coverage() -> None:
114 """Cover seed path, step callbacks, parallelism init, and compile-disabled path."""
115 model = FluxKreaGeneration()
116 model.init()
117 assert model.status == "ok"
119 image = await model.generate(
120 width=256,
121 height=256,
122 prompt="Seed coverage test",
123 seed=42)
124 assert image is not None
126 pipeline_instance = model.pipeline
128 def _pipeline_with_callback(*args: Any, **kwargs: Any) -> Any:
129 n_steps = kwargs.get("num_inference_steps", 2)
130 callback = kwargs.get("callback_on_step_end")
131 if callback:
132 for step in range(n_steps):
133 callback(pipeline_instance, step, 0, {})
134 out = MagicMock()
135 out.images = [MagicMock()]
136 return out
138 pipeline_instance.side_effect = _pipeline_with_callback
139 image = await model.generate(
140 width=256,
141 height=256,
142 prompt="Callback coverage test",
143 sampling_steps=2)
144 assert image is not None
146 model.world_size = 2
147 with patch.object(_fluxkrea_module, 'parallelize_transformer'):
148 model.init_model_parallelism()
150 model.torch_compile = False
151 model.model_compile()
153 del model
156def test_model_compile_no_pipeline() -> None:
157 """model_compile() with pipeline=None returns early (pipeline not yet loaded)."""
158 model = FluxKreaGeneration()
159 assert model.pipeline is None
160 model.model_compile()
161 del model