Coverage for wrapper/flux2/wrapper_flux2.py: 100%

83 statements  

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

1""" 

2Wrapper class for FLUX.2-dev 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 Flux2Pipeline 

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 Flux2Generation(FluxGeneration): 

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

31 

32 HF_MODEL_NAME = "black-forest-labs/FLUX.2-dev" 

33 

34 def __init__( 

35 self, 

36 model_name: str = "flux2", 

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[Flux2Pipeline] = None 

47 

48 def load_model(self) -> None: 

49 """Load the FLUX.2-dev model.""" 

50 assert torch.cuda.is_available() 

51 

52 self.load_timer.start("pipeline") 

53 # Use device_map="balanced" to shard the large transformer across all 

54 # available GPUs instead of loading it onto a single device (OOM risk). 

55 transformer = xFuserFlux2Transformer2DWrapper.from_pretrained( 

56 pretrained_model_name_or_path=self.HF_MODEL_NAME, 

57 torch_dtype=self.param_dtype, 

58 subfolder="transformer", 

59 device_map="balanced", 

60 ) # nosec B615 

61 # device_map="balanced" distributes the remaining pipeline components 

62 # (VAE, text encoders) across all available GPUs. The transformer is 

63 # already sharded via its own device_map above; providing it here 

64 # prevents diffusers from loading it a second time from disk. 

65 self.pipeline = Flux2Pipeline.from_pretrained( 

66 pretrained_model_name_or_path=self.HF_MODEL_NAME, 

67 torch_dtype=self.param_dtype, 

68 transformer=transformer, 

69 device_map="balanced", 

70 ) 

71 self.load_timer.end("pipeline") 

72 

73 logging.info( 

74 "Loaded Flux2Pipeline: %s device:%s dtype:%s.", 

75 self.HF_MODEL_NAME, self.device, self.param_dtype) 

76 

77 def init_model_parallelism(self) -> None: 

78 """Initialize model parallelism using xfuser.""" 

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

80 return 

81 

82 self.load_timer.start("dit_parallel") 

83 initialize_runtime_state(self.pipeline, self.engine_config) 

84 get_runtime_state().set_input_parameters( 

85 batch_size=1, 

86 max_condition_sequence_length=512, 

87 split_text_embed_in_sp=get_pipeline_parallel_world_size() == 1, 

88 ) 

89 self.load_timer.end("dit_parallel") 

90 

91 def model_compile(self) -> None: 

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

93 if not self.torch_compile: 

94 return 

95 if self.pipeline is None: 

96 return 

97 

98 self.load_timer.start("dit_compile") 

99 torch._inductor.config.reorder_for_compute_comm_overlap = True 

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

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

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

103 ) 

104 self.load_timer.end("dit_compile") 

105 

106 @inference_mode() 

107 async def generate( 

108 self, 

109 width: int, 

110 height: int, 

111 prompt: str, 

112 neg_prompt: str = "", 

113 sampling_steps: int = 25, 

114 seed: Optional[int] = None, 

115 job_id: Optional[str] = None, 

116 ) -> Image.Image: 

117 """Generate an image from a prompt using the FLUX.2-dev model. 

118 

119 Args: 

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

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

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

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

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

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

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

127 """ 

128 gen_timer = self._new_gen_timer(job_id) 

129 

130 self._assert_model_init() 

131 self._assert_args(height, width) 

132 assert self.pipeline is not None 

133 

134 self.running = True 

135 

136 try: 

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

138 self.set_seed(seed) 

139 else: 

140 self.reset_seed() 

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

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

143 seed_g.manual_seed(seed) 

144 

145 def callback_gen_timer( 

146 pipeline: Flux2Pipeline, 

147 step: int, 

148 timestep: int, 

149 callback_kwargs: dict 

150 ) -> dict: 

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

152 if step < sampling_steps - 1: 

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

154 self.check_interrupted() 

155 return callback_kwargs 

156 

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

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

159 height=height, 

160 width=width, 

161 prompt=prompt, 

162 num_inference_steps=sampling_steps, 

163 output_type="pil", 

164 generator=seed_g, 

165 callback_on_step_end=callback_gen_timer, 

166 ) 

167 

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

169 

170 return output.images[0] 

171 finally: 

172 self.running = False 

173 gen_timer.end("total") 

174 

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

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

177 if data_json is None: 

178 raise ValueError("Missing JSON body") 

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

180 if prompt is None: 

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

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

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

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

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

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

187 return { 

188 "task": self.model_name, 

189 "args": { 

190 "prompt": prompt, 

191 "neg_prompt": neg_prompt, 

192 "height": height, 

193 "width": width, 

194 "sampling_steps": steps, 

195 "seed": seed, 

196 } 

197 }