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

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 

13mock_torch = TorchMock() 

14mock_diffusers = DiffusersMock() 

15 

16sys.path.append("fluxkrea") 

17sys.path.append("flux") 

18 

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

36 

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

38 from fluxkrea.wrapper_fluxkrea import FluxKreaGeneration 

39 _fluxkrea_module = sys.modules['fluxkrea.wrapper_fluxkrea'] 

40 

41 

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" 

48 

49 with pytest.raises(ValueError, match="Model not initialized"): 

50 await model.generate( 

51 width=128, 

52 height=80, 

53 prompt="test prompt") 

54 

55 model.init() 

56 assert model.status == "ok" 

57 

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 

63 

64 health = model.get_health() 

65 assert health is not None 

66 timestamps = model.get_timestamps() 

67 assert timestamps is not None 

68 

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

79 

80 await model.warmup() 

81 

82 # TODO this should fail 

83 await model.generate( 

84 width=17, 

85 height=13, 

86 prompt="Test prompt") 

87 

88 image = await model.generate( 

89 width=256, 

90 height=256, 

91 prompt="Test prompt") 

92 assert image is not None 

93 

94 image = await model.generate( 

95 width=256, 

96 height=160, 

97 prompt="Test prompt") 

98 assert image is not None 

99 

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

108 

109 del model 

110 

111 

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" 

118 

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 

125 

126 pipeline_instance = model.pipeline 

127 

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 

137 

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 

145 

146 model.world_size = 2 

147 with patch.object(_fluxkrea_module, 'parallelize_transformer'): 

148 model.init_model_parallelism() 

149 

150 model.torch_compile = False 

151 model.model_compile() 

152 

153 del model 

154 

155 

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