MCPcopy Create free account
hub / github.com/modelscope/FunASR / main

Function main

funasr/bin/train_ds.py:64–255  ·  view source on GitHub ↗

Main. Args: **kwargs: Additional keyword arguments.

(**kwargs)

Source from the content-addressed store, hash-verified

62
63
64def main(**kwargs):
65
66 # set random seed
67 """Main.
68
69 Args:
70 **kwargs: Additional keyword arguments.
71 """
72 set_all_random_seed(kwargs.get("seed", 0))
73 torch.backends.cudnn.enabled = kwargs.get("cudnn_enabled", torch.backends.cudnn.enabled)
74 torch.backends.cudnn.benchmark = kwargs.get("cudnn_benchmark", torch.backends.cudnn.benchmark)
75 torch.backends.cudnn.deterministic = kwargs.get("cudnn_deterministic", True)
76 # open tf32
77 torch.backends.cuda.matmul.allow_tf32 = kwargs.get("enable_tf32", True)
78
79 rank = int(os.environ.get("RANK", 0))
80 local_rank = int(os.environ.get("LOCAL_RANK", 0))
81 world_size = int(os.environ.get("WORLD_SIZE", 1))
82
83 if local_rank == 0:
84 tables.print()
85
86 use_ddp = world_size > 1
87 use_fsdp = kwargs.get("use_fsdp", False)
88 use_deepspeed = kwargs.get("use_deepspeed", False)
89 if use_deepspeed:
90 logging.info(f"use_deepspeed: {use_deepspeed}")
91 deepspeed.init_distributed(dist_backend=kwargs.get("backend", "nccl"))
92 elif use_ddp or use_fsdp:
93 logging.info(f"use_ddp: {use_ddp}, use_fsdp: {use_fsdp}")
94 dist.init_process_group(
95 backend=kwargs.get("backend", "nccl"),
96 init_method="env://",
97 )
98 torch.cuda.set_device(local_rank)
99
100 # rank = dist.get_rank()
101
102 logging.info("Build model, frontend, tokenizer")
103 device = kwargs.get("device", "cuda")
104 kwargs["device"] = "cpu"
105 model = AutoModel(**kwargs)
106
107 # save config.yaml
108 if rank == 0:
109 prepare_model_dir(**kwargs)
110
111 # parse kwargs
112 kwargs = model.kwargs
113 kwargs["device"] = device
114 tokenizer = kwargs["tokenizer"]
115 frontend = kwargs["frontend"]
116 model = model.model
117 del kwargs["model"]
118
119 # freeze_param
120 freeze_param = kwargs.get("freeze_param", None)
121 if freeze_param is not None:

Callers 1

main_hydraFunction · 0.70

Calls 15

warp_modelMethod · 0.95
warp_optim_schedulerMethod · 0.95
resume_checkpointMethod · 0.95
train_epochMethod · 0.95
validate_epochMethod · 0.95
save_checkpointMethod · 0.95
closeMethod · 0.95
set_all_random_seedFunction · 0.90
AutoModelClass · 0.90
prepare_model_dirFunction · 0.90
model_summaryFunction · 0.90

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…