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

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

17sys.path.append("wrapper/flux2klein") 

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

19 

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

38 

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

40 from flux2klein.wrapper_flux2klein import Flux2KleinGeneration 

41 

42 

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" 

49 

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

51 await model.generate( 

52 width=128, 

53 height=80, 

54 prompt="test prompt") 

55 

56 model.init() 

57 assert model.status == "ok" 

58 

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 

64 

65 health = model.get_health() 

66 assert health is not None 

67 timestamps = model.get_timestamps() 

68 assert timestamps is not None 

69 

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

80 

81 await model.warmup() 

82 

83 image = await model.generate( 

84 width=256, 

85 height=256, 

86 prompt="Test prompt") 

87 assert image is not None 

88 

89 image = await model.generate( 

90 width=256, 

91 height=160, 

92 prompt="Test prompt") 

93 assert image is not None 

94 

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

101 

102 del model 

103 

104 

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" 

111 

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 

118 

119 pipeline_instance = model.pipeline 

120 

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 

130 

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 

138 

139 model.world_size = 2 

140 model.init_model_parallelism() 

141 

142 model.torch_compile = False 

143 model.model_compile() 

144 

145 del model 

146 

147 

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