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
« 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
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 Flux2KleinPipeline
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 Flux2KleinGeneration(FluxGeneration):
30 """Wrapper class for FLUX.2-klein-9B image generation using Hugging Face Diffusers and Xfuser."""
32 HF_MODEL_NAME = "black-forest-labs/FLUX.2-klein-9B"
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 )
46 self.pipeline: Optional[Flux2KleinPipeline] = None
48 def load_model(self) -> None:
49 """Load the FLUX.2-klein-9B model."""
50 assert torch.cuda.is_available()
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")
69 logging.info(
70 "Loaded Flux2KleinPipeline: %s device:%s dtype:%s.",
71 self.HF_MODEL_NAME, self.device, self.param_dtype)
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
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")
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
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")
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.
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)
126 self._assert_model_init()
127 self._assert_args(height, width)
128 assert self.pipeline is not None
130 self.running = True
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)
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
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 )
164 assert len(output.images) == 1, f"Expected 1 image, but got {len(output.images)} images."
166 return output.images[0]
167 finally:
168 self.running = False
169 gen_timer.end("total")
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 }