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

Function generate

generate_multitalk.py:430–633  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

428 return s1, s2, save_path_sum
429
430def 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:

Callers 1

Calls 10

generateMethod · 0.95
save_video_ffmpegFunction · 0.90
_init_loggingFunction · 0.70
custom_initFunction · 0.70
audio_prepare_multiFunction · 0.70
get_embeddingFunction · 0.70
audio_prepare_singleFunction · 0.70
process_tts_singleFunction · 0.70
process_tts_multiFunction · 0.70

Tested by

no test coverage detected