Coverage for wrapper/fluxkontext/wrapper_fluxkontext.py: 100%
96 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 Kontext model generation.
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
14from PIL import Image
16import torch
17import torch.distributed as dist
18from torch import inference_mode
20from image_utils import base64_to_img
21from wrapper_flux import FluxGeneration
23from flux_xfuser import parallelize_transformer
25from diffusers import FluxKontextPipeline
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 FluxKontextGeneration(FluxGeneration):
34 """Class for generating images using the Flux Kontext model."""
36 def __init__(
37 self,
38 model_name: str = "fluxkontext",
39 engine_config: EngineConfig = None,
40 param_dtype: torch.dtype = torch.bfloat16,
41 ) -> None:
42 super().__init__(
43 model_name=model_name,
44 engine_config=engine_config,
45 param_dtype=param_dtype)
47 def load_model(self) -> None:
48 """Load the Flux Kontext model from Hugging Face."""
49 assert torch.cuda.is_available()
51 self.load_timer.start("pipeline")
52 cache_args = None
53 self.MODEL_NAME = "black-forest-labs/FLUX.1-Kontext-dev"
54 self.pipeline = FluxKontextPipeline.from_pretrained(
55 pretrained_model_name_or_path=self.MODEL_NAME,
56 engine_config=self.engine_config,
57 cache_args=cache_args,
58 torch_dtype=self.param_dtype,
59 # device_map="auto", # TODO check if needed
60 )
61 self.pipeline = self.pipeline.to(self.device)
62 self.load_timer.end("pipeline")
64 logging.info(
65 f"Loaded FluxKontextPipeline: {self.MODEL_NAME} device:{self.device} dtype:{self.param_dtype}.")
67 def init_model_parallelism(self) -> None:
68 """Initialize model parallelism using xfuser."""
69 if not dist.is_initialized() or self.world_size <= 1:
70 return
72 self.load_timer.start("dit_parallel")
73 initialize_runtime_state(self.pipeline, self.engine_config)
74 get_runtime_state().set_input_parameters(
75 batch_size=1,
76 # height=self.input_config.height,
77 # width=self.input_config.width,
78 # num_inference_steps=self.input_config.num_inference_steps,
79 max_condition_sequence_length=512,
80 split_text_embed_in_sp=get_pipeline_parallel_world_size() == 1,
81 )
83 parallelize_transformer(self.pipeline)
84 self.load_timer.end("dit_parallel")
86 def model_compile(self) -> None:
87 """Compile the model using torch compile if enabled."""
88 if not self.torch_compile:
89 return
91 self.load_timer.start("dit_compile")
92 torch._inductor.config.reorder_for_compute_comm_overlap = True
93 self.pipeline.transformer = torch.compile(
94 self.pipeline.transformer,
95 mode="max-autotune-no-cudagraphs"
96 )
97 self.load_timer.end("dit_compile")
99 @inference_mode()
100 async def warmup(self) -> None:
101 """Warmup the model with a dummy generation to initialize everything."""
102 logging.info(f"[{self.rank}] Warmup for Flux Kontext generation.")
103 empty_img = Image.new("RGB", (512, 512), (255, 255, 255))
104 await self.generate(
105 empty_img,
106 width=1280,
107 height=800,
108 prompt="A warmup image to initialize the model.",
109 neg_prompt="",
110 sampling_steps=5)
112 @override
113 @inference_mode()
114 async def generate(
115 self,
116 img: Image.Image,
117 height: int,
118 width: int,
119 prompt: str,
120 neg_prompt: str = "",
121 sampling_steps: int = 25, # 10
122 seed: Optional[int] = None,
123 job_id: Optional[str] = None,
124 ) -> Image.Image:
125 """
126 Generate an image from another image using the Flux Kontext model.
127 Args:
128 img (Image.Image): Input image to guide the generation.
129 height (int): Height of the generated image.
130 width (int): Width of the generated image.
131 prompt (str): Text prompt to guide the image generation.
132 negative_prompt (str, optional): Negative prompt to avoid certain features in the image.
133 sampling_steps (int, optional): Number of inference steps for sampling. Default is 25.
134 """
135 gen_timer = self._new_gen_timer(job_id)
137 self._assert_model_init()
138 self._assert_args(height, width)
140 gen_timer.start("image_preprocess")
141 img = img.resize((width, height), Image.Resampling.LANCZOS)
142 gen_timer.end("image_preprocess")
144 self.running = True # Mark running to avoid concurrent calls
146 try:
147 if seed is not None and seed >= 0:
148 self.set_seed(seed)
149 seed = self.base_seed if self.base_seed >= 0 else random.randint(0, sys.maxsize)
150 seed_g = torch.Generator(device=self.device)
151 seed_g.manual_seed(seed)
153 def callback_gen_timer(
154 pipeline: FluxKontextPipeline,
155 step: int,
156 timestep: int,
157 callback_kwargs: dict
158 ) -> dict:
159 gen_timer.end(f"step_{step:03d}")
160 if step < sampling_steps - 1:
161 gen_timer.start(f"step_{step + 1:03d}")
162 self.check_interrupted()
163 return callback_kwargs
165 gen_timer.start(f"step_{0:03d}")
166 output = await asyncio.to_thread(
167 self.pipeline,
168 image=img,
169 height=height,
170 width=width,
171 prompt=prompt,
172 negative_prompt=neg_prompt,
173 num_inference_steps=sampling_steps,
174 output_type="pil",
175 generator=seed_g,
176 callback_on_step_end=callback_gen_timer,
177 )
179 assert len(output.images) == 1, f"Expected 1 image, but got {len(output.images)} images."
181 return output.images[0]
182 finally:
183 self.running = False
184 gen_timer.end("total")
186 async def get_rest_args(self, data_json: Dict[str, str]) -> Dict[str, Any]:
187 if data_json is None:
188 raise ValueError("Missing JSON body")
189 img_base64 = data_json.get("img", None)
190 if img_base64 is None:
191 raise ValueError("Missing 'img' parameter")
192 img = base64_to_img(img_base64)
193 prompt = data_json.get("prompt", None)
194 if prompt is None:
195 raise ValueError("Missing 'prompt' parameter")
196 neg_prompt = data_json.get("neg_prompt", "")
197 height = int(data_json.get("height", 480))
198 width = int(data_json.get("width", 640))
199 steps = int(data_json.get("sampling_steps", 20))
200 seed = data_json.get("seed", None)
201 return {
202 "task": self.model_name,
203 "args": {
204 "img": img,
205 "prompt": prompt,
206 "neg_prompt": neg_prompt,
207 "width": width,
208 "height": height,
209 "sampling_steps": steps,
210 "seed": seed,
211 }
212 }