MCPcopy Create free account
hub / github.com/MeiGen-AI/MultiTalk / run_graio_demo

Function run_graio_demo

app.py:421–779  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

419 return s1, s2, save_path_sum
420
421def 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

Callers 1

app.pyFile · 0.85

Calls 3

_init_loggingFunction · 0.70
custom_initFunction · 0.70

Tested by

no test coverage detected