Coverage for wrapper/flux2klein/wrapper_flux2klein.py: 99%

87 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-09 04:47 +0000

1""" 

2Wrapper class for FLUX.2-klein-9B image generation using Hugging Face Diffusers and Xfuser. 

3""" 

4import logging 

5import sys 

6import random 

7 

8from typing import Optional 

9from typing import Dict 

10from typing import Any 

11 

12from PIL import Image 

13 

14import torch 

15import torch.distributed as dist 

16from torch import inference_mode 

17 

18from wrapper_flux import FluxGeneration 

19 

20from diffusers import Flux2KleinPipeline 

21 

22from xfuser.config import EngineConfig 

23from xfuser.core.distributed import get_runtime_state 

24from xfuser.core.distributed import initialize_runtime_state 

25from xfuser.core.distributed import get_pipeline_parallel_world_size 

26from xfuser.model_executor.models.transformers.transformer_flux2 import xFuserFlux2Transformer2DWrapper 

27 

28 

29class Flux2KleinGeneration(FluxGeneration): 

30 """Wrapper class for FLUX.2-klein-9B image generation using Hugging Face Diffusers and Xfuser.""" 

31 

32 HF_MODEL_NAME = "black-forest-labs/FLUX.2-klein-9B" 

33 

34 def __init__( 

35 self, 

36 model_name: str = "flux2klein", 

37 engine_config: EngineConfig = None, 

38 param_dtype: torch.dtype = torch.bfloat16, 

39 ) -> None: 

40 super().__init__( 

41 model_name=model_name, 

42 engine_config=engine_config, 

43 param_dtype=param_dtype, 

44 ) 

45 

46 self.pipeline: Optional[Flux2KleinPipeline] = None 

47 

48 def load_model(self) -> None: 

49 """Load the FLUX.2-klein-9B model.""" 

50 assert torch.cuda.is_available() 

51 

52 self.load_timer.start("pipeline") 

53 transformer = xFuserFlux2Transformer2DWrapper.from_pretrained( 

54 pretrained_model_name_or_path=self.HF_MODEL_NAME, 

55 torch_dtype=self.param_dtype, 

56 subfolder="transformer", 

57 ) # nosec B615 

58 self.pipeline = Flux2KleinPipeline.from_pretrained( 

59 pretrained_model_name_or_path=self.HF_MODEL_NAME, 

60 torch_dtype=self.param_dtype, 

61 transformer=transformer, 

62 ) 

63 if not self.pipeline: 

64 raise ValueError("Failed to load Flux2Klein pipeline.") 

65 assert isinstance(self.pipeline, Flux2KleinPipeline) 

66 self.pipeline = self.pipeline.to(self.device) # type: ignore[attr-defined] 

67 self.load_timer.end("pipeline") 

68 

69 logging.info( 

70 "Loaded Flux2KleinPipeline: %s device:%s dtype:%s.", 

71 self.HF_MODEL_NAME, self.device, self.param_dtype) 

72 

73 def init_model_parallelism(self) -> None: 

74 """Initialize model parallelism using xfuser.""" 

75 if not dist.is_initialized() or self.world_size <= 1: 

76 return 

77 

78 self.load_timer.start("dit_parallel") 

79 initialize_runtime_state(self.pipeline, self.engine_config) 

80 get_runtime_state().set_input_parameters( 

81 batch_size=1, 

82 max_condition_sequence_length=512, 

83 split_text_embed_in_sp=get_pipeline_parallel_world_size() == 1, 

84 ) 

85 self.load_timer.end("dit_parallel") 

86 

87 def model_compile(self) -> None: 

88 """Compile the model using torch.compile if enabled.""" 

89 if not self.torch_compile: 

90 return 

91 if self.pipeline is None: 

92 return 

93 

94 self.load_timer.start("dit_compile") 

95 torch._inductor.config.reorder_for_compute_comm_overlap = True 

96 self.pipeline.transformer = torch.compile( # type: ignore[attr-defined] 

97 self.pipeline.transformer, # type: ignore[attr-defined] 

98 mode="max-autotune-no-cudagraphs" 

99 ) 

100 self.load_timer.end("dit_compile") 

101 

102 @inference_mode() 

103 async def generate( 

104 self, 

105 width: int, 

106 height: int, 

107 prompt: str, 

108 neg_prompt: str = "", 

109 sampling_steps: int = 25, 

110 seed: Optional[int] = None, 

111 job_id: Optional[str] = None, 

112 ) -> Image.Image: 

113 """Generate an image from a prompt using the FLUX.2-klein-9B model. 

114 

115 Args: 

116 width (int): Width of the generated image. 

117 height (int): Height of the generated image. 

118 prompt (str): Text prompt to guide the image generation. 

119 neg_prompt (str, optional): Negative prompt to avoid certain features. 

120 sampling_steps (int, optional): Number of inference steps. Default is 25. 

121 seed (int, optional): Random seed for reproducibility. 

122 job_id (str, optional): Job identifier for logging and timing. 

123 """ 

124 gen_timer = self._new_gen_timer(job_id) 

125 

126 self._assert_model_init() 

127 self._assert_args(height, width) 

128 assert self.pipeline is not None 

129 

130 self.running = True 

131 

132 try: 

133 if seed is not None and seed >= 0: 

134 self.set_seed(seed) 

135 else: 

136 self.reset_seed() 

137 seed = self.base_seed if self.base_seed >= 0 else random.randint(0, sys.maxsize) 

138 seed_g = torch.Generator(device=self.device) 

139 seed_g.manual_seed(seed) 

140 

141 def callback_gen_timer( 

142 pipeline: Flux2KleinPipeline, 

143 step: int, 

144 timestep: int, 

145 callback_kwargs: dict 

146 ) -> dict: 

147 gen_timer.end(f"step_{step:03d}") 

148 if step < sampling_steps - 1: 

149 gen_timer.start(f"step_{step + 1:03d}") 

150 self.check_interrupted() 

151 return callback_kwargs 

152 

153 gen_timer.start(f"step_{0:03d}") 

154 output = self.pipeline( # type: ignore[operator] 

155 height=height, 

156 width=width, 

157 prompt=prompt, 

158 num_inference_steps=sampling_steps, 

159 output_type="pil", 

160 generator=seed_g, 

161 callback_on_step_end=callback_gen_timer, 

162 ) 

163 

164 assert len(output.images) == 1, f"Expected 1 image, but got {len(output.images)} images." 

165 

166 return output.images[0] 

167 finally: 

168 self.running = False 

169 gen_timer.end("total") 

170 

171 async def get_rest_args(self, data_json: Dict[str, Any]) -> Dict[str, Any]: 

172 """Extract and validate arguments from the REST API request.""" 

173 if data_json is None: 

174 raise ValueError("Missing JSON body") 

175 prompt = data_json.get("prompt", None) 

176 if prompt is None: 

177 raise ValueError("Missing 'prompt' parameter") 

178 neg_prompt = data_json.get("neg_prompt", "") 

179 height = int(data_json.get("height", 480)) 

180 width = int(data_json.get("width", 640)) 

181 steps = int(data_json.get("sampling_steps", 25)) 

182 seed = data_json.get("seed", None) 

183 return { 

184 "task": self.model_name, 

185 "args": { 

186 "prompt": prompt, 

187 "neg_prompt": neg_prompt, 

188 "height": height, 

189 "width": width, 

190 "sampling_steps": steps, 

191 "seed": seed, 

192 } 

193 }