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

Function generate

generate_infinitetalk.py:453–658  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

451 return s1, s2, save_path_sum
452
453def generate(args):
454 rank = int(os.getenv("RANK", 0))
455 world_size = int(os.getenv("WORLD_SIZE", 1))
456 local_rank = int(os.getenv("LOCAL_RANK", 0))
457 device = local_rank
458 _init_logging(rank)
459
460 if args.offload_model is None:
461 args.offload_model = False if world_size > 1 else True
462 logging.info(
463 f"offload_model is not specified, set to {args.offload_model}.")
464 if world_size > 1:
465 torch.cuda.set_device(local_rank)
466 dist.init_process_group(
467 backend="nccl",
468 init_method="env://",
469 rank=rank,
470 world_size=world_size)
471 else:
472 assert not (
473 args.t5_fsdp or args.dit_fsdp
474 ), f"t5_fsdp and dit_fsdp are not supported in non-distributed environments."
475 assert not (
476 args.ulysses_size > 1 or args.ring_size > 1
477 ), f"context parallel are not supported in non-distributed environments."
478
479 if args.ulysses_size > 1 or args.ring_size > 1:
480 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."
481 from xfuser.core.distributed import (
482 init_distributed_environment,
483 initialize_model_parallel,
484 )
485 init_distributed_environment(
486 rank=dist.get_rank(), world_size=dist.get_world_size())
487
488 initialize_model_parallel(
489 sequence_parallel_degree=dist.get_world_size(),
490 ring_degree=args.ring_size,
491 ulysses_degree=args.ulysses_size,
492 )
493
494 # TODO: use prompt refine
495 # if args.use_prompt_extend:
496 # if args.prompt_extend_method == "dashscope":
497 # prompt_expander = DashScopePromptExpander(
498 # model_name=args.prompt_extend_model,
499 # is_vl="i2v" in args.task or "flf2v" in args.task)
500 # elif args.prompt_extend_method == "local_qwen":
501 # prompt_expander = QwenPromptExpander(
502 # model_name=args.prompt_extend_model,
503 # is_vl="i2v" in args.task,
504 # device=rank)
505 # else:
506 # raise NotImplementedError(
507 # f"Unsupport prompt_extend_method: {args.prompt_extend_method}")
508
509 cfg = WAN_CONFIGS[args.task]
510 if args.ulysses_size > 1:

Callers 1

Calls 11

generate_infinitetalkMethod · 0.95
is_videoFunction · 0.90
shot_detectFunction · 0.90
split_wav_librosaFunction · 0.90
save_video_ffmpegFunction · 0.90
_init_loggingFunction · 0.70
custom_initFunction · 0.70
audio_prepare_multiFunction · 0.70
audio_prepare_singleFunction · 0.70
get_embeddingFunction · 0.70

Tested by

no test coverage detected