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
« 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
9from typing import override
10from typing import Optional
11from typing import Dict
12from typing import Any
13from typing import Union
15from PIL import Image
17import torch
18import torch.distributed as dist
19from torch import inference_mode
21from wrapper_usp import USPGeneration
23from flux_xfuser import parallelize_transformer
25from diffusers import FluxPipeline
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
33class FluxGeneration(USPGeneration):
34 """Wrapper class for FLUX model generation using Hugging Face Diffusers and Xfuser."""
36 MAX_LOG_TEXT_LEN = 64
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 )
50 # Model components
51 self.pipeline: Optional[FluxPipeline] = None
53 def __del__(self) -> None:
54 # Clean models
55 if self.pipeline is not None:
56 self.pipeline = None
57 super().__del__()
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")
91 logging.info(
92 f"Loaded FluxPipeline: {self.MODEL_NAME} device:{self.device} dtype:{self.param_dtype}.")
94 def init_model_parallelism(self) -> None:
95 if not dist.is_initialized() or self.world_size <= 1:
96 return
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 )
109 parallelize_transformer(self.pipeline)
110 self.load_timer.end("dit_parallel")
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.")
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")
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.")
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
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.")
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
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)
181 self._assert_model_init()
182 self._assert_args(height, width)
183 assert self.pipeline is not None
185 self.running = True # Mark running to avoid concurrent calls
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)
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}.")
205 if step < sampling_steps - 1:
206 gen_timer.start(f"step_{step + 1:03d}")
207 self.check_interrupted()
208 return callback_kwargs
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 )
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")
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
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")
249 job_id = data_json.get("job_id", None)
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", "")
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))
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