Skip to content

Commit 2e6f7c9

Browse files
Add original_path and model_name parameters to cross-site validation functions
1 parent 9f9261d commit 2e6f7c9

4 files changed

Lines changed: 85 additions & 25 deletions

File tree

monai/nvflare/nnunet_executor.py

Lines changed: 54 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -130,7 +130,8 @@ def __init__(
130130
train_extra_configs=None,
131131
exclude_vars=None,
132132
continue_training=False,
133-
label_dict = None
133+
label_dict = None,
134+
original_path = None,
134135
):
135136
super().__init__()
136137

@@ -159,6 +160,7 @@ def __init__(
159160
self.cross_site_validation_task_name = cross_site_validation_task_name
160161
self.monai_deploy_config = monai_deploy_config
161162
self.label_dict = label_dict
163+
self.original_path = original_path
162164

163165
def handle_event(self, event_type: str, fl_ctx: FLContext):
164166
if event_type == EventType.START_RUN:
@@ -173,19 +175,20 @@ def initialize(self, fl_ctx: FLContext):
173175

174176
with open("init_logfile_out.log", "w") as f_o:
175177
with open("init_logfile_err.log", "w") as f_e:
176-
subprocess.call(
177-
[
178-
sys.executable,
179-
"-m",
180-
"pip",
181-
"install",
182-
"--user",
183-
"-r",
184-
str(Path(self.custom_app_dir).joinpath("requirements.txt")),
185-
],
186-
stdout=f_o,
187-
stderr=f_e,
188-
)
178+
...
179+
#subprocess.call(
180+
# [
181+
# sys.executable,
182+
# "-m",
183+
# "pip",
184+
# "install",
185+
# "--user",
186+
# "-r",
187+
# str(Path(self.custom_app_dir).joinpath("requirements.txt")),
188+
# ],
189+
# stdout=f_o,
190+
# stderr=f_e,
191+
#)
189192

190193
def execute(self, task_name: str, shareable: Shareable, fl_ctx: FLContext, abort_signal: Signal) -> Shareable:
191194
self.run_dir = fl_ctx.get_engine().get_workspace().get_run_dir(fl_ctx.get_job_id())
@@ -223,6 +226,10 @@ def prepare_dataset(self) -> Shareable:
223226
dataset_name_or_id = self.nnunet_config["dataset_name_or_id"]
224227
if isinstance(dataset_name_or_id, dict):
225228
dataset_name_or_id = dataset_name_or_id.get(self.client_name, None)
229+
if isinstance(dataset_name_or_id, dict):
230+
dataset_name = dataset_name_or_id.get("name", None)
231+
dataset_name_or_id = dataset_name_or_id.get("id", None)
232+
226233
data_list = prepare_data_folder(
227234
data_dir=self.data_dir,
228235
nnunet_root_dir=self.nnunet_root_folder,
@@ -237,6 +244,7 @@ def prepare_dataset(self) -> Shareable:
237244
subfolder_suffix=self.subfolder_suffix,
238245
trainer_class_name=nnunet_trainer_name,
239246
modality_list=self.modality_list,
247+
dataset_name=dataset_name
240248
)
241249

242250
outgoing_dxo = DXO(data_kind=DataKind.COLLECTION, data=data_list, meta={})
@@ -274,6 +282,10 @@ def plan_and_preprocess(self):
274282
dataset_name_or_id = self.nnunet_config["dataset_name_or_id"]
275283
if isinstance(dataset_name_or_id, dict):
276284
dataset_name_or_id = dataset_name_or_id.get(self.client_name, None)
285+
if isinstance(dataset_name_or_id, dict):
286+
dataset_name = dataset_name_or_id.get("name", None)
287+
dataset_name_or_id = dataset_name_or_id.get("id", None)
288+
277289
nnunet_plans = plan_and_preprocess(
278290
self.nnunet_root_folder,
279291
dataset_name_or_id,
@@ -282,6 +294,7 @@ def plan_and_preprocess(self):
282294
self.tracking_uri,
283295
nnunet_plans_name=nnunet_plans_name,
284296
trainer_class_name=nnunet_trainer_name,
297+
dataset_name=dataset_name,
285298
)
286299

287300
outgoing_dxo = DXO(data_kind=DataKind.COLLECTION, data=nnunet_plans, meta={})
@@ -301,6 +314,10 @@ def preprocess(self):
301314
dataset_name_or_id = self.nnunet_config["dataset_name_or_id"]
302315
if isinstance(dataset_name_or_id, dict):
303316
dataset_name_or_id = dataset_name_or_id.get(self.client_name, None)
317+
if isinstance(dataset_name_or_id, dict):
318+
dataset_name = dataset_name_or_id.get("name", None)
319+
dataset_name_or_id = dataset_name_or_id.get("id", None)
320+
304321
nnunet_plans = preprocess(
305322
self.nnunet_root_folder,
306323
dataset_name_or_id,
@@ -324,6 +341,10 @@ def train(self):
324341
dataset_name_or_id = self.nnunet_config["dataset_name_or_id"]
325342
if isinstance(dataset_name_or_id, dict):
326343
dataset_name_or_id = dataset_name_or_id.get(self.client_name, None)
344+
if isinstance(dataset_name_or_id, dict):
345+
dataset_name = dataset_name_or_id.get("name", None)
346+
dataset_name_or_id = dataset_name_or_id.get("id", None)
347+
327348
validation_summary = train(
328349
self.nnunet_root_folder,
329350
trainer_class_name=nnunet_trainer_name,
@@ -354,6 +375,10 @@ def prepare_bundle(self):
354375
dataset_name_or_id = self.nnunet_config["dataset_name_or_id"]
355376
if isinstance(dataset_name_or_id, dict):
356377
dataset_name_or_id = dataset_name_or_id.get(self.client_name, None)
378+
if isinstance(dataset_name_or_id, dict):
379+
dataset_name = dataset_name_or_id.get("name", None)
380+
dataset_name_or_id = dataset_name_or_id.get("id", None)
381+
357382
bundle_config = {
358383
"bundle_root": self.bundle_root,
359384
"tracking_uri": self.tracking_uri,
@@ -363,6 +388,7 @@ def prepare_bundle(self):
363388
"nnunet_trainer_class_name": nnunet_trainer_name,
364389
"dataset_name_or_id": dataset_name_or_id,
365390
"label_dict": self.label_dict,
391+
"dataset_name": dataset_name,
366392
}
367393

368394
bundle_config = prepare_bundle(bundle_config, self.train_extra_configs)
@@ -386,6 +412,10 @@ def finalize_bundle(self):
386412
dataset_name_or_id = self.nnunet_config["dataset_name_or_id"]
387413
if isinstance(dataset_name_or_id, dict):
388414
dataset_name_or_id = dataset_name_or_id.get(self.client_name, None)
415+
if isinstance(dataset_name_or_id, dict):
416+
dataset_name = dataset_name_or_id.get("name", None)
417+
dataset_name_or_id = dataset_name_or_id.get("id", None)
418+
389419
validation_summary = finalize_bundle(
390420
self.bundle_root,
391421
self.nnunet_root_folder,
@@ -395,7 +425,8 @@ def finalize_bundle(self):
395425
client_name=self.client_name,
396426
tracking_uri=self.tracking_uri,
397427
nnunet_plans_name=nnunet_plans_name,
398-
dataset_name_or_id=dataset_name_or_id
428+
dataset_name_or_id=dataset_name_or_id,
429+
dataset_name=dataset_name,
399430
)
400431
outgoing_dxo = DXO(data_kind=DataKind.COLLECTION, data=validation_summary, meta={})
401432
return outgoing_dxo.to_shareable()
@@ -414,18 +445,25 @@ def run_cross_site_validation(self):
414445
dataset_name_or_id = self.nnunet_config["dataset_name_or_id"]
415446
if isinstance(dataset_name_or_id, dict):
416447
dataset_name_or_id = dataset_name_or_id.get(self.client_name, None)
448+
if isinstance(dataset_name_or_id, dict):
449+
dataset_name = dataset_name_or_id.get("name", None)
450+
dataset_name_or_id = dataset_name_or_id.get("id", None)
451+
417452
validation_summary = run_cross_site_validation(
418453
self.nnunet_root_folder,
419454
dataset_name_or_id,
420455
self.monai_deploy_config["app_path"],
421456
self.monai_deploy_config["app_model_path"],
422457
self.monai_deploy_config["app_output_path"],
458+
self.monai_deploy_config["model_name"],
423459
trainer_class_name=nnunet_trainer_name,
424460
fold=0,
425461
experiment_name=self.nnunet_config["experiment_name"],
426462
client_name=self.client_name,
427463
tracking_uri=self.tracking_uri,
428-
nnunet_plans_name=nnunet_plans_name
464+
nnunet_plans_name=nnunet_plans_name,
465+
dataset_name=dataset_name,
466+
original_path=self.original_path,
429467
)
430468

431469
outgoing_dxo = DXO(data_kind=DataKind.COLLECTION, data=validation_summary, meta={})

monai/nvflare/nvflare_generate_job_configs.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1322,6 +1322,7 @@ def cross_site_validation_config(clients, experiment, root_dir, script_dir, nvfl
13221322
"monai_deploy_config": {
13231323
"app_path": clients[client_id]["app_path"],
13241324
"app_model_path": clients[client_id]["app_model_path"],
1325+
"model_name": clients[client_id]["model_name"],
13251326
"app_output_path": clients[client_id]["app_output_path"],
13261327
},
13271328
"client_name": clients[client_id]["client_name"],
@@ -1341,6 +1342,11 @@ def cross_site_validation_config(clients, experiment, root_dir, script_dir, nvfl
13411342
if "bundle_root" in clients[client_id]:
13421343
client["executors"][0]["executor"]["args"]["bundle_root"] = clients[client_id]["bundle_root"]
13431344

1345+
if "original_path" in clients[client_id]:
1346+
client["executors"][0]["executor"]["args"]["original_path"] = clients[client_id][
1347+
"original_path"
1348+
]
1349+
13441350
Path(root_dir).joinpath(task_name).joinpath(f"{task_name}-client-{client_id}").mkdir(
13451351
parents=True, exist_ok=True
13461352
)

monai/nvflare/nvflare_nnunet.py

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -541,7 +541,7 @@ def prepare_bundle(bundle_config, train_extra_configs=None):
541541
None
542542
"""
543543

544-
prepare_bundle_api(bundle_config, train_extra_configs=train_extra_configs, is_federated=True)
544+
return prepare_bundle_api(bundle_config, train_extra_configs=train_extra_configs, is_federated=True)
545545

546546

547547

@@ -644,9 +644,9 @@ def finalize_bundle(bundle_root, nnunet_root_dir=None, validate_with_nnunet=True
644644
return validation_summary_dict
645645

646646

647-
def run_cross_site_validation(nnunet_root_dir, dataset_name_or_id, app_path, app_model_path, app_output_path, trainer_class_name="nnUNetTrainer", fold=0,
647+
def run_cross_site_validation(nnunet_root_dir, dataset_name_or_id, app_path, app_model_path, app_output_path, model_name, trainer_class_name="nnUNetTrainer", fold=0,
648648
experiment_name=None, client_name=None, tracking_uri=None,
649-
nnunet_plans_name="nnUNetPlans", mlflow_token=None, skip_prediction=False, dataset_name=None):
649+
nnunet_plans_name="nnUNetPlans", mlflow_token=None, skip_prediction=False, dataset_name=None, original_path=None):
650650

651651
validation_summary_dict, labels = cross_site_evaluation_api(
652652
nnunet_root_dir,
@@ -658,6 +658,8 @@ def run_cross_site_validation(nnunet_root_dir, dataset_name_or_id, app_path, app
658658
fold=fold,
659659
nnunet_plans_name=nnunet_plans_name,
660660
skip_prediction=skip_prediction,
661+
original_path = original_path,
662+
661663
)
662664
if mlflow_token is not None:
663665
os.environ["MLFLOW_TRACKING_TOKEN"] = mlflow_token
@@ -670,28 +672,28 @@ def run_cross_site_validation(nnunet_root_dir, dataset_name_or_id, app_path, app
670672
print(e)
671673
mlflow.set_experiment(experiment_id=(mlflow.get_experiment_by_name(experiment_name).experiment_id))
672674

673-
run_name = f"run_cross_site_validation_{client_name}"
675+
run_name = f"run_cross_site_validation_{client_name}_Model_{model_name}"
674676

675677
runs = mlflow.search_runs(
676678
experiment_names=[experiment_name],
677679
filter_string=f"tags.mlflow.runName = '{run_name}'",
678680
order_by=["start_time DESC"]
679681
)
680-
tags = {"client": client_name}
682+
tags = {"client": client_name,"model": model_name}
681683
if dataset_name is not None:
682684
tags["dataset_name"] = dataset_name
683685

684686

685687
if len(runs) == 0:
686-
with mlflow.start_run(run_name=f"run_{client_name}", tags={"client": client_name}):
688+
with mlflow.start_run(run_name=run_name, tags=tags):
687689
mlflow.log_dict(validation_summary_dict, "validation_summary.json")
688690
for label in validation_summary_dict["mean"]:
689691
for metric in validation_summary_dict["mean"][label]:
690692
label_name = labels[label]
691693
mlflow.log_metric(f"{label_name}_{metric}", float(validation_summary_dict["mean"][label][metric]))
692694

693695
else:
694-
with mlflow.start_run(run_id=runs.iloc[0].run_id, tags={"client": client_name}):
696+
with mlflow.start_run(run_id=runs.iloc[0].run_id, tags=tags):
695697
mlflow.log_dict(validation_summary_dict, "validation_summary.json")
696698
for label in validation_summary_dict["mean"]:
697699
for metric in validation_summary_dict["mean"][label]:

monai/nvflare/utils.py

Lines changed: 16 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -256,7 +256,7 @@ def prepare_data_folder_api(data_dir,
256256
return data_list
257257

258258

259-
def cross_site_evaluation_api(nnunet_root_dir, dataset_name_or_id, app_path, app_model_path, app_output_path, fold=0, trainer_class_name="nnUNetTrainer", nnunet_plans_name="nnUNetPlans", skip_prediction=False):
259+
def cross_site_evaluation_api(nnunet_root_dir, dataset_name_or_id, app_path, app_model_path, app_output_path, fold=0, trainer_class_name="nnUNetTrainer", nnunet_plans_name="nnUNetPlans", skip_prediction=False, original_path=None):
260260
data_src_cfg = os.path.join(nnunet_root_dir, f"Task{dataset_name_or_id}_data_src_cfg.yaml")
261261

262262
runner = nnUNetV2Runner(input_config=data_src_cfg, trainer_class_name=trainer_class_name, work_dir=nnunet_root_dir)
@@ -285,14 +285,28 @@ def cross_site_evaluation_api(nnunet_root_dir, dataset_name_or_id, app_path, app
285285
filename = Path(app_input_path).name
286286

287287
new_id = None
288+
updated_image_path = False
288289
for case in nnunet_datalist["training"]:
289290
if case["image"].endswith(filename):
290291
new_id = case["new_name"]
291292
break
293+
if filename.startswith(case["new_name"]+"_"):
294+
new_id = case["new_name"]
295+
data["image"] = os.path.join(original_path, Path(case["image"]).name)
296+
app_input_path = data["image"]
297+
updated_image_path = True
298+
break
292299
if new_id in nnunet_splits[fold]["val"]:
293-
id_mapping[Path(data["image"]).name.split("_")[0].split(".")[0]] = new_id
300+
if updated_image_path:
301+
id_mapping[Path(data["image"]).name.split(".")[0]] = new_id
302+
else:
303+
id_mapping[Path(data["image"]).name.split("_")[0].split(".")[0]] = new_id
294304
if skip_prediction:
295305
continue
306+
print(f"Processing case: {new_id}")
307+
print(f"App input path: {app_input_path}")
308+
mapped_filename = Path(data["image"]).name.split(".")[0]
309+
print(f"Mapping: {mapped_filename} -> {new_id}")
296310
subprocess.run(
297311
[
298312
"python",

0 commit comments

Comments
 (0)