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
« 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
8from typing import Optional
9from typing import Dict
10from typing import Any
12from PIL import Image
14import torch
15import torch.distributed as dist
16from torch import inference_mode
18from wrapper_flux import FluxGeneration
20from diffusers import Flux2Pipeline
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
29class Flux2Generation(FluxGeneration):
30 """Wrapper class for FLUX.2-dev image generation using Hugging Face Diffusers and Xfuser."""
32 HF_MODEL_NAME = "black-forest-labs/FLUX.2-dev"
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 )
46 self.pipeline: Optional[Flux2Pipeline] = None
48 def load_model(self) -> None:
49 """Load the FLUX.2-dev model."""
50 assert torch.cuda.is_available()
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")
73 logging.info(
74 "Loaded Flux2Pipeline: %s device:%s dtype:%s.",
75 self.HF_MODEL_NAME, self.device, self.param_dtype)
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
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")
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
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")
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.
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)
130 self._assert_model_init()
131 self._assert_args(height, width)
132 assert self.pipeline is not None
134 self.running = True
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)
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
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 )
168 assert len(output.images) == 1, f"Expected 1 image, but got {len(output.images)} images."
170 return output.images[0]
171 finally:
172 self.running = False
173 gen_timer.end("total")
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 }