Coverage for tests/test_wrapper_flux2klein.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
13mock_torch = TorchMock()
14mock_diffusers = DiffusersMock()
16sys.path.append("wrapper")
17sys.path.append("wrapper/flux2klein")
18sys.path.append("wrapper/flux")
20mock_modules = {
21 'nvidia_smi': MagicMock(),
22 'imageio': MagicMock(),
23 'cv2': MagicMock(),
24 'torch': mock_torch,
25 'xfuser': MagicMock(),
26 'xfuser.config': MagicMock(),
27 'xfuser.core': MagicMock(),
28 'xfuser.core.distributed': MagicMock(),
29 'xfuser.model_executor': MagicMock(),
30 'xfuser.model_executor.models': MagicMock(),
31 'xfuser.model_executor.models.transformers.transformer_flux': MagicMock(),
32 'xfuser.model_executor.models.transformers.transformer_flux2': 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 flux2klein.wrapper_flux2klein import Flux2KleinGeneration
43@pytest.mark.asyncio
44async def test_wrapper_flux2klein() -> None:
45 model = Flux2KleinGeneration()
46 assert model is not None
47 assert model.model_name == "flux2klein"
48 assert model.status == "initializing"
50 with pytest.raises(ValueError, match="Model not initialized"):
51 await model.generate(
52 width=128,
53 height=80,
54 prompt="test prompt")
56 model.init()
57 assert model.status == "ok"
59 # Mock pipeline return object
60 mock_output = MagicMock()
61 mock_output.images = ["image"]
62 model.pipeline = MagicMock(return_value=mock_output)
63 model.pipeline.vae_scale_factor = 8
65 health = model.get_health()
66 assert health is not None
67 timestamps = model.get_timestamps()
68 assert timestamps is not None
70 with pytest.raises(ValueError):
71 await model.get_rest_args(None)
72 with pytest.raises(ValueError):
73 await model.get_rest_args({})
74 await model.get_rest_args({
75 "job_id": "unittest",
76 "prompt": "Test prompt",
77 "width": 80,
78 "height": 60,
79 })
81 await model.warmup()
83 image = await model.generate(
84 width=256,
85 height=256,
86 prompt="Test prompt")
87 assert image is not None
89 image = await model.generate(
90 width=256,
91 height=160,
92 prompt="Test prompt")
93 assert image is not None
95 # 15x17 not supported for 2 GPUs.
96 model.world_size = 2
97 image = await model.generate(
98 width=15,
99 height=17,
100 prompt="Test prompt")
102 del model
105@pytest.mark.asyncio
106async def test_additional_coverage() -> None:
107 """Cover seed path, step callbacks, parallelism init, and compile-disabled path."""
108 model = Flux2KleinGeneration()
109 model.init()
110 assert model.status == "ok"
112 image = await model.generate(
113 width=256,
114 height=320,
115 prompt="Seed coverage test",
116 seed=42)
117 assert image is not None
119 pipeline_instance = model.pipeline
121 def _pipeline_with_callback(*args: Any, **kwargs: Any) -> Any:
122 n_steps = kwargs.get("num_inference_steps", 2)
123 callback = kwargs.get("callback_on_step_end")
124 if callback:
125 for step in range(n_steps):
126 callback(pipeline_instance, step, 0, {})
127 out = MagicMock()
128 out.images = [MagicMock()]
129 return out
131 pipeline_instance.side_effect = _pipeline_with_callback
132 image = await model.generate(
133 width=256,
134 height=320,
135 prompt="Callback coverage test",
136 sampling_steps=2)
137 assert image is not None
139 model.world_size = 2
140 model.init_model_parallelism()
142 model.torch_compile = False
143 model.model_compile()
145 del model
148def test_model_compile_no_pipeline() -> None:
149 """model_compile() with pipeline=None returns early (pipeline not yet loaded)."""
150 model = Flux2KleinGeneration()
151 assert model.pipeline is None
152 model.model_compile()
153 del model