Coverage for wrapper/fluxkrea/wrapper_fluxkrea.py: 99%
88 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
1import logging
2import sys
3import random
5from typing import Optional
6from typing import Dict
7from typing import Any
9from PIL import Image
11import torch
12import torch.distributed as dist
13from torch import inference_mode
15from wrapper_flux import FluxGeneration
17from flux_xfuser import parallelize_transformer
19from diffusers import FluxPipeline
21from xfuser.config import EngineConfig
22from xfuser.core.distributed import get_runtime_state
23from xfuser.core.distributed import initialize_runtime_state
24from xfuser.core.distributed import get_pipeline_parallel_world_size
27class FluxKreaGeneration(FluxGeneration):
28 """Handle image generation using the Flux Krea model."""
29 HF_MODEL_NAME = "black-forest-labs/FLUX.1-Krea-dev"
31 def __init__(
32 self,
33 model_name: str = "fluxkrea",
34 engine_config: EngineConfig = None,
35 param_dtype: torch.dtype = torch.bfloat16,
36 ) -> None:
37 super().__init__(
38 model_name=model_name,
39 engine_config=engine_config,
40 param_dtype=param_dtype)
42 self.pipeline: Optional[FluxPipeline] = None
44 def load_model(self) -> None:
45 """Load the Flux Krea model."""
46 assert torch.cuda.is_available()
48 self.load_timer.start("pipeline")
49 cache_args = None
50 self.pipeline = FluxPipeline.from_pretrained(
51 pretrained_model_name_or_path=self.HF_MODEL_NAME,
52 engine_config=self.engine_config,
53 cache_args=cache_args,
54 torch_dtype=self.param_dtype,
55 # device_map="auto", # TODO check if needed
56 )
57 if not self.pipeline:
58 raise ValueError("Failed to load FluxKrea pipeline.")
59 assert isinstance(self.pipeline, FluxPipeline)
60 self.pipeline = self.pipeline.to(self.device) # type: ignore[attr-defined]
61 self.load_timer.end("pipeline")
63 logging.info(
64 "Loaded FluxKreaPipeline: %s device:%s dtype:%s.",
65 self.HF_MODEL_NAME, self.device, 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
90 if self.pipeline is None:
91 return
93 self.load_timer.start("dit_compile")
94 torch._inductor.config.reorder_for_compute_comm_overlap = True
95 self.pipeline.transformer = torch.compile( # type: ignore[attr-defined]
96 self.pipeline.transformer, # type: ignore[attr-defined]
97 mode="max-autotune-no-cudagraphs"
98 )
99 self.load_timer.end("dit_compile")
101 @inference_mode()
102 async def generate(
103 self,
104 width: int,
105 height: int,
106 prompt: str,
107 neg_prompt: str = "",
108 sampling_steps: int = 25, # 10
109 seed: Optional[int] = None,
110 job_id: Optional[str] = None,
111 ) -> Image.Image:
112 """
113 Generate an image from another image using the Flux Krea model.
114 Args:
115 height (int): Height of the generated image.
116 width (int): Width of the generated image.
117 prompt (str): Text prompt to guide the image generation.
118 negative_prompt (str, optional): Negative prompt to avoid certain features in the image.
119 sampling_steps (int, optional): Number of inference steps for sampling. Default is 25.
120 """
121 gen_timer = self._new_gen_timer(job_id)
123 self._assert_model_init()
124 # Check if the image size is supported for the current parallelism setting
125 # https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/flux/pipeline_flux.py
126 self._assert_args(height, width)
127 assert self.pipeline is not None
129 self.running = True # Mark running to avoid concurrent calls
131 try:
132 if seed is not None and seed >= 0:
133 self.set_seed(seed)
134 else:
135 self.reset_seed()
136 seed = self.base_seed if self.base_seed >= 0 else random.randint(0, sys.maxsize)
137 seed_g = torch.Generator(device=self.device)
138 seed_g.manual_seed(seed)
140 def callback_gen_timer(
141 pipeline: FluxPipeline,
142 step: int,
143 timestep: int,
144 callback_kwargs: dict
145 ) -> dict:
146 gen_timer.end(f"step_{step:03d}")
147 if step < sampling_steps - 1:
148 gen_timer.start(f"step_{step + 1:03d}")
149 self.check_interrupted()
150 return callback_kwargs
152 gen_timer.start(f"step_{0:03d}")
153 output = self.pipeline( # type: ignore[operator]
154 height=height,
155 width=width,
156 prompt=prompt,
157 negative_prompt=neg_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, str]) -> 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", 20))
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 }