Coverage for wrapper/flux/wrapper_flux.py: 96%

134 statements  

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

1""" 

2Wrapper class for FLUX model generation using Hugging Face Diffusers and Xfuser. 

3""" 

4import logging 

5import sys 

6import random 

7import asyncio 

8 

9from typing import override 

10from typing import Optional 

11from typing import Dict 

12from typing import Any 

13from typing import Union 

14 

15from PIL import Image 

16 

17import torch 

18import torch.distributed as dist 

19from torch import inference_mode 

20 

21from wrapper_usp import USPGeneration 

22 

23from flux_xfuser import parallelize_transformer 

24 

25from diffusers import FluxPipeline 

26 

27from xfuser.config import EngineConfig 

28from xfuser.core.distributed import get_runtime_state 

29from xfuser.core.distributed import initialize_runtime_state 

30from xfuser.core.distributed import get_pipeline_parallel_world_size 

31 

32 

33class FluxGeneration(USPGeneration): 

34 """Wrapper class for FLUX model generation using Hugging Face Diffusers and Xfuser.""" 

35 

36 MAX_LOG_TEXT_LEN = 64 

37 

38 def __init__( 

39 self, 

40 model_name: str = "flux", 

41 engine_config: EngineConfig = None, 

42 param_dtype: torch.dtype = torch.bfloat16, 

43 ) -> None: 

44 super().__init__( 

45 model_name=model_name, 

46 engine_config=engine_config, 

47 param_dtype=param_dtype, 

48 ) 

49 

50 # Model components 

51 self.pipeline: Optional[FluxPipeline] = None 

52 

53 def __del__(self) -> None: 

54 # Clean models 

55 if self.pipeline is not None: 

56 self.pipeline = None 

57 super().__del__() 

58 

59 def load_model(self) -> None: 

60 self.load_timer.start("pipeline") 

61 cache_args = None 

62 """ 

63 cache_args = { 

64 "use_teacache": engine_args.use_teacache, 

65 "use_fbcache": engine_args.use_fbcache, 

66 "rel_l1_thresh": 0.12, 

67 "return_hidden_states_first": False, 

68 "num_steps": input_config.num_inference_steps, 

69 } 

70 """ 

71 self.MODEL_NAME = "black-forest-labs/FLUX.1-dev" 

72 self.pipeline = FluxPipeline.from_pretrained( 

73 pretrained_model_name_or_path=self.MODEL_NAME, 

74 engine_config=self.engine_config, 

75 cache_args=cache_args, 

76 torch_dtype=self.param_dtype, 

77 # device_map="auto", # TODO check if needed 

78 ) 

79 if not self.pipeline: 

80 raise ValueError("Failed to load FLUX pipeline.") 

81 assert isinstance(self.pipeline, FluxPipeline) 

82 # TODO save some memory for V100 32GB 

83 # https://huggingface.co/docs/diffusers/main/en/optimization/memory 

84 # https://huggingface.co/docs/diffusers/main/en/optimization/memory#reduce-memory-usage 

85 # self.pipeline.enable_sequential_cpu_offload() 

86 # self.pipeline.enable_model_cpu_offload() 

87 # https://huggingface.co/docs/diffusers/en/training/distributed_inference#model-sharding 

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

89 self.load_timer.end("pipeline") 

90 

91 logging.info( 

92 f"Loaded FluxPipeline: {self.MODEL_NAME} device:{self.device} dtype:{self.param_dtype}.") 

93 

94 def init_model_parallelism(self) -> None: 

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

96 return 

97 

98 self.load_timer.start("dit_parallel") 

99 initialize_runtime_state(self.pipeline, self.engine_config) 

100 get_runtime_state().set_input_parameters( 

101 batch_size=1, 

102 # height=self.input_config.height, 

103 # width=self.input_config.width, 

104 # num_inference_steps=self.input_config.num_inference_steps, 

105 max_condition_sequence_length=512, 

106 split_text_embed_in_sp=get_pipeline_parallel_world_size() == 1, 

107 ) 

108 

109 parallelize_transformer(self.pipeline) 

110 self.load_timer.end("dit_parallel") 

111 

112 def model_compile(self) -> None: 

113 if not self.torch_compile: 

114 return 

115 if not self.pipeline: 

116 raise ValueError("FLUX pipeline not initialized.") 

117 

118 self.load_timer.start("dit_compile") 

119 torch._inductor.config.reorder_for_compute_comm_overlap = True 

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

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

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

123 ) 

124 self.load_timer.end("dit_compile") 

125 

126 def _assert_model_init(self) -> None: 

127 super()._assert_model_init() 

128 if self.pipeline is None: 

129 raise ValueError("FLUX pipeline not initialized.") 

130 

131 def _get_vae_scale_factor(self) -> int: 

132 if not self.pipeline: 

133 raise ValueError("Model not initialized.") 

