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

1import logging 

2import sys 

3import random 

4 

5from typing import Optional 

6from typing import Dict 

7from typing import Any 

8 

9from PIL import Image 

10 

11import torch 

12import torch.distributed as dist 

13from torch import inference_mode 

14 

15from wrapper_flux import FluxGeneration 

16 

17from flux_xfuser import parallelize_transformer 

18 

19from diffusers import FluxPipeline 

20 

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 

25 

26 

27class FluxKreaGeneration(FluxGeneration): 

28 """Handle image generation using the Flux Krea model.""" 

29 HF_MODEL_NAME = "black-forest-labs/FLUX.1-Krea-dev" 

30 

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) 

41 

42 self.pipeline: Optional[FluxPipeline] = None 

43 

44 def load_model(self) -> None: 

45 """Load the Flux Krea model.""" 

46 assert torch.cuda.is_available() 

47 

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

62 

63 logging.info( 

64 "Loaded FluxKreaPipeline: %s device:%s dtype:%s.", 

65 self.HF_MODEL_NAME, self.device, 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 if self.pipeline is None: 

91 return 

92 

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

100 

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) 

122 

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 

128 

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

130 

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) 

139 

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 

151 

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 ) 

163 

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

165 

166 return output.images[0] 

167 finally: 

168 self.running = False 

169 gen_timer.end("total") 

170 

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 }