(args)
| 419 | return s1, s2, save_path_sum |
| 420 | |
| 421 | def run_graio_demo(args): |
| 422 | rank = int(os.getenv("RANK", 0)) |
| 423 | world_size = int(os.getenv("WORLD_SIZE", 1)) |
| 424 | local_rank = int(os.getenv("LOCAL_RANK", 0)) |
| 425 | device = local_rank |
| 426 | _init_logging(rank) |
| 427 | |
| 428 | if args.offload_model is None: |
| 429 | args.offload_model = False if world_size > 1 else True |
| 430 | logging.info( |
| 431 | f"offload_model is not specified, set to {args.offload_model}.") |
| 432 | if world_size > 1: |
| 433 | torch.cuda.set_device(local_rank) |
| 434 | dist.init_process_group( |
| 435 | backend="nccl", |
| 436 | init_method="env://", |
| 437 | rank=rank, |
| 438 | world_size=world_size) |
| 439 | else: |
| 440 | assert not ( |
| 441 | args.t5_fsdp or args.dit_fsdp |
| 442 | ), f"t5_fsdp and dit_fsdp are not supported in non-distributed environments." |
| 443 | assert not ( |
| 444 | args.ulysses_size > 1 or args.ring_size > 1 |
| 445 | ), f"context parallel are not supported in non-distributed environments." |
| 446 | |
| 447 | if args.ulysses_size > 1 or args.ring_size > 1: |
| 448 | assert args.ulysses_size * args.ring_size == world_size, f"The number of ulysses_size and ring_size should be equal to the world size." |
| 449 | from xfuser.core.distributed import ( |
| 450 | init_distributed_environment, |
| 451 | initialize_model_parallel, |
| 452 | ) |
| 453 | init_distributed_environment( |
| 454 | rank=dist.get_rank(), world_size=dist.get_world_size()) |
| 455 | |
| 456 | initialize_model_parallel( |
| 457 | sequence_parallel_degree=dist.get_world_size(), |
| 458 | ring_degree=args.ring_size, |
| 459 | ulysses_degree=args.ulysses_size, |
| 460 | ) |
| 461 | |
| 462 | |
| 463 | cfg = WAN_CONFIGS[args.task] |
| 464 | if args.ulysses_size > 1: |
| 465 | assert cfg.num_heads % args.ulysses_size == 0, f"`{cfg.num_heads=}` cannot be divided evenly by `{args.ulysses_size=}`." |
| 466 | |
| 467 | logging.info(f"Generation job args: {args}") |
| 468 | logging.info(f"Generation model config: {cfg}") |
| 469 | |
| 470 | if dist.is_initialized(): |
| 471 | base_seed = [args.base_seed] if rank == 0 else [None] |
| 472 | dist.broadcast_object_list(base_seed, src=0) |
| 473 | args.base_seed = base_seed[0] |
| 474 | |
| 475 | assert args.task == "multitalk-14B", 'You should choose multitalk in args.task.' |
| 476 | |
| 477 | |
| 478 |
no test coverage detected