Coverage for wrapper/bagel/wrapper_bagel.py: 37%
309 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
1from __future__ import annotations
3import logging
4import os
5import asyncio
7from typing import override
8from typing import List
9from typing import Optional
10from typing import Dict
11from typing import Any
13from PIL import Image
14from copy import deepcopy
16import torch
17import torch.distributed as dist
18from torch import inference_mode
20from wrapper_model import ModelGeneration
21from image_utils import base64_to_img
23from modeling.bagel import BagelConfig
24from modeling.bagel import Bagel
25from modeling.bagel import Qwen2Config
26from modeling.bagel import Qwen2ForCausalLM
27from modeling.bagel import SiglipVisionConfig
28from modeling.bagel import SiglipVisionModel
29from modeling.bagel.qwen2_navit import NaiveCache
30from modeling.autoencoder import load_ae
31from modeling.qwen2 import Qwen2Tokenizer
33from data.transforms import ImageTransform
34from data.data_utils import add_special_tokens
36from accelerate import load_checkpoint_and_dispatch
37from accelerate import init_empty_weights
39from xfuser.config import EngineConfig
42class BagelGeneration(ModelGeneration):
43 def __init__(
44 self,
45 engine_config: EngineConfig = None,
46 param_dtype: torch.dtype = torch.bfloat16,
47 ) -> None:
48 super().__init__("bagel")
50 self.engine_config = engine_config
51 self.param_dtype = param_dtype
53 # Parallelism
54 self.gpu = torch.cuda.get_device_name(0)
55 self.rank = -1
56 self.world_size = -1
57 self.local_rank = -1
58 self.device: Optional[torch.device] = None
60 # Model components
61 self.tokenizer: Optional[Qwen2Tokenizer] = None
62 self.new_token_ids: Optional[List[int]] = None
63 self.model: Optional[Bagel] = None
64 self.language_model: Optional[Qwen2ForCausalLM] = None
65 self.vit_model: Optional[SiglipVisionModel] = None
66 self.vae_model: Optional[Any] = None
67 self.vae_transform: Optional[ImageTransform] = None
68 self.vit_transform: Optional[ImageTransform] = None
70 def __del__(self) -> None:
71 if self.tokenizer is not None:
72 self.tokenizer = None
73 if self.model is not None:
74 self.model = None
75 if self.language_model is not None:
76 self.language_model = None
77 if self.vit_model is not None:
78 self.vit_model = None
79 if dist.is_initialized():
80 dist.destroy_process_group()
82 def init_parallelism(self) -> None:
83 self.load_timer.start("torch_dist")
84 self.rank = int(os.getenv("RANK", 0))
85 self.local_rank = int(os.getenv("LOCAL_RANK", 0))
86 self.world_size = int(os.getenv("WORLD_SIZE", 1))
88 self.device_id = self.local_rank
89 self.device = torch.device(f"cuda:{self.device_id}")
91 torch.cuda.set_device(self.local_rank)
92 self.load_timer.end("torch_dist")
94 # TODO implement xfuser parallelism
95 if self.world_size > 1:
96 logging.warning("Parallelism is not supported in Bagel generation, running on single device.")
98 def load_model(self) -> None:
99 assert torch.cuda.is_available()
101 model_path = "BAGEL-7B-MoT" # TODO temporary
103 # VAE
104 # TODO check the dtype
105 self.load_timer.start("vae")
106 self.vae_model, vae_config = load_ae(local_path=os.path.join(model_path, "ae.safetensors"))
107 self.vae_model.to(self.device)
108 self.load_timer.end("vae")
110 # LLM
111 self.load_timer.start("llm")
112 llm_config = Qwen2Config.from_json_file(os.path.join(model_path, "llm_config.json"))
113 llm_config.qk_norm = True
114 llm_config.tie_word_embeddings = False
115 llm_config.layer_module = "Qwen2MoTDecoderLayer"
116 self.load_timer.end("llm")
118 # ViT
119 self.load_timer.start("vit")
120 vit_config = SiglipVisionConfig.from_json_file(os.path.join(model_path, "vit_config.json"))
121 vit_config.rope = False
122 vit_config.num_hidden_layers = vit_config.num_hidden_layers - 1
123 self.load_timer.end("vit")
125 # Transformer
126 self.load_timer.start("bagel")
127 config = BagelConfig(
128 visual_gen=True,
129 visual_und=True,
130 llm_config=llm_config,
131 vit_config=vit_config,
132 vae_config=vae_config,
133 vit_max_num_patch_per_side=70,
134 connector_act="gelu_pytorch_tanh",
135 latent_patch_size=2,
136 max_latent_size=64,
137 )
139 with init_empty_weights():
140 self.language_model = Qwen2ForCausalLM(llm_config)
141 self.vit_model = SiglipVisionModel(vit_config)
142 self.model = Bagel(self.language_model, self.vit_model, config)
143 self.model.vit_model.vision_model.embeddings.convert_conv2d_to_linear(vit_config, meta=True)
145 self.model = load_checkpoint_and_dispatch(
146 self.model,
147 checkpoint=os.path.join(model_path, "ema.safetensors"),
148 offload_buffers=True,
149 dtype=self.param_dtype,
150 device_map={"": self.device_id}, # TODO fix
151 force_hooks=True,
152 offload_folder="/tmp/offload"
153 )
154 # self.model.to(self.device) # done through accelerate#load_checkpoint_and_dispatch()
155 self.model.eval()
156 # self.model.require_grad_(False)
157 self.load_timer.end("bagel")
159 # Tokenizer
160 self.load_timer.start("tokenizer")
161 self.tokenizer = Qwen2Tokenizer.from_pretrained(
162 model_path,
163 torch_dtype=self.param_dtype,
164 )
165 # self.tokenizer.to(self.device)
166 self.tokenizer, self.new_token_ids, _ = add_special_tokens(self.tokenizer)
167 self.load_timer.end("tokenizer")
169 # Image Transform
170 self.vae_transform = ImageTransform(1024, 512, 16)
171 self.vit_transform = ImageTransform(980, 224, 14)
173 def init_model_parallelism(self) -> None:
174 if self.world_size > 1:
175 logging.warning("Parallelism not supported.")
177 def model_compile(self) -> None:
178 if not self.torch_compile:
179 return
181 self.load_timer.start("dit_compile")
182 torch._inductor.config.reorder_for_compute_comm_overlap = True
183 self.model = torch.compile(
184 self.model,
185 mode="max-autotune-no-cudagraphs"
186 )
187 self.load_timer.end("dit_compile")
189 @inference_mode()
190 def update_context_text(
191 self,
192 text: str,
193 gen_context: Dict[str, Any],
194 ) -> Dict[str, Any]:
195 assert self.model is not None
196 past_key_values = gen_context['past_key_values']
197 kv_lens = gen_context['kv_lens']
198 ropes = gen_context['ropes']
199 generation_input, kv_lens, ropes = self.model.prepare_prompts(
200 curr_kvlens=kv_lens,
201 curr_rope=ropes,
202 prompts=[text],
203 tokenizer=self.tokenizer,
204 new_token_ids=self.new_token_ids,
205 )
206 # TODO make the tokenizer stuff the right way
207 generation_input["text_token_lens"] = generation_input["text_token_lens"].to(self.device)
208 generation_input["packed_text_ids"] = generation_input["packed_text_ids"].to(self.device)
209 generation_input["packed_text_position_ids"] = generation_input["packed_text_position_ids"].to(self.device)
210 generation_input["packed_text_indexes"] = generation_input["packed_text_indexes"].to(self.device)
211 generation_input["packed_key_value_indexes"] = generation_input["packed_key_value_indexes"].to(self.device)
212 generation_input["key_values_lens"] = generation_input["key_values_lens"].to(self.device)
214 past_key_values = self.model.forward_cache_update_text(past_key_values, **generation_input)
215 gen_context['kv_lens'] = kv_lens
216 gen_context['ropes'] = ropes
217 gen_context['past_key_values'] = past_key_values
218 return gen_context
220 @inference_mode()
221 def update_context_images(
222 self,
223 images: List[Image.Image],
224 gen_context: Dict[str, Any],
225 vae: bool = True,
226 vit: bool = True
227 ) -> Dict[str, Any]:
228 past_key_values = gen_context['past_key_values']
229 kv_lens = gen_context['kv_lens']
230 ropes = gen_context['ropes']
232 # VAE
233 assert self.model is not None
234 generation_input, kv_lens, ropes = self.model.prepare_vae_images(
235 curr_kvlens=kv_lens,
236 curr_rope=ropes,
237 images=images,
238 transforms=self.vae_transform,
239 new_token_ids=self.new_token_ids,
240 )
241 generation_input["padded_images"] = generation_input["padded_images"].to(self.device)
242 past_key_values = self.model.forward_cache_update_vae(self.vae_model, past_key_values, **generation_input)
244 # ViT
245 generation_input, kv_lens, ropes = self.model.prepare_vit_images(
246 curr_kvlens=kv_lens,
247 curr_rope=ropes,
248 images=images,
249 transforms=self.vit_transform,
250 new_token_ids=self.new_token_ids,
251 )
252 past_key_values = self.model.forward_cache_update_vit(past_key_values, **generation_input)
254 # Output
255 gen_context['past_key_values'] = past_key_values
256 gen_context['kv_lens'] = kv_lens
257 gen_context['ropes'] = ropes
259 return gen_context
261 def decode_image(
262 self,
263 latent: torch.Tensor,
264 image_shape: tuple[int, int]
265 ) -> Image.Image:
266 assert self.model is not None
267 assert self.vae_model is not None
268 H, W = image_shape
269 h, w = H // self.model.latent_downsample, W // self.model.latent_downsample
270 latent = latent.reshape(1, h, w, self.model.latent_patch_size,
271 self.model.latent_patch_size, self.model.latent_channel)
272 latent = torch.einsum("nhwpqc->nchpwq", latent)
273 latent = latent.reshape(1, self.model.latent_channel, h
274 * self.model.latent_patch_size, w * self.model.latent_patch_size)
275 # TODO do this right away instead of float32 to bfloat16
276 latent = latent.to(self.param_dtype).to(self.device) # Ensure dtype matches model
277 image = self.vae_model.decode(latent)
278 image = (image * 0.5 + 0.5).clamp(0, 1)[0].permute(1, 2, 0) * 255
279 image = Image.fromarray((image).to(torch.uint8).cpu().numpy())
280 return image
282 def _assert_model_init(self) -> None:
283 super()._assert_model_init()
284 if self.model is None:
285 raise ValueError("Model not initialized.")
287 @inference_mode()
288 async def warmup(self) -> None:
289 logging.info(f"[{self.rank}] Warmup for Bagel generation")
290 empty_img_0 = Image.new("RGB", (512, 512), (255, 255, 255))
291 empty_img_1 = Image.new("RGB", (512, 512), (255, 255, 255))
292 await self.generate(
293 width=1024,
294 height=1024,
295 prompt="Warmup.",
296 neg_prompt="",
297 imgs=[empty_img_0, empty_img_1],
298 sampling_steps=2)
300 @override
301 @inference_mode()
302 async def generate(
303 self,
304 height: int,
305 width: int,
306 prompt: str,
307 neg_prompt: str = "", # TODO not used
308 imgs: List[Image.Image] = [],
309 sampling_steps: int = 50,
310 understanding_output: bool = False,
311 job_id: Optional[str] = None,
312 ) -> Image.Image:
313 gen_timer = self._new_gen_timer(job_id)
315 self._assert_model_init()
316 assert self.model is not None
317 assert self.vae_transform is not None
319 self.running = True # Mark running to avoid concurrent calls
321 try:
322 # Other arguments
323 cfg_renorm_min = 0.0
324 cfg_renorm_type = "global"
325 cfg_text_scale = 3.0
326 cfg_img_scale = 1.5
327 cfg_type = "parallel"
328 cfg_interval = (0.4, 1.0)
329 timestep_shift = 3.0
331 image_shape = (height, width)
333 # https://github.com/ByteDance-Seed/Bagel/blob/main/inferencer.py#L119
334 num_hidden_layers = self.model.config.llm_config.num_hidden_layers
335 # num_hidden_layers = 32 # Set this up properly
336 gen_context = {
337 'kv_lens': [0],
338 'ropes': [0],
339 'past_key_values': NaiveCache(num_hidden_layers),
340 }
341 cfg_text_context = deepcopy(gen_context)
342 cfg_img_context = deepcopy(gen_context)
344 with torch.autocast(device_type="cuda", enabled=True, dtype=self.param_dtype):
345 if imgs is not None and len(imgs) > 0:
346 img_latents = []
347 for img in imgs:
348 img_rgb = img.convert("RGB")
349 img_latent = self.vae_transform.resize_transform(img_rgb)
350 img_latents.append(img_latent)
351 gen_context = self.update_context_images(img_latents, gen_context, vae=not understanding_output)
352 # image_shapes = img_latent.size[::-1]
353 cfg_text_context = deepcopy(gen_context)
355 # TODO figure why the order of this matters
356 # Text (we can add multiple)
357 cfg_text_context = deepcopy(gen_context)
358 gen_context = self.update_context_text(prompt, gen_context)
359 cfg_img_context = self.update_context_text(prompt, cfg_img_context)
361 # VAE latent
362 gen_timer.start("vae_encoder")
363 past_key_values = gen_context['past_key_values']
364 kv_lens = gen_context['kv_lens']
365 ropes = gen_context['ropes']
366 generation_input = self.model.prepare_vae_latent(
367 curr_kvlens=kv_lens,
368 curr_rope=ropes,
369 image_sizes=[image_shape],
370 new_token_ids=self.new_token_ids,
371 )
372 packed_vae_token_indexes = generation_input['packed_vae_token_indexes']
373 packed_vae_position_ids = generation_input['packed_vae_position_ids']
374 packed_text_ids = generation_input['packed_text_ids']
375 packed_text_indexes = generation_input['packed_text_indexes']
376 packed_position_ids = generation_input['packed_position_ids']
377 packed_indexes = generation_input['packed_indexes']
378 packed_seqlens = generation_input['packed_seqlens']
379 key_values_lens = generation_input['key_values_lens']
380 packed_key_value_indexes = generation_input['packed_key_value_indexes']
381 gen_timer.end("vae_encoder")
383 # Text cfg
384 gen_timer.start("text_encoder")
385 cfg_text_past_key_values = cfg_text_context['past_key_values']
386 kv_lens_cfg = cfg_text_context['kv_lens']
387 ropes_cfg = cfg_text_context['ropes']
388 generation_input_cfg_text = self.model.prepare_vae_latent_cfg(
389 curr_kvlens=kv_lens_cfg,
390 curr_rope=ropes_cfg,
391 image_sizes=[image_shape],
392 )
393 gen_timer.end("text_encoder")
395 # Image cfg
396 gen_timer.start("image_encoder")
397 cfg_img_past_key_values = cfg_img_context['past_key_values']
398 kv_lens_cfg = cfg_img_context['kv_lens']
399 ropes_cfg = cfg_img_context['ropes']
400 generation_input_cfg_img = self.model.prepare_vae_latent_cfg(
401 curr_kvlens=kv_lens_cfg,
402 curr_rope=ropes_cfg,
403 image_sizes=[image_shape],
404 )
405 # TODO Try to get the VAE to generate in the GPU directly
406 x_t = generation_input['packed_init_noises']
407 x_t = x_t.to(self.device)
408 x_t = x_t.to(self.param_dtype)
409 gen_timer.end("image_encoder")
411 # Diffusion sampling
412 # https://github.com/ByteDance-Seed/Bagel/blob/main/modeling/bagel/bagel.py#L643
413 timesteps = torch.linspace(1, 0, sampling_steps, device=x_t.device)
414 timesteps = timestep_shift * timesteps / (1 + (timestep_shift - 1) * timesteps)
415 dts = timesteps[:-1] - timesteps[1:]
416 timesteps = timesteps[:-1]
418 for it, timestep in enumerate(timesteps):
419 gen_timer.start(f"dit_{it:03d}")
420 timestep_tensor = torch.tensor([timestep] * x_t.shape[0], device=x_t.device)
421 if timestep > cfg_interval[0] and timestep <= cfg_interval[1]:
422 cfg_text_scale_ = cfg_text_scale
423 cfg_img_scale_ = cfg_img_scale
424 else:
425 cfg_text_scale_ = 1.0
426 cfg_img_scale_ = 1.0
427 v_t = self.model._forward_flow(
428 x_t=x_t,
429 timestep=timestep_tensor,
430 packed_vae_token_indexes=packed_vae_token_indexes,
431 packed_vae_position_ids=packed_vae_position_ids,
432 packed_text_ids=packed_text_ids,
433 packed_text_indexes=packed_text_indexes,
434 packed_position_ids=packed_position_ids,
435 packed_indexes=packed_indexes,
436 packed_seqlens=packed_seqlens,
437 key_values_lens=key_values_lens,
438 past_key_values=past_key_values,
439 packed_key_value_indexes=packed_key_value_indexes,
440 cfg_renorm_min=cfg_renorm_min,
441 cfg_renorm_type=cfg_renorm_type,
442 # cfg_text
443 cfg_text_scale=cfg_text_scale_,
444 cfg_text_past_key_values=cfg_text_past_key_values,
445 # cfg_text_packed_position_ids,
446 cfg_text_packed_position_ids=generation_input_cfg_text['cfg_packed_position_ids'],
447 # cfg_text_packed_query_indexes,
448 cfg_text_packed_query_indexes=generation_input_cfg_text['cfg_packed_query_indexes'],
449 # cfg_text_key_values_lens,
450 cfg_text_key_values_lens=generation_input_cfg_text['cfg_key_values_lens'],
451 # cfg_text_packed_key_value_indexes,
452 cfg_text_packed_key_value_indexes=generation_input_cfg_text['cfg_packed_key_value_indexes'],
453 # cfg_img
454 cfg_img_scale=cfg_img_scale_,
455 cfg_img_past_key_values=cfg_img_past_key_values,
456 # cfg_img_packed_position_ids,
457 cfg_img_packed_position_ids=generation_input_cfg_img['cfg_packed_position_ids'],
458 # cfg_img_packed_query_indexes,
459 cfg_img_packed_query_indexes=generation_input_cfg_img['cfg_packed_query_indexes'],
460 # cfg_img_key_values_lens,
461 cfg_img_key_values_lens=generation_input_cfg_img['cfg_key_values_lens'],
462 # cfg_img_packed_key_value_indexes,
463 cfg_img_packed_key_value_indexes=generation_input_cfg_img['cfg_packed_key_value_indexes'],
464 cfg_type=cfg_type,
465 )
466 x_t = x_t - v_t.to(x_t.device) * dts[it] # velocity pointing from data to noise
467 gen_timer.end(f"dit_{it:03d}")
469 unpacked_latent = x_t.split((packed_seqlens - 2).tolist())
471 # Decode the image
472 gen_timer.start("image_decoder")
473 output_image = self.decode_image(unpacked_latent[0], image_shape)
474 gen_timer.end("image_decoder")
475 return output_image
476 finally:
477 self.running = False
478 gen_timer.end("total")
480 def get_health(self) -> Dict[str, Any]:
481 ret = super().get_health()
482 ret.update({
483 "gpu": self.gpu,
484 "rank": self.rank,
485 "world_size": self.world_size,
486 "dtype": str(self.param_dtype),
487 })
488 return ret
490 async def get_rest_args(
491 self,
492 data_json: Dict[str, str]
493 ) -> Dict[str, Any]:
494 if data_json is None:
495 raise ValueError("Missing JSON body")
496 imgs_base64 = data_json.get("imgs", None)
497 imgs = []
498 if imgs_base64 is not None:
499 for img_base64 in imgs_base64:
500 img = base64_to_img(img_base64)
501 imgs.append(img)
502 prompt = data_json.get("prompt", None)
503 if prompt is None:
504 raise ValueError("Missing 'prompt' parameter")
505 neg_prompt = data_json.get("neg_prompt", "")
506 height = int(data_json.get("height", 400))
507 width = int(data_json.get("width", 640))
508 steps = int(data_json.get("sampling_steps", 20))
509 return {
510 "task": self.model_name,
511 "args": {
512 "imgs": [img] if img is not None else [],
513 "prompt": prompt,
514 "neg_prompt": neg_prompt,
515 "height": height,
516 "width": width,
517 "sampling_steps": steps,
518 }
519 }
522async def main() -> None:
523 bagel = BagelGeneration()
525 img_size = (720, 1280)
526 img_size = (1024, 1024)
527 img_size = (512, 512)
528 num_steps = 50
530 # Generate original image 1
531 img_prompt = "A full-body photo of a young woman standing on a white background. "
532 img_prompt += "She has straight blonde hair, parted slightly to the side, and is smiling at the camera. "
533 img_prompt += "She is wearing a white sleeveless tank top, blue skinny jeans, and blue socks. "
534 img_prompt += "Her arms are relaxed at her sides, and she is standing with her feet together. "
535 img_prompt += "The lighting is bright and even, with no visible shadows, giving the image a clean studio feel."
537 output_image_0 = await bagel.generate(
538 width=img_size[0],
539 height=img_size[1],
540 prompt=img_prompt,
541 imgs=[],
542 sampling_steps=num_steps,
543 )
544 output_image_0.save("output/output_image_0.png")
546 # Generate original image 2
547 img_prompt = "A full-body photo of a young man standing on a white studio background. "
548 img_prompt += "He has short black hair styled upward and is smiling slightly at the camera. "
549 img_prompt += "He is wearing a fitted light gray V-neck T-shirt, dark blue jeans, and brown suede shoes. "
550 img_prompt += "His arms are relaxed by his sides, and his stance is casual with feet slightly apart. "
551 img_prompt += "The lighting is bright and even, with no shadows, creating a clean and professional look."
552 output_image_1 = await bagel.generate(
553 width=img_size[0],
554 height=img_size[1],
555 prompt=img_prompt,
556 imgs=[],
557 sampling_steps=num_steps,
558 )
559 output_image_1.save("output/output_image_1.png")
561 # Merge the two images
562 img_prompt += "Using the two people wearing from the original two images."
563 img_prompt += "A photorealistic podcast setup featuring the woman in the right and the man in the left sitting "
564 img_prompt += "across each other at a wooden table in a professional recording studio."
565 img_prompt += "The studio has red acoustic panels on the walls, warm lighting, and a large off screen in the "
566 img_prompt += "background."
567 img_prompt += "Both wear headphones and speak into high-quality podcast microphones mounted on adjustable arms."
568 img_prompt += "Scene captured from a slightly elevated front-facing perspective, showing their upper bodies and "
569 img_prompt += "expressive gestures as they engage in conversation."
570 img_prompt += "Table equipped with coffee mugs, water bottles, and recording equipment."
572 output_image_2 = await bagel.generate(
573 img_size[0],
574 img_size[1],
575 img_prompt,
576 imgs=[
577 output_image_0,
578 output_image_1,
579 ],
580 sampling_steps=num_steps,
581 )
582 output_image_2.save("output/output_image_2.png")
585if __name__ == "__main__":
586 asyncio.run(main())