funasr/bin/train_ds.py
@@ -84,6 +84,8 @@ dist.init_process_group(backend=kwargs.get("backend", "nccl"), init_method="env://") torch.cuda.set_device(local_rank) # rank = dist.get_rank() logging.info("Build model, frontend, tokenizer") device = kwargs.get("device", "cuda") kwargs["device"] = "cpu" @@ -124,6 +126,7 @@ use_ddp=use_ddp, use_fsdp=use_fsdp, device=kwargs["device"], excludes=kwargs.get("excludes", None), output_dir=kwargs.get("output_dir", "./exp"), **kwargs.get("train_conf"), )