134 vae_scale_factor = getattr(self.pipeline, "vae_scale_factor", None) 

135 if vae_scale_factor is None: 

136 raise ValueError("Model does not have vae_scale_factor.") 

137 return vae_scale_factor 

138 

139 def _assert_args( 

140 self, 

141 height: int, 

142 width: int, 

143 ) -> None: 

144 # Check if the image size is supported for the current parallelism setting 

145 # https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/flux/pipeline_flux.py 

146 vae_scale_factor = self._get_vae_scale_factor() 

147 height_latent = height // vae_scale_factor 

148 width_latent = width // vae_scale_factor 

149 img_latent_shape = (height_latent // 2) * (width_latent // 2) 

150 if img_latent_shape % self.world_size != 0: 

151 raise ValueError(f"{width}x{height} not supported for {self.world_size} GPUs.") 

152 

153 @inference_mode() 

154 async def warmup(self) -> None: 

155 """Warmup the model with a sample generation.""" 

156 logging.info(f"[{self.rank}] Warmup for FLUX generation.") 

157 await self.generate( 

158 # Ideally, we would use smaller sizes, but it has issues with 8 GPUs 

159 job_id="warmup", 

160 width=1280, 

161 height=800, 

162 prompt="A warmup image to initialize the model.", 

163 neg_prompt="", 

164 sampling_steps=5) # It needs at least 5 steps to warm up properly 

165 

166 @override 

167 @inference_mode() 

168 async def generate( 

169 self, 

170 height: int, 

171 width: int, 

172 prompt: str, 

173 neg_prompt: str = "", 

174 sampling_steps: int = 25, 

175 seed: Optional[int] = None, 

176 job_id: Optional[str] = None, 

177 ) -> Image.Image: 

178 """Generate an image from a prompt using the FLUX model.""" 

179 gen_timer = self._new_gen_timer(job_id) 

180 

181 self._assert_model_init() 

182 self._assert_args(height, width) 

183 assert self.pipeline is not None 

184 

185 self.running = True # Mark running to avoid concurrent calls 

186 

187 try: 

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

189 self.set_seed(seed) 

190 seed = random.randint(0, sys.maxsize) 

191 if self.base_seed is not None and self.base_seed >= 0: 

192 seed = self.base_seed 

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

194 seed_g.manual_seed(seed) 

195 

196 def callback_gen_timer( 

197 pipeline: FluxPipeline, 

198 step: int, 

199 timestep: int, 

200 callback_kwargs: Dict[str, Any], 

201 ) -> Dict[str, Any]: 

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

203 logging.info(f"[{self.rank}] Step {step + 1}/{sampling_steps}.") 

204 

205 if step < sampling_steps - 1: 

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

207 self.check_interrupted() 

208 return callback_kwargs 

209 

210 logging.info( 

211 f"[{self.rank}] Generating image with {width}x{height} and '{prompt[:self.MAX_LOG_TEXT_LEN]}'...") 

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

213 output: Any = await asyncio.to_thread( 

214 lambda: self.pipeline( # type: ignore[operator, misc] 

215 width=width, 

216 height=height, 

217 prompt=prompt, 

218 negative_prompt=neg_prompt, 

219 num_inference_steps=sampling_steps, 

220 output_type="pil", 

221 generator=seed_g, 

222 callback_on_step_end=callback_gen_timer, 

223 ) 

224 ) 

225 

226 if not output or len(output.images) != 1: 

227 raise ValueError(f"Expected 1 image, but got {len(output.images)} images") 

228 image = output.images[0] 

229 return image 

230 finally: 

231 self.running = False 

232 torch.cuda.empty_cache() 

233 gen_timer.end("total") 

234 

235 def get_health(self) -> Dict[str, Any]: 

236 ret = super().get_health() 

237 ret.update({ 

238 "device_map": getattr(self.pipeline, "hf_device_map", None) if self.pipeline else None, 

239 }) 

240 return ret 

241 

242 async def get_rest_args( 

243 self, 

244 data_json: Dict[str, Union[str, int, float]] 

245 ) -> Dict[str, Any]: 

246 if data_json is None or not isinstance(data_json, dict): 

247 raise ValueError("Missing JSON body") 

248 

249 job_id = data_json.get("job_id", None) 

250 

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

252 if prompt is None: 

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

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

255 

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

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

258 steps = int(data_json.get("sampling_steps", 20)) 

259 

260 rest_args: Dict[str, Any] = { 

261 "task": self.model_name, 

262 "args": { 

263 "job_id": job_id, 

264 "prompt": prompt, 

265 "neg_prompt": neg_prompt, 

266 "height": height, 

267 "width": width, 

268 "sampling_steps": steps, 

269 } 

270 } 

271 if "seed" in data_json: 

272 seed = data_json.get("seed", -1) 

273 if seed is not None: 

274 rest_args["args"]["seed"] = int(seed) 

275 return rest_args