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

1from __future__ import annotations 

2 

3import logging 

4import os 

5import asyncio 

6 

7from typing import override 

8from typing import List 

9from typing import Optional 

10from typing import Dict 

11from typing import Any 

12 

13from PIL import Image 

14from copy import deepcopy 

15 

16import torch 

17import torch.distributed as dist 

18from torch import inference_mode 

19 

20from wrapper_model import ModelGeneration 

21from image_utils import base64_to_img 

22 

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 

32 

33from data.transforms import ImageTransform 

34from data.data_utils import add_special_tokens 

35 

36from accelerate import load_checkpoint_and_dispatch 

37from accelerate import init_empty_weights 

38 

39from xfuser.config import EngineConfig 

40 

41 

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

49 

50 self.engine_config = engine_config 

51 self.param_dtype = param_dtype 

52 

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 

59 

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 

69 

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

81 

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

87 

88 self.device_id = self.local_rank 

89 self.device = torch.device(f"cuda:{self.device_id}") 

90 

91 torch.cuda.set_device(self.local_rank) 

92 self.load_timer.end("torch_dist") 

93 

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

97 

98 def load_model(self) -> None: 

99 assert torch.cuda.is_available() 

100 

101 model_path = "BAGEL-7B-MoT" # TODO temporary 

102 

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

109 

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

117 

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

124 

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 ) 

138 

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) 

144 

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

158 

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

168 

169 # Image Transform 

170 self.vae_transform = ImageTransform(1024, 512, 16) 

171 self.vit_transform = ImageTransform(980, 224, 14) 

172 

173 def init_model_parallelism(self) -> None: 

174 if self.world_size > 1: 

175 logging.warning("Parallelism not supported.") 

176 

177 def model_compile(self) -> None: 

178 if not self.torch_compile: 

179 return 

180 

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

188 

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) 

213 

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 

219 

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'] 

231 

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) 

243 

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) 

253 

254 # Output 

255 gen_context['past_key_values'] = past_key_values 

256 gen_context['kv_lens'] = kv_lens 

257 gen_context['ropes'] = ropes 

258 

259 return gen_context 

260 

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 

281 

282 def _assert_model_init(self) -> None: 

283 super()._assert_model_init() 

284 if self.model is None: 

285 raise ValueError("Model not initialized.") 

286 

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) 

299 

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) 

314 

315 self._assert_model_init() 

316 assert self.model is not None 

317 assert self.vae_transform is not None 

318 

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

320 

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 

330 

331 image_shape = (height, width) 

332 

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) 

343 

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) 

354 

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) 

360 

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

382 

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

394 

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

410 

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] 

417 

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

468 

469 unpacked_latent = x_t.split((packed_seqlens - 2).tolist()) 

470 

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

479 

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 

489 

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 } 

520 

521 

522async def main() -> None: 

523 bagel = BagelGeneration() 

524 

525 img_size = (720, 1280) 

526 img_size = (1024, 1024) 

527 img_size = (512, 512) 

528 num_steps = 50 

529 

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

536 

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

545 

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

560 

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

571 

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

583 

584 

585if __name__ == "__main__": 

586 asyncio.run(main())