Skip to content

Commit 2d09694

Browse files
Add dataset_name_or_id to nnunet_config and update data source configuration
1 parent 2cf94c7 commit 2d09694

3 files changed

Lines changed: 11 additions & 2 deletions

File tree

monai/nvflare/nnunet_executor.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -335,6 +335,7 @@ def prepare_bundle(self):
335335
"mlflow_run_name": self.client_name,
336336
"nnunet_plans_identifier": nnunet_plans_name,
337337
"nnunet_trainer_class_name": nnunet_trainer_name,
338+
"dataset_name_or_id": self.nnunet_config["dataset_name_or_id"]
338339
}
339340

340341
prepare_bundle(bundle_config, self.train_extra_configs)

monai/nvflare/nvflare_generate_job_configs.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -833,7 +833,7 @@ def prepare_bundle_config(clients, experiment, root_dir, script_dir, nvflare_exe
833833
"executor": {
834834
"path": "monai.nvflare.nnunet_executor.nnUNetExecutor",
835835
"args": {
836-
"nnunet_config": {"experiment_name": experiment["experiment_name"]},
836+
"nnunet_config": {"experiment_name": experiment["experiment_name"], "dataset_name_or_id": experiment["dataset_name_or_id"]},
837837
"client_name": clients[client_id]["client_name"],
838838
"tracking_uri": experiment["tracking_uri"],
839839
},

monai/nvflare/nvflare_nnunet.py

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -83,6 +83,8 @@ def run_job(sess, task_name, job_folder, clients=None):
8383
f.write("\n}")
8484

8585
job_id = sess.submit_job(str(Path(job_folder).joinpath(task_name)))
86+
87+
return job_id
8688

8789
def train(
8890
nnunet_root_dir,
@@ -681,7 +683,8 @@ def prepare_bundle(bundle_config, train_extra_configs=None):
681683
train_config["mlflow_experiment_name"] = bundle_config["mlflow_experiment_name"]
682684
train_config["mlflow_run_name"] = bundle_config["mlflow_run_name"]
683685

684-
train_config["data_src_cfg"] = "$@nnunet_root_folder+'/data_src_cfg.yaml'"
686+
train_config["dataset_name_or_id"] = bundle_config["dataset_name_or_id"]
687+
train_config["data_src_cfg"] = "$@nnunet_root_folder+'/Task'+@dataset_name_or_id+'_data_src_cfg.yaml'"
685688
train_config["nnunet_root_folder"] = "."
686689
train_config["runner"] = {
687690
"_target_": "nnUNetV2Runner",
@@ -710,6 +713,11 @@ def prepare_bundle(bundle_config, train_extra_configs=None):
710713
]
711714
else:
712715
train_config["initialize"] = ["$monai.utils.set_determinism(seed=123)", "$@runner.dataset_name_or_id"]
716+
717+
if train_extra_configs is not None:
718+
for key in train_extra_configs:
719+
if key != "resume_epoch":
720+
train_config[key] = train_extra_configs[key]
713721

714722
if "Val_Dice" in train_config["val_key_metric"]:
715723
train_config["val_key_metric"] = {"Val_Dice_Local": train_config["val_key_metric"]["Val_Dice"]}

0 commit comments

Comments
 (0)