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

1""" 

2Wrapper class for Flux Kontext model generation. 

3""" 

4import logging 

5import sys 

6import random 

7import asyncio 

8 

9from typing import override 

10from typing import Optional 

11from typing import Dict 

12from typing import Any 

13 

14from PIL import Image 

15 

16import torch 

17import torch.distributed as dist 

18from torch import inference_mode 

19 

20from image_utils import base64_to_img 

21from wrapper_flux import FluxGeneration 

22 

23from flux_xfuser import parallelize_transformer 

24 

25from diffusers import FluxKontextPipeline 

26 

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 

31 

32 

33class FluxKontextGeneration(FluxGeneration): 

34 """Class for generating images using the Flux Kontext model.""" 

35 

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) 

46 

47 def load_model(self) -> None: 

48 """Load the Flux Kontext model from Hugging Face.""" 

49 assert torch.cuda.is_available() 

50 

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") 

63 

64 logging.info( 

65 f"Loaded FluxKontextPipeline: {self.MODEL_NAME} device:{self.device} dtype:{self.param_dtype}.") 

66 

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 

71 

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 ) 

82 

83 parallelize_transformer(self.pipeline) 

84 self.load_timer.end("dit_parallel") 

85 

86 def model_compile(self) -> None: 

87 """Compile the model using torch compile if enabled.""" 

88 if not self.torch_compile: 

89 return 

90 

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") 

98 

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) 

111 

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) 

136 

137 self._assert_model_init() 

138 self._assert_args(height, width) 

139 

140 gen_timer.start("image_preprocess") 

141 img = img.resize((width, height), Image.Resampling.LANCZOS) 

142 gen_timer.end("image_preprocess") 

143 

144 self.running = True # Mark running to avoid concurrent calls 

145 

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) 

152 

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 

164 

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 ) 

178 

179 assert len(output.images) == 1, f"Expected 1 image, but got {len(output.images)} images." 

180 

181 return output.images[0] 

182 finally: 

183 self.running = False 

184 gen_timer.end("total") 

185 

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 }