(args)
| 428 | return s1, s2, save_path_sum |
| 429 | |
| 430 | def generate(args): |
| 431 | rank = int(os.getenv("RANK", 0)) |
| 432 | world_size = int(os.getenv("WORLD_SIZE", 1)) |
| 433 | local_rank = int(os.getenv("LOCAL_RANK", 0)) |
| 434 | device = local_rank |
| 435 | _init_logging(rank) |
| 436 | |
| 437 | if args.offload_model is None: |
| 438 | args.offload_model = False if world_size > 1 else True |
| 439 | logging.info( |
| 440 | f"offload_model is not specified, set to {args.offload_model}.") |
| 441 | if world_size > 1: |
| 442 | torch.cuda.set_device(local_rank) |
| 443 | dist.init_process_group( |
| 444 | backend="nccl", |
| 445 | init_method="env://", |
| 446 | rank=rank, |
| 447 | world_size=world_size) |
| 448 | else: |
| 449 | assert not ( |
| 450 | args.t5_fsdp or args.dit_fsdp |
| 451 | ), f"t5_fsdp and dit_fsdp are not supported in non-distributed environments." |
| 452 | assert not ( |
| 453 | args.ulysses_size > 1 or args.ring_size > 1 |
| 454 | ), f"context parallel are not supported in non-distributed environments." |
| 455 | |
| 456 | if args.ulysses_size > 1 or args.ring_size > 1: |
| 457 | 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." |
| 458 | from xfuser.core.distributed import ( |
| 459 | init_distributed_environment, |
| 460 | initialize_model_parallel, |
| 461 | ) |
| 462 | init_distributed_environment( |
| 463 | rank=dist.get_rank(), world_size=dist.get_world_size()) |
| 464 | |
| 465 | initialize_model_parallel( |
| 466 | sequence_parallel_degree=dist.get_world_size(), |
| 467 | ring_degree=args.ring_size, |
| 468 | ulysses_degree=args.ulysses_size, |
| 469 | ) |
| 470 | |
| 471 | # TODO: use prompt refine |
| 472 | # if args.use_prompt_extend: |
| 473 | # if args.prompt_extend_method == "dashscope": |
| 474 | # prompt_expander = DashScopePromptExpander( |
| 475 | # model_name=args.prompt_extend_model, |
| 476 | # is_vl="i2v" in args.task or "flf2v" in args.task) |
| 477 | # elif args.prompt_extend_method == "local_qwen": |
| 478 | # prompt_expander = QwenPromptExpander( |
| 479 | # model_name=args.prompt_extend_model, |
| 480 | # is_vl="i2v" in args.task, |
| 481 | # device=rank) |
| 482 | # else: |
| 483 | # raise NotImplementedError( |
| 484 | # f"Unsupport prompt_extend_method: {args.prompt_extend_method}") |
| 485 | |
| 486 | cfg = WAN_CONFIGS[args.task] |
| 487 | if args.ulysses_size > 1: |
no test coverage detected