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"