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

1#!/usr/bin/env python3 

2 

3import sys 

4import pytest 

5 

6from typing import Any 

7 

8from unittest.mock import patch 

9from unittest.mock import MagicMock 

10from tests.torch_mock import TorchMock 

11from tests.diffusers_mock import DiffusersMock 

12 

13from PIL import Image 

14 

15mock_torch = TorchMock() 

16mock_diffusers = DiffusersMock() 

17 

18sys.path.append("wrapper") 

19sys.path.append("wrapper/flux") 

20 

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()) 

38 

39with patch.dict(sys.modules, mock_modules): 

40 from flux.wrapper_flux import FluxGeneration 

41 _flux_module = sys.modules['flux.wrapper_flux'] 

42 

43 

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" 

50 

51 with pytest.raises(ValueError, match="Model not initialized. Current status: initializing."): 

52 await model.generate(64, 48, "test prompt") 

53 

54 model.init() 

55 assert model.status == "ok" 

56 

57 health = model.get_health() 

58 assert health is not None 

59 timestamps = model.get_timestamps() 

60 assert timestamps is not None 

61 

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 }) 

73 

74 await model.warmup() 

75 

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) 

83 

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") 

91 

92 del model 

93 

94 

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" 

101 

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) 

108 

109 pipeline_instance = model.pipeline 

110 

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 

120 

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) 

128 

129 model.world_size = 2 

130 with patch.object(_flux_module, 'parallelize_transformer'): 

131 model.init_model_parallelism() 

132 

133 model.torch_compile = False 

134 model.model_compile() 

135 

136 del model 

137 

138 

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