OSCR

Traumatic brain injury recovery prediction by harmonizing real brain CT and synthetic brain MRI: a pilot study.

Code ↔ Paper

2 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 2 matches
  1. [1] § Materials and methods › Statistical analysis ↔ Stage_2_CLS/5fold/analysis.ipynb, lines 270–397 · score 0.73 · receiver operating characteristic, DeLong, correlated, curves, ROC, AUCs
  2. [2] § Materials and methods › GAN as harmonization technique for synthetic imaging rendering ↔ Stage_1_FPGAN/model.py, lines 22–59 · score 0.60 · residual blocks, domain information, network, layers, models

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Jupyter notebook · 3,122 lines · 119 KB · no license · 1 match

  1. # %%
  2. %cd
  3. # %% [markdown]
  4. # # Binary classification
  5. # %%
  6. import logging
  7. import os
  8. import sys
  9. from pathlib import Path
  10. import pandas as pd
  11. import monai
  12. import argparse
  13. import numpy as np
  14. import torch
  15. import torch.nn as nn
  16. from torch.utils.data import DataLoader as _TorchDataLoader
  17. from torch.utils.data import Dataset
  18. from torch.utils.tensorboard import SummaryWriter
  19. from monai.data import decollate_batch, CSVSaver, DataLoader
  20. from monai.data.utils import list_data_collate, set_rnd, worker_init_fn
  21. from monai.metrics import ROCAUCMetric
  22. from monai.transforms import ScaleIntensityRanged, NormalizeIntensityd, Activations, AsDiscrete, Compose, LoadImaged, RandRotated, Resized, ScaleIntensityd, RandFlipd
  23. from monai.networks.nets import resnet10, resnet18, resnet34, resnet50
  24. from sklearn.model_selection import StratifiedKFold
  25. from statistics import mean, mode
  26. import matplotlib.pyplot as plt
  27. from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay, classification_report, accuracy_score, roc_auc_score
  28. # %%
  29. def calculate_scores(y_true, y_pred_act_all):
  30. print("="*20)
  31. print("Test Set Results:\n")
  32. y_pred_act_all = np.array(y_pred_act_all)
  33. sum5 = np.sum(np.array(y_pred_act_all), axis=0)
  34. probs = sum5/5
  35. ypred_soft_votes = np.argmax(probs, axis=1)
  36. acc_score = accuracy_score(y_true, ypred_soft_votes)
  37. probs_auc = y_pred_act_all[:,:,1]
  38. sum5 = np.sum(np.array(probs_auc), axis=0)
  39. probs_auc = sum5/5
  40. auc_score = roc_auc_score(y_true, probs_auc)
  41. cm = confusion_matrix(y_true, ypred_soft_votes)
  42. # Calculate confusion matrix
  43. tn, fp, fn, tp = cm.ravel()
  44. # Calculate Precision
  45. precision = precision_score(y_true, ypred_soft_votes)
  46. # Calculate Recall (Sensitivity)
  47. recall = recall_score(y_true, ypred_soft_votes)
  48. # Calculate Specificity
  49. specificity = tn / (tn + fp)
  50. # Calculate F1 Score
  51. f1 = f1_score(y_true, ypred_soft_votes)
  52. return acc_score, auc_score, precision, recall, specificity, f1, cm, ypred_soft_votes
  53. # %%
  54. import sys
  55. from mri_classification.scripts.sbatch.bi_5fold_solo_eval import main
  56. from scipy.stats import mode
  57. from sklearn.metrics import roc_auc_score
  58. from sklearn.metrics import confusion_matrix, precision_score, recall_score, f1_score
  59. sys.argv = ['']
  60. parser = argparse.ArgumentParser()
  61. parser.add_argument('--modality', type=str, default='MR')
  62. parser.add_argument('--opt', type=str, default='Adam')
  63. parser.add_argument('--lr', type=float, default=1e-3)
  64. parser.add_argument('--wdecay', type=float, default=1e-4)
  65. config = parser.parse_args()
  66. y_true, y_pred_act_all = main(config)
  67. acc_score, auc_score, precision, recall, specificity, f1, cm, ypred_soft_votes = calculate_scores(y_true, y_pred_act_all)
  68. # Print the results
  69. print(f"Accuracy: {round(acc_score, 4)}")
  70. print(f"AUC Score: {round(auc_score, 4)}")
  71. print(f"Precision: {round(precision,4)}")
  72. print(f"Recall (Sensitivity): {round(recall,4)}")
  73. print(f"Specificity: {round(specificity, 4)}")
  74. print(f"F1 Score: {round(f1, 4)}")
  75. target_names = ['class 0+1', 'class 2']
  76. print(classification_report(y_true, ypred_soft_votes, target_names=target_names))
  77. disp = ConfusionMatrixDisplay(confusion_matrix=cm)
  78. disp.plot(cmap = 'Blues')
  79. plt.show()
  80. # %%
  81. import sys
  82. from mri_classification.scripts.sbatch.bi_5fold_solo_eval import main
  83. from scipy.stats import mode
  84. from sklearn.metrics import roc_auc_score
  85. from sklearn.metrics import confusion_matrix, precision_score, recall_score, f1_score
  86. sys.argv = ['']
  87. parser = argparse.ArgumentParser()
  88. parser.add_argument('--modality', type=str, default='CT')
  89. parser.add_argument('--opt', type=str, default='Adam')
  90. parser.add_argument('--lr', type=float, default=1e-5)
  91. parser.add_argument('--wdecay', type=float, default=1e-6)
  92. config = parser.parse_args()
  93. y_true, y_pred_act_all = main(config)
  94. acc_score, auc_score, precision, recall, specificity, f1, cm, ypred_soft_votes = calculate_scores(y_true, y_pred_act_all)
  95. # Print the results
  96. print(f"Accuracy: {round(acc_score, 4)}")
  97. print(f"AUC Score: {round(auc_score, 4)}")
  98. print(f"Precision: {round(precision,4)}")
  99. print(f"Recall (Sensitivity): {round(recall,4)}")
  100. print(f"Specificity: {round(specificity, 4)}")
  101. print(f"F1 Score: {round(f1, 4)}")
  102. target_names = ['class 0+1', 'class 2']
  103. print(classification_report(y_true, ypred_soft_votes, target_names=target_names))
  104. disp = ConfusionMatrixDisplay(confusion_matrix=cm)
  105. disp.plot(cmap = 'Blues')
  106. plt.show()
  107. # %%
  108. import sys
  109. from mri_classification.scripts.sbatch.bi_5fold_ensemble_eval import main
  110. from scipy.stats import mode
  111. from sklearn.metrics import roc_auc_score
  112. from sklearn.metrics import confusion_matrix, precision_score, recall_score, f1_score
  113. sys.argv = ['']
  114. parser = argparse.ArgumentParser()
  115. parser.add_argument('--ct_type', type=str, default='Real')
  116. parser.add_argument('--mr_type', type=str, default='Real')
  117. parser.add_argument('--opt', type=str, default='Adam')
  118. parser.add_argument('--epochs', type=int, default=60)
  119. parser.add_argument('--lr', type=float, default=1e-5)
  120. parser.add_argument('--wdecay', type=float, default=1e-3)
  121. config = parser.parse_args()
  122. y_true, y_pred_act_all = main(config)
  123. acc_score, auc_score, precision, recall, specificity, f1, cm, ypred_soft_votes = calculate_scores(y_true, y_pred_act_all)
  124. # Print the results
  125. print(f"Accuracy: {round(acc_score, 4)}")
  126. print(f"AUC Score: {round(auc_score, 4)}")
  127. print(f"Precision: {round(precision,4)}")
  128. print(f"Recall (Sensitivity): {round(recall,4)}")
  129. print(f"Specificity: {round(specificity, 4)}")
  130. print(f"F1 Score: {round(f1, 4)}")
  131. target_names = ['class 0+1', 'class 2']
  132. print(classification_report(y_true, ypred_soft_votes, target_names=target_names))
  133. disp = ConfusionMatrixDisplay(confusion_matrix=cm)
  134. disp.plot(cmap = 'Blues')
  135. plt.show()
  136. # %%
  137. import sys
  138. from mri_classification.scripts.sbatch.bi_5fold_ensemble_eval import main
  139. from scipy.stats import mode
  140. from sklearn.metrics import roc_auc_score
  141. from sklearn.metrics import confusion_matrix, precision_score, recall_score, f1_score
  142. sys.argv = ['']
  143. parser = argparse.ArgumentParser()
  144. parser.add_argument('--ct_type', type=str, default='Fake')
  145. parser.add_argument('--mr_type', type=str, default='Real')
  146. parser.add_argument('--opt', type=str, default='Adam')
  147. parser.add_argument('--epochs', type=int, default=60)
  148. parser.add_argument('--lr', type=float, default=1e-5)
  149. parser.add_argument('--wdecay', type=float, default=1e-2)
  150. config = parser.parse_args()
  151. y_true, y_pred_act_all = main(config)
  152. acc_score, auc_score, precision, recall, specificity, f1, cm, ypred_soft_votes = calculate_scores(y_true, y_pred_act_all)
  153. # Print the results
  154. print(f"Accuracy: {round(acc_score, 4)}")
  155. print(f"AUC Score: {round(auc_score, 4)}")
  156. print(f"Precision: {round(precision,4)}")
  157. print(f"Recall (Sensitivity): {round(recall,4)}")
  158. print(f"Specificity: {round(specificity, 4)}")
  159. print(f"F1 Score: {round(f1, 4)}")
  160. target_names = ['class 0+1', 'class 2']
  161. print(classification_report(y_true, ypred_soft_votes, target_names=target_names))
  162. disp = ConfusionMatrixDisplay(confusion_matrix=cm)
  163. disp.plot(cmap = 'Blues')
  164. plt.show()
  165. # %%
  166. import sys
  167. from mri_classification.scripts.sbatch.bi_5fold_ensemble_eval import main
  168. from scipy.stats import mode
  169. from sklearn.metrics import roc_auc_score
  170. from sklearn.metrics import confusion_matrix, precision_score, recall_score, f1_score
  171. sys.argv = ['']
  172. parser = argparse.ArgumentParser()
  173. parser.add_argument('--ct_type', type=str, default='Fake')
  174. parser.add_argument('--mr_type', type=str, default='Real')
  175. parser.add_argument('--opt', type=str, default='Adam')
  176. parser.add_argument('--epochs', type=int, default=60)
  177. parser.add_argument('--lr', type=float, default=1e-5)
  178. parser.add_argument('--wdecay', type=float, default=1e-4)
  179. config = parser.parse_args()
  180. y_true, y_pred_act_all = main(config)
  181. acc_score, auc_score, precision, recall, specificity, f1, cm, ypred_soft_votes = calculate_scores(y_true, y_pred_act_all)
  182. # Print the results
  183. print(f"Accuracy: {round(acc_score, 4)}")
  184. print(f"AUC Score: {round(auc_score, 4)}")
  185. print(f"Precision: {round(precision,4)}")
  186. print(f"Recall (Sensitivity): {round(recall,4)}")
  187. print(f"Specificity: {round(specificity, 4)}")
  188. print(f"F1 Score: {round(f1, 4)}")
  189. target_names = ['class 0+1', 'class 2']
  190. print(classification_report(y_true, ypred_soft_votes, target_names=target_names))
  191. disp = ConfusionMatrixDisplay(confusion_matrix=cm)
  192. disp.plot(cmap = 'Blues')
  193. plt.show()
  194. # %% [markdown]
  195. # ### DeLong Test
  196. # %%
  197. import pandas as pd
  198. import numpy as np
  199. import scipy.stats
  200. # AUC comparison adapted from
  201. # https://github.com/Netflix/vmaf/
  202. def compute_midrank(x):
  203. """Computes midranks.
  204. Args:
  205. x - a 1D numpy array
  206. Returns:
  207. array of midranks
  208. """
  209. J = np.argsort(x)
  210. Z = x[J]
  211. N = len(x)
  212. T = np.zeros(N, dtype=float)
  213. i = 0
  214. while i < N:
  215. j = i
  216. while j < N and Z[j] == Z[i]:
  217. j += 1
  218. T[i:j] = 0.5*(i + j - 1)
  219. i = j
  220. T2 = np.empty(N, dtype=float)
  221. # Note(kazeevn) +1 is due to Python using 0-based indexing
  222. # instead of 1-based in the AUC formula in the paper
  223. T2[J] = T + 1
  224. return T2
  225. def fastDeLong(predictions_sorted_transposed, label_1_count):
  226. """
  227. The fast version of DeLong's method for computing the covariance of
  228. unadjusted AUC.
  229. Args:
  230. predictions_sorted_transposed: a 2D numpy.array[n_classifiers, n_examples]
  231. sorted such as the examples with label "1" are first
  232. Returns:
  233. (AUC value, DeLong covariance)
  234. Reference:
  235. @article{sun2014fast,
  236. title={Fast Implementation of DeLong's Algorithm for
  237. Comparing the Areas Under Correlated Receiver Operating Characteristic Curves},
  238. author={Xu Sun and Weichao Xu},
  239. journal={IEEE Signal Processing Letters},
  240. volume={21},
  241. number={11},
  242. pages={1389--1393},
  243. year={2014},
  244. publisher={IEEE}
  245. }
  246. """
  247. # Short variables are named as they are in the paper
  248. m = label_1_count
  249. n = predictions_sorted_transposed.shape[1] - m
  250. positive_examples = predictions_sorted_transposed[:, :m]
  251. negative_examples = predictions_sorted_transposed[:, m:]
  252. k = predictions_sorted_transposed.shape[0]
  253. tx = np.empty([k, m], dtype=float)
  254. ty = np.empty([k, n], dtype=float)
  255. tz = np.empty([k, m + n], dtype=np.float)
  256. for r in range(k):
  257. tx[r, :] = compute_midrank(positive_examples[r, :])
  258. ty[r, :] = compute_midrank(negative_examples[r, :])
  259. tz[r, :] = compute_midrank(predictions_sorted_transposed[r, :])
  260. aucs = tz[:, :m].sum(axis=1) / m / n - float(m + 1.0) / 2.0 / n
  261. v01 = (tz[:, :m] - tx[:, :]) / n
  262. v10 = 1.0 - (tz[:, m:] - ty[:, :]) / m
  263. sx = np.cov(v01)
  264. sy = np.cov(v10)
  265. delongcov = sx / m + sy / n
  266. return aucs, delongcov
  267. def calc_pvalue(aucs, sigma):
  268. """Computes log(10) of p-values.
  269. Args:
  270. aucs: 1D array of AUCs
  271. sigma: AUC DeLong covariances
  272. Returns:
  273. log10(pvalue)
  274. """
  275. l = np.array([[1, -1]])
  276. z = np.abs(np.diff(aucs)) / np.sqrt(np.dot(np.dot(l, sigma), l.T))
  277. return np.log10(2) + scipy.stats.norm.logsf(z, loc=0, scale=1) / np.log(10)
  278. def compute_ground_truth_statistics(ground_truth):
  279. assert np.array_equal(np.unique(ground_truth), [0, 1])
  280. order = (-ground_truth).argsort()
  281. label_1_count = int(ground_truth.sum())
  282. return order, label_1_count
  283. def delong_roc_variance(ground_truth, predictions):
  284. """
  285. Computes ROC AUC variance for a single set of predictions
  286. Args:
  287. ground_truth: np.array of 0 and 1
  288. predictions: np.array of floats of the probability of being class 1
  289. """
  290. order, label_1_count = compute_ground_truth_statistics(ground_truth)
  291. predictions_sorted_transposed = predictions[np.newaxis, order]
  292. aucs, delongcov = fastDeLong(predictions_sorted_transposed, label_1_count)
  293. assert len(aucs) == 1, "There is a bug in the code, please forward this to the developers"
  294. return aucs[0], delongcov
  295. def delong_roc_test(ground_truth, predictions_one, predictions_two):
  296. """
  297. Computes log(p-value) for hypothesis that two ROC AUCs are different
  298. Args:
  299. ground_truth: np.array of 0 and 1
  300. predictions_one: predictions of the first model,
  301. np.array of floats of the probability of being class 1
  302. predictions_two: predictions of the second model,
  303. np.array of floats of the probability of being class 1
  304. """
  305. order, label_1_count = compute_ground_truth_statistics(ground_truth)
  306. predictions_sorted_transposed = np.vstack((predictions_one, predictions_two))[:, order]
  307. aucs, delongcov = fastDeLong(predictions_sorted_transposed, label_1_count)
  308. return calc_pvalue(aucs, delongcov)
  309. # %%
  310. import sys
  311. from mri_classification.scripts.sbatch.bi_5fold_ensemble_eval import main
  312. from scipy.stats import mode
  313. from sklearn.metrics import roc_auc_score
  314. from sklearn.metrics import confusion_matrix, precision_score, recall_score, f1_score
  315. # Fake CT + Real MR
  316. sys.argv = ['']
  317. parser = argparse.ArgumentParser()
  318. parser.add_argument('--ct_type', type=str, default='Fake')
  319. parser.add_argument('--mr_type', type=str, default='Real')
  320. parser.add_argument('--opt', type=str, default='Adam')
  321. parser.add_argument('--epochs', type=int, default=60)
  322. parser.add_argument('--lr', type=float, default=1e-5)
  323. parser.add_argument('--wdecay', type=float, default=1e-4)
  324. config = parser.parse_args()
  325. y_true, y_pred_act_all = main(config)
  326. acc_score, auc_score, precision, recall, specificity, f1, cm, fakect_ypred = calculate_scores(y_true, y_pred_act_all)
  327. # Real MR only
  328. from mri_classification.scripts.sbatch.bi_5fold_solo_eval import main
  329. sys.argv = ['']
  330. parser = argparse.ArgumentParser()
  331. parser.add_argument('--modality', type=str, default='MR')
  332. parser.add_argument('--opt', type=str, default='Adam')
  333. parser.add_argument('--lr', type=float, default=1e-3)
  334. parser.add_argument('--wdecay', type=float, default=1e-4)
  335. config = parser.parse_args()
  336. y_true, y_pred_act_all = main(config)
  337. acc_score, auc_score, precision, recall, specificity, f1, cm, mr_ypred = calculate_scores(y_true, y_pred_act_all)
  338. # %%
  339. import pandas as pd
  340. import numpy as np
  341. import scipy.stats
  342. # AUC comparison adapted from
  343. # https://github.com/Netflix/vmaf/
  344. def compute_midrank(x):
  345. """Computes midranks.
  346. Args:
  347. x - a 1D numpy array
  348. Returns:
  349. array of midranks
  350. """
  351. J = np.argsort(x)
  352. Z = x[J]
  353. N = len(x)
  354. T = np.zeros(N, dtype=float)
  355. i = 0
  356. while i < N:
  357. j = i
  358. while j < N and Z[j] == Z[i]:
  359. j += 1
  360. T[i:j] = 0.5*(i + j - 1)
  361. i = j
  362. T2 = np.empty(N, dtype=float)
  363. # Note(kazeevn) +1 is due to Python using 0-based indexing
  364. # instead of 1-based in the AUC formula in the paper
  365. T2[J] = T + 1
  366. return T2
  367. def fastDeLong(predictions_sorted_transposed, label_1_count):
  368. """
  369. The fast version of DeLong's method for computing the covariance of
  370. unadjusted AUC.
  371. Args:
  372. predictions_sorted_transposed: a 2D numpy.array[n_classifiers, n_examples]
  373. sorted such as the examples with label "1" are first
  374. Returns:
  375. (AUC value, DeLong covariance)
  376. Reference:
  377. @article{sun2014fast,
  378. title={Fast Implementation of DeLong's Algorithm for
  379. Comparing the Areas Under Correlated Receiver Operating Characteristic Curves},
  380. author={Xu Sun and Weichao Xu},
  381. journal={IEEE Signal Processing Letters},
  382. volume={21},
  383. number={11},
  384. pages={1389--1393},
  385. year={2014},
  386. publisher={IEEE}
  387. }
  388. """
  389. # Short variables are named as they are in the paper
  390. m = label_1_count
  391. n = predictions_sorted_transposed.shape[1] - m
  392. positive_examples = predictions_sorted_transposed[:, :m]
  393. negative_examples = predictions_sorted_transposed[:, m:]
  394. k = predictions_sorted_transposed.shape[0]
  395. tx = np.empty([k, m], dtype=float)
  396. ty = np.empty([k, n], dtype=float)
  397. tz = np.empty([k, m + n], dtype=float)
  398. for r in range(k):
  399. tx[r, :] = compute_midrank(positive_examples[r, :])
  400. ty[r, :] = compute_midrank(negative_examples[r, :])
  401. tz[r, :] = compute_midrank(predictions_sorted_transposed[r, :])
  402. aucs = tz[:, :m].sum(axis=1) / m / n - float(m + 1.0) / 2.0 / n
  403. v01 = (tz[:, :m] - tx[:, :]) / n
  404. v10 = 1.0 - (tz[:, m:] - ty[:, :]) / m
  405. sx = np.cov(v01)
  406. sy = np.cov(v10)
  407. delongcov = sx / m + sy / n
  408. return aucs, delongcov
  409. def calc_pvalue(aucs, sigma):
  410. """Computes log(10) of p-values.
  411. Args:
  412. aucs: 1D array of AUCs
  413. sigma: AUC DeLong covariances
  414. Returns:
  415. log10(pvalue)
  416. """
  417. l = np.array([[1, -1]])
  418. z = np.abs(np.diff(aucs)) / np.sqrt(np.dot(np.dot(l, sigma), l.T))
  419. return np.log10(2) + scipy.stats.norm.logsf(z, loc=0, scale=1) / np.log(10)
  420. def compute_ground_truth_statistics(ground_truth):
  421. assert np.array_equal(np.unique(ground_truth), [0, 1])
  422. order = (-ground_truth).argsort()
  423. label_1_count = int(ground_truth.sum())
  424. return order, label_1_count
  425. def delong_roc_variance(ground_truth, predictions):
  426. """
  427. Computes ROC AUC variance for a single set of predictions
  428. Args:
  429. ground_truth: np.array of 0 and 1
  430. predictions: np.array of floats of the probability of being class 1
  431. """
  432. order, label_1_count = compute_ground_truth_statistics(ground_truth)
  433. predictions_sorted_transposed = predictions[np.newaxis, order]
  434. aucs, delongcov = fastDeLong(predictions_sorted_transposed, label_1_count)
  435. assert len(aucs) == 1, "There is a bug in the code, please forward this to the developers"
  436. return aucs[0], delongcov
  437. def delong_roc_test(ground_truth, predictions_one, predictions_two):
  438. """
  439. Computes log(p-value) for hypothesis that two ROC AUCs are different
  440. Args:
  441. ground_truth: np.array of 0 and 1
  442. predictions_one: predictions of the first model,
  443. np.array of floats of the probability of being class 1
  444. predictions_two: predictions of the second model,
  445. np.array of floats of the probability of being class 1
  446. """
  447. order, label_1_count = compute_ground_truth_statistics(ground_truth)
  448. predictions_sorted_transposed = np.vstack((predictions_one, predictions_two))[:, order]
  449. aucs, delongcov = fastDeLong(predictions_sorted_transposed, label_1_count)
  450. return calc_pvalue(aucs, delongcov)
  451. # %%
  452. pvalue = delong_roc_test(y_true, fakect_ypred, mr_ypred)
  453. pvalue
  454. # %%
  455. p_value_test = np.exp(np.log(10)*pvalue)
  456. p_value_test
  457. # %%
  458. # %%
  459. # %%
  460. import pandas as pd
  461. import numpy as np
  462. import scipy.stats
  463. from scipy import stats
  464. # AUC comparison adapted from
  465. # https://github.com/Netflix/vmaf/
  466. def compute_midrank(x):
  467. """Computes midranks.
  468. Args:
  469. x - a 1D numpy array
  470. Returns:
  471. array of midranks
  472. """
  473. J = np.argsort(x)
  474. Z = x[J]
  475. N = len(x)
  476. T = np.zeros(N, dtype=float)
  477. i = 0
  478. while i < N:
  479. j = i
  480. while j < N and Z[j] == Z[i]:
  481. j += 1
  482. T[i:j] = 0.5*(i + j - 1)
  483. i = j
  484. T2 = np.empty(N, dtype=float)
  485. # Note(kazeevn) +1 is due to Python using 0-based indexing
  486. # instead of 1-based in the AUC formula in the paper
  487. T2[J] = T + 1
  488. return T2
  489. def compute_midrank_weight(x, sample_weight):
  490. """Computes midranks.
  491. Args:
  492. x - a 1D numpy array
  493. Returns:
  494. array of midranks
  495. """
  496. J = np.argsort(x)
  497. Z = x[J]
  498. cumulative_weight = np.cumsum(sample_weight[J])
  499. N = len(x)
  500. T = np.zeros(N, dtype=float)
  501. i = 0
  502. while i < N:
  503. j = i
  504. while j < N and Z[j] == Z[i]:
  505. j += 1
  506. T[i:j] = cumulative_weight[i:j].mean()
  507. i = j
  508. T2 = np.empty(N, dtype=float)
  509. T2[J] = T
  510. return T2
  511. def fastDeLong(predictions_sorted_transposed, label_1_count, sample_weight=None):
  512. if sample_weight is None:
  513. return fastDeLong_no_weights(predictions_sorted_transposed, label_1_count)
  514. else:
  515. return fastDeLong_weights(predictions_sorted_transposed, label_1_count, sample_weight)
  516. def fastDeLong_weights(predictions_sorted_transposed, label_1_count, sample_weight):
  517. """
  518. The fast version of DeLong's method for computing the covariance of
  519. unadjusted AUC.
  520. Args:
  521. predictions_sorted_transposed: a 2D numpy.array[n_classifiers, n_examples]
  522. sorted such as the examples with label "1" are first
  523. Returns:
  524. (AUC value, DeLong covariance)
  525. Reference:
  526. @article{sun2014fast,
  527. title={Fast Implementation of DeLong's Algorithm for
  528. Comparing the Areas Under Correlated Receiver Oerating Characteristic Curves},
  529. author={Xu Sun and Weichao Xu},
  530. journal={IEEE Signal Processing Letters},
  531. volume={21},
  532. number={11},
  533. pages={1389--1393},
  534. year={2014},
  535. publisher={IEEE}
  536. }
  537. """
  538. # Short variables are named as they are in the paper
  539. m = label_1_count
  540. n = predictions_sorted_transposed.shape[1] - m
  541. positive_examples = predictions_sorted_transposed[:, :m]
  542. negative_examples = predictions_sorted_transposed[:, m:]
  543. k = predictions_sorted_transposed.shape[0]
  544. tx = np.empty([k, m], dtype=float)
  545. ty = np.empty([k, n], dtype=float)
  546. tz = np.empty([k, m + n], dtype=float)
  547. for r in range(k):
  548. tx[r, :] = compute_midrank_weight(positive_examples[r, :], sample_weight[:m])
  549. ty[r, :] = compute_midrank_weight(negative_examples[r, :], sample_weight[m:])
  550. tz[r, :] = compute_midrank_weight(predictions_sorted_transposed[r, :], sample_weight)
  551. total_positive_weights = sample_weight[:m].sum()
  552. total_negative_weights = sample_weight[m:].sum()
  553. pair_weights = np.dot(sample_weight[:m, np.newaxis], sample_weight[np.newaxis, m:])
  554. total_pair_weights = pair_weights.sum()
  555. aucs = (sample_weight[:m]*(tz[:, :m] - tx)).sum(axis=1) / total_pair_weights
  556. v01 = (tz[:, :m] - tx[:, :]) / total_negative_weights
  557. v10 = 1. - (tz[:, m:] - ty[:, :]) / total_positive_weights
  558. sx = np.cov(v01)
  559. sy = np.cov(v10)
  560. delongcov = sx / m + sy / n
  561. return aucs, delongcov
  562. def fastDeLong_no_weights(predictions_sorted_transposed, label_1_count):
  563. """
  564. The fast version of DeLong's method for computing the covariance of
  565. unadjusted AUC.
  566. Args:
  567. predictions_sorted_transposed: a 2D numpy.array[n_classifiers, n_examples]
  568. sorted such as the examples with label "1" are first
  569. Returns:
  570. (AUC value, DeLong covariance)
  571. Reference:
  572. @article{sun2014fast,
  573. title={Fast Implementation of DeLong's Algorithm for
  574. Comparing the Areas Under Correlated Receiver Oerating
  575. Characteristic Curves},
  576. author={Xu Sun and Weichao Xu},
  577. journal={IEEE Signal Processing Letters},
  578. volume={21},
  579. number={11},
  580. pages={1389--1393},
  581. year={2014},
  582. publisher={IEEE}
  583. }
  584. """
  585. # Short variables are named as they are in the paper
  586. m = label_1_count
  587. n = predictions_sorted_transposed.shape[1] - m
  588. positive_examples = predictions_sorted_transposed[:, :m]
  589. negative_examples = predictions_sorted_transposed[:, m:]
  590. k = predictions_sorted_transposed.shape[0]
  591. tx = np.empty([k, m], dtype=float)
  592. ty = np.empty([k, n], dtype=float)
  593. tz = np.empty([k, m + n], dtype=float)
  594. for r in range(k):
  595. tx[r, :] = compute_midrank(positive_examples[r, :])
  596. ty[r, :] = compute_midrank(negative_examples[r, :])
  597. tz[r, :] = compute_midrank(predictions_sorted_transposed[r, :])
  598. aucs = tz[:, :m].sum(axis=1) / m / n - float(m + 1.0) / 2.0 / n
  599. v01 = (tz[:, :m] - tx[:, :]) / n
  600. v10 = 1.0 - (tz[:, m:] - ty[:, :]) / m
  601. sx = np.cov(v01)
  602. sy = np.cov(v10)
  603. delongcov = sx / m + sy / n
  604. return aucs, delongcov
  605. def calc_pvalue(aucs, sigma):
  606. """Computes log(10) of p-values.
  607. Args:
  608. aucs: 1D array of AUCs
  609. sigma: AUC DeLong covariances
  610. Returns:
  611. log10(pvalue)
  612. """
  613. l = np.array([[1, -1]])
  614. z = np.abs(np.diff(aucs)) / (np.sqrt(np.dot(np.dot(l, sigma), l.T)) + 1e-8)
  615. pvalue = 2 * (1 - scipy.stats.norm.cdf(np.abs(z)))
  616. # print(10**(np.log10(2) + scipy.stats.norm.logsf(z, loc=0, scale=1) / np.log(10)))
  617. return pvalue
  618. def compute_ground_truth_statistics(ground_truth, sample_weight=None):
  619. assert np.array_equal(np.unique(ground_truth), [0, 1])
  620. order = (-ground_truth).argsort()
  621. label_1_count = int(ground_truth.sum())
  622. if sample_weight is None:
  623. ordered_sample_weight = None
  624. else:
  625. ordered_sample_weight = sample_weight[order]
  626. return order, label_1_count, ordered_sample_weight
  627. def delong_roc_variance(ground_truth, predictions):
  628. """
  629. Computes ROC AUC variance for a single set of predictions
  630. Args:
  631. ground_truth: np.array of 0 and 1
  632. predictions: np.array of floats of the probability of being class 1
  633. """
  634. sample_weight = None
  635. order, label_1_count, ordered_sample_weight = compute_ground_truth_statistics(
  636. ground_truth, sample_weight)
  637. predictions_sorted_transposed = predictions[np.newaxis, order]
  638. aucs, delongcov = fastDeLong(predictions_sorted_transposed, label_1_count)
  639. assert len(aucs) == 1, "There is a bug in the code, please forward this to the developers"
  640. return aucs[0], delongcov
  641. def delong_roc_test(ground_truth, predictions_one, predictions_two):
  642. """
  643. Computes log(p-value) for hypothesis that two ROC AUCs are different
  644. Args:
  645. ground_truth: np.array of 0 and 1
  646. predictions_one: predictions of the first model,
  647. np.array of floats of the probability of being class 1
  648. predictions_two: predictions of the second model,
  649. np.array of floats of the probability of being class 1
  650. """
  651. sample_weight = None
  652. order, label_1_count,ordered_sample_weight = compute_ground_truth_statistics(ground_truth)
  653. predictions_sorted_transposed = np.vstack((predictions_one, predictions_two))[:, order]
  654. aucs, delongcov = fastDeLong(predictions_sorted_transposed, label_1_count,sample_weight)
  655. return calc_pvalue(aucs, delongcov)
  656. # def delong_roc_ci(y_true,y_pred):
  657. # aucs, auc_cov = delong_roc_variance(y_true, y_pred)
  658. # auc_std = np.sqrt(auc_cov)
  659. # lower_upper_q = np.abs(np.array([0, 1]) - (1 - alpha) / 2)
  660. # ci = stats.norm.ppf(
  661. # lower_upper_q,
  662. # loc=aucs,
  663. # scale=auc_std)
  664. # ci[ci > 1] = 1
  665. # return aucs,ci
  666. def delong_roc_ci(y_true,y_pred):
  667. aucs, auc_cov = delong_roc_variance(y_true, y_pred)
  668. auc_std = np.sqrt(auc_cov)
  669. lower_upper_q = np.abs(np.array([0, 1]) - (1 - alpha) / 2)
  670. ci = stats.norm.ppf(
  671. lower_upper_q,
  672. loc=aucs,
  673. scale=auc_std)
  674. ci[ci > 1] = 1
  675. return aucs,ci
  676. def get_95CI(y_true,y_pred_1):
  677. """
  678. Return the 95% CI and AUC of prediction
  679. Args:
  680. labels: array (n,) the ground truth
  681. scores1: array (n,) the predicted probability
  682. """
  683. alpha = .95
  684. auc_1, auc_cov_1 = delong_roc_variance(y_true, y_pred_1)
  685. auc_std = np.sqrt(auc_cov_1)
  686. lower_upper_q = np.abs(np.array([0, 1]) - (1 - alpha) / 2)
  687. # 95% CI
  688. ci = stats.norm.ppf(
  689. lower_upper_q,
  690. loc=auc_1,
  691. scale=auc_std)
  692. ci[ci > 1] = 1
  693. return ci,auc_1
  694. # threshold
  695. from sklearn import metrics
  696. def get_optimal_threshold(labels,y_pred_1):
  697. '''
  698. get the threshold according to youden index
  699. Args:
  700. labels:<numpy.ndarray> (n,) groundtruth
  701. y_pred_1:<numpy.ndarray> (n,) predicted probabilities
  702. '''
  703. fpr, tpr, thresholds = metrics.roc_curve(labels,y_pred_1)
  704. optimal_index = np.argmax(+tpr-fpr)
  705. optimal_thresholds = thresholds[optimal_index]
  706. # print(thresholds)
  707. # print(optimal_index)
  708. # print(fpr)
  709. return optimal_thresholds
  710. def get_metric(y_true, y_prob,threshold,verbose = True):
  711. '''
  712. Return the commanly used metric value according to the given threshold
  713. Args:
  714. y_true:<numpy.ndarray> (n,) groundtruth
  715. y_prob:<numpy.ndarray> (n,) predicted probabilities
  716. threshold: <float>
  717. Return:
  718. scores: <dict> commanly used metric
  719. '''
  720. scores = {}
  721. y_pred = (y_prob>=(threshold-1E-4)).astype(int)
  722. # print report
  723. target_names = ['class 0', 'class 1']
  724. text = metrics.classification_report(y_true, y_pred, target_names=target_names)
  725. conf_mat=pd.crosstab(y_true, y_pred,rownames=['label'],colnames=['pre'])
  726. if verbose:
  727. print(conf_mat)
  728. print(text)
  729. # accuracy
  730. scores['accuracy'] = metrics.accuracy_score(y_true, y_pred)
  731. # precision
  732. try:
  733. scores['PPV'] = metrics.precision_score(y_true, y_pred)
  734. except:
  735. scores['PPV'] = None
  736. #NPV
  737. try:
  738. scores['NPV'] = conf_mat[0][0]/(conf_mat[0][0]+conf_mat[0][1])
  739. except:
  740. scores['NPV'] = None
  741. # recall
  742. scores['recall (sensitivity)'] = metrics.recall_score(y_true, y_pred)
  743. scores['recall_neg (specificity)'] = metrics.recall_score(y_true==0, y_pred==0)
  744. # F1-score
  745. scores['f1_score'] = metrics.f1_score(y_true, y_pred)
  746. # ROC/AUC
  747. scores['AUC95%CI'], scores['AUC'] = get_95CI(y_true,y_pred)
  748. return scores
  749. # %%
  750. pvalue = delong_roc_test(y_true, fakect_ypred, mr_ypred)
  751. print('p_value:', pvalue)
  752. # %%
  753. # %%
  754. # %%
  755. # %% [markdown]
  756. # # Three-class classification
  757. # %%
  758. import logging
  759. import os
  760. import sys
  761. from pathlib import Path
  762. import pandas as pd
  763. import monai
  764. import numpy as np
  765. import torch
  766. import torch.nn as nn
  767. from torch.utils.data import DataLoader as _TorchDataLoader
  768. from torch.utils.data import Dataset
  769. from torch.utils.tensorboard import SummaryWriter
  770. from monai.data import decollate_batch, CSVSaver, DataLoader
  771. from monai.data.utils import list_data_collate, set_rnd, worker_init_fn
  772. from monai.metrics import ROCAUCMetric
  773. from monai.transforms import ScaleIntensityRanged, NormalizeIntensityd, Activations, AsDiscrete, Compose, LoadImaged, RandRotated, Resized, ScaleIntensityd, RandFlipd
  774. from monai.networks.nets import resnet10, resnet18, resnet34, resnet50
  775. from sklearn.model_selection import StratifiedKFold
  776. from statistics import mean, mode
  777. import matplotlib.pyplot as plt
  778. from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay, classification_report, accuracy_score, roc_auc_score
  779. class CustomDataset(monai.data.Dataset):
  780. def __init__(self, d1, d2):
  781. self.d1 = d1
  782. self.d2 = d2
  783. def __getitem__(self, idx):
  784. dict1 = self.d1.__getitem__(idx)
  785. image1, label1 = dict1["img"], dict1["label"]
  786. dict2 = self.d2.__getitem__(idx)
  787. image2, label2 = dict2["img"], dict2["label"]
  788. assert label1==label2
  789. dict_1 = dict()
  790. dict_1["img"] = image1
  791. dict_1["label"] = label1
  792. dict_2 = dict()
  793. dict_2["img"] = image2
  794. dict_2["label"] = label2
  795. return dict_1, dict_2
  796. def __len__(self):
  797. return len(self.d1)
  798. class MyEnsemble(nn.Module):
  799. def __init__(self, modelA, modelB):
  800. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  801. super(MyEnsemble, self).__init__()
  802. self.modelA = modelA
  803. self.modelB = modelB
  804. # Remove last linear layer
  805. self.modelA.fc = nn.Identity()
  806. self.modelB.fc = nn.Identity()
  807. # Create new classifier
  808. self.mlp1 = nn.Linear(1024,256).to(device)
  809. self.mlp2 = nn.Linear(256,32).to(device)
  810. self.classifier = nn.Linear(32,3).to(device)
  811. def forward(self, i1, i2):
  812. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  813. x1 = self.modelA(i1) # clone to make sure x is not changed by inplace methods
  814. x1 = x1.view(x1.size(0), -1).to(device)
  815. x2 = self.modelB(i2)
  816. x2 = x2.view(x2.size(0), -1).to(device)
  817. x = torch.cat((x1, x2), dim=1).to(device)
  818. x = nn.functional.relu(self.mlp1(x)).to(device)
  819. x = nn.functional.relu(self.mlp2(x)).to(device)
  820. x = self.classifier(x).to(device)
  821. return x
  822. def load_data(train_images, val_images, train_labels, val_labels, test_images, test_labels):
  823. train_labels_array = np.array(train_labels, dtype=np.int64)
  824. val_labels_array = np.array(val_labels, dtype=np.int64)
  825. test_labels_array = np.array(test_labels, dtype=np.int64)
  826. #loading the CT images
  827. ct_path = Path('/scratch/ajoshi83/Data/Reg_CT')
  828. #ct_path = Path('/scratch/ajoshi83/generated_ct')
  829. #ct_path = Path('/data/amciilab/ajoshi83/generated_ct2')
  830. ct_train_images_path = [os.path.join(ct_path, f) for f in train_images]
  831. ct_val_images_path = [os.path.join(ct_path, f) for f in val_images]
  832. ct_test_images_path = [os.path.join(ct_path, f) for f in test_images]
  833. ct_train_files = [{"img": img, "label": label} for img, label in zip(ct_train_images_path, train_labels_array)]
  834. ct_val_files = [{"img": img, "label": label} for img, label in zip(ct_val_images_path, val_labels_array)]
  835. ct_test_files = [{"img": img, "label": label} for img, label in zip(ct_test_images_path, test_labels_array)]
  836. #loading the MR images
  837. mr_path = Path('/scratch/ajoshi83/Data/Reg_MR')
  838. #mr_path = Path('/data/amciilab/ajoshi83/generated_mr2')
  839. mr_train_images_path = [os.path.join(mr_path, f) for f in train_images]
  840. mr_val_images_path = [os.path.join(mr_path, f) for f in val_images]
  841. mr_test_images_path = [os.path.join(mr_path, f) for f in test_images]
  842. mr_train_files = [{"img": img, "label": label} for img, label in zip(mr_train_images_path, train_labels_array)]
  843. mr_val_files = [{"img": img, "label": label} for img, label in zip(mr_val_images_path, val_labels_array)]
  844. mr_test_files = [{"img": img, "label": label} for img, label in zip(mr_test_images_path, test_labels_array)]
  845. return ct_train_files, ct_val_files, mr_train_files, mr_val_files, ct_test_files, mr_test_files
  846. #monai.config.print_config()
  847. # old_stdout = sys.stdout
  848. # log_file = open("mri_classification/logs/5fold_eval_fakemr.log","w")
  849. # sys.stdout = log_file
  850. print("this will be written to message.log")
  851. # file_handler = logging.FileHandler(filename='logs/tmp.log')
  852. # #stdout_handler = logging.StreamHandler(stream=sys.stdout)
  853. # handlers = [file_handler]
  854. # logging.basicConfig(
  855. # level=logging.DEBUG,
  856. # format='[%(asctime)s] {%(filename)s:%(lineno)d} %(levelname)s - %(message)s',
  857. # handlers=handlers
  858. # )
  859. # logger = logging.getLogger('LOGGER_NAME')
  860. # edit the path accordingly
  861. gose = pd.read_csv('mri_classification/Pilot_169_train_val_test_split.csv')
  862. X = []
  863. Y = []
  864. test_subjects = []
  865. test_labels = []
  866. torch.cuda.empty_cache()
  867. for i in range(gose.shape[0]):
  868. subj = gose['Main.GUID'][i]
  869. subj_id = str(subj)[4:]
  870. name = str(subj_id) + '.nii'
  871. if gose['Set'][i]=='Train':
  872. X.append(name)
  873. Y.append(gose['Class'][i])
  874. elif gose['Set'][i]=='Val':
  875. X.append(name)
  876. Y.append(gose['Class'][i])
  877. elif gose['Set'][i]=='Test':
  878. test_subjects.append(name)
  879. test_labels.append(gose['Class'][i])
  880. else:
  881. print("Unknown Set: ", gose['Set'][i])
  882. print("Total Subjects for 5-Fold (Train+Val):", len(X))
  883. print("Test Subjects (kept separate):", len(test_subjects))
  884. print("5-Fold Subjects Information:\n")
  885. print("Subjects of Class 0: ", Y.count(0))
  886. print("Subjects of Class 1: ", Y.count(1))
  887. print("Subjects of Class 2: ", Y.count(2))
  888. skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=1)
  889. fold_number = 1
  890. lr =1e-5
  891. decay = 1e-5
  892. X = np.array(X)
  893. Y = np.array(Y)
  894. train_acc, train_auc = [], []
  895. val_acc, val_auc = [], []
  896. test_acc, test_auc = [], []
  897. y_pred_test = []
  898. y_pred_auc = []
  899. for train_index, test_index in skf.split(X, Y):
  900. print("Starting Fold {}..!".format(fold_number))
  901. x_train_fold, x_test_fold = X[train_index], X[test_index]
  902. y_train_fold, y_test_fold = Y[train_index], Y[test_index]
  903. y_train_fold = list(y_train_fold)
  904. y_test_fold = list(y_test_fold)
  905. x_train_fold = list(x_train_fold)
  906. x_test_fold = list(x_test_fold)
  907. # print("\n")
  908. # print("Fold {} statistics:\n".format(fold_number))
  909. # print("Train Subjects: {}".format(len(y_train_fold)))
  910. # print("Subjects of Class 0: ", y_train_fold.count(0))
  911. # print("Subjects of Class 1: ", y_train_fold.count(1))
  912. # print("Subjects of Class 2: {}\n".format(y_train_fold.count(2)))
  913. # print("Val Subjects: {}".format(len(y_test_fold)))
  914. # print("Subjects of Class 0: ", y_test_fold.count(0))
  915. # print("Subjects of Class 1: ", y_test_fold.count(1))
  916. # print("Subjects of Class 2: ", y_test_fold.count(2))
  917. ct_train_files, ct_val_files, mr_train_files, mr_val_files, ct_test_files, mr_test_files = load_data(x_train_fold, x_test_fold, y_train_fold, y_test_fold, test_subjects, test_labels)
  918. # Define transforms for CT and MR respectively
  919. ct_transforms = Compose(
  920. [
  921. LoadImaged(keys=["img"], ensure_channel_first=True),
  922. ScaleIntensityRanged(keys=["img"],a_min=0, a_max=85, b_min=0, b_max=1),
  923. NormalizeIntensityd(keys=["img"]),
  924. #Resized(keys=["img"], spatial_size=(96, 96, 96)),
  925. ]
  926. )
  927. mr_transforms = Compose(
  928. [
  929. LoadImaged(keys=["img"], ensure_channel_first=True),
  930. #ScaleIntensityRanged(keys=["img"],a_min=0, a_max=85, b_min=0, b_max=1),
  931. NormalizeIntensityd(keys=["img"]),
  932. # RandFlipd(keys=["img"], prob=1, spatial_axis=2),
  933. # RandRotated(keys=["img"], prob=1, range_x=[0.4,0.4])
  934. #ScaleIntensityRanged(keys=["img"],a_min=0, a_max=85, b_min=0, b_max=1),
  935. ]
  936. )
  937. post_pred = Compose([Activations(softmax=True)])
  938. post_label = Compose([AsDiscrete(to_onehot=3)])
  939. # create a training data loader
  940. ct_train_ds = monai.data.Dataset(data=ct_train_files, transform=ct_transforms)
  941. mr_train_ds = monai.data.Dataset(data=mr_train_files, transform=mr_transforms)
  942. combined_train_ds = CustomDataset(d1 = mr_train_ds, d2 = ct_train_ds)
  943. train_loader = DataLoader(combined_train_ds, batch_size=1, shuffle=False, pin_memory=torch.cuda.is_available())
  944. check_data1, check_data2 = monai.utils.misc.first(train_loader)
  945. #print(check_data1["img"].shape, check_data1["label"], check_data2["img"].shape, check_data2["label"])
  946. # create a validation data loader
  947. ct_val_ds = monai.data.Dataset(data=ct_val_files, transform=ct_transforms)
  948. mr_val_ds = monai.data.Dataset(data=mr_val_files, transform=mr_transforms)
  949. combined_val_ds = CustomDataset(d1 = mr_val_ds, d2 = ct_val_ds)
  950. val_loader = DataLoader(combined_val_ds, batch_size=1, shuffle=False, pin_memory=torch.cuda)
  951. # create a test data loader
  952. ct_test_ds = monai.data.Dataset(data=ct_test_files, transform=ct_transforms)
  953. mr_test_ds = monai.data.Dataset(data=mr_test_files, transform=mr_transforms)
  954. combined_test_ds = CustomDataset(d1 = mr_test_ds, d2 = ct_test_ds)
  955. test_loader = DataLoader(combined_test_ds, batch_size=1, shuffle=False, pin_memory=torch.cuda)
  956. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  957. #print(device)
  958. model_ct = resnet18(pretrained=False, spatial_dims=3, n_input_channels=1, num_classes=3).to(device)
  959. #model_ct.load_state_dict(torch.load("models/best_val_resnet18.pth"))
  960. model_mr = resnet18(pretrained=False, spatial_dims=3, n_input_channels=1, num_classes=3).to(device)
  961. #model_mr.load_state_dict(torch.load("models/best_mr_val_resnet18.pth"))
  962. # Freeze these models
  963. for param in model_mr.parameters():
  964. param.requires_grad_(True)
  965. for param in model_ct.parameters():
  966. param.requires_grad_(True)
  967. # Create ensemble model
  968. model = MyEnsemble(model_ct, model_mr)
  969. #model.load_state_dict(torch.load("/data/amciilab/ajoshi83/models/CT_Fold{}.pth".format(fold_number)))
  970. #model.load_state_dict(torch.load("/scratch/ajoshi83/models_august/Fake_CT_Fold{}_ensemble_adam_slower.pth".format(int(fold_number))))
  971. model.load_state_dict(torch.load("/scratch/ajoshi83/models_sbatch/Both_Real_Fold{}_ensemble_Adam_{}_{}.pth".format(int(fold_number), lr, decay)))
  972. loss_function = torch.nn.CrossEntropyLoss()
  973. optimizer = torch.optim.Adam(model.parameters(),lr= 1e-5, weight_decay=1e-3)
  974. auc_metric = ROCAUCMetric(average="weighted")
  975. # starting evaluation
  976. val_interval = 1
  977. best_metric = -1
  978. best_metric_epoch = -1
  979. best_val_loss = 2
  980. writer = SummaryWriter()
  981. with torch.no_grad():
  982. num_correct = 0.0
  983. metric_count = 0
  984. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  985. y = torch.tensor([], dtype=torch.long, device=device)
  986. saver = CSVSaver(output_dir="mri_classification/train_output")
  987. for batch_data in train_loader:
  988. #step += 1
  989. #print(batch_data["img"])
  990. ct_batch, mr_batch = batch_data[0], batch_data[1]
  991. #print(ct_batch)
  992. #print(mr_batch)
  993. ct_inputs, ct_labels = ct_batch["img"].to(device), ct_batch["label"].to(device)
  994. mr_inputs, mr_labels = mr_batch["img"].to(device), mr_batch["label"].to(device)
  995. train_outputs = model(ct_inputs, mr_inputs).argmax(dim=1)
  996. y_pred = torch.cat([y_pred, model(ct_inputs, mr_inputs)], dim=0)
  997. y = torch.cat([y, mr_labels], dim=0)
  998. value = torch.eq(train_outputs, mr_labels)
  999. metric_count += len(value)
  1000. num_correct += value.sum().item()
  1001. #saver.save_batch(train_outputs, train_data["img"].meta)
  1002. metric = num_correct / metric_count
  1003. # print("val evaluation metric:", metric)
  1004. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  1005. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred, detach = True)]
  1006. y_onehot = torch.stack(y_onehot, dim=0)
  1007. y_pred_act = torch.stack(y_pred_act, dim=0)
  1008. y_onehot = y_onehot.to(device="cpu")
  1009. y_pred_act = y_pred_act.to(device="cpu")
  1010. auc_metric(y_pred_act, y_onehot)
  1011. auc_result = auc_metric.aggregate()
  1012. auc_metric.reset()
  1013. saver.finalize()
  1014. train_acc.append(round(metric,3))
  1015. train_auc.append(round(auc_result,3))
  1016. print("Train Set: Accuracy: {:.4f}, AUC: {:.4f}".format(metric, auc_result))
  1017. with torch.no_grad():
  1018. num_correct = 0.0
  1019. metric_count = 0
  1020. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  1021. y = torch.tensor([], dtype=torch.long, device=device)
  1022. saver = CSVSaver(output_dir="mri_classification/val_output")
  1023. for val_data in val_loader:
  1024. ct_val_data, mr_val_data = val_data[0], val_data[1]
  1025. ct_val_images, ct_val_labels = ct_val_data["img"].to(device), ct_val_data["label"].to(device)
  1026. mr_val_images, mr_val_labels = mr_val_data["img"].to(device), mr_val_data["label"].to(device)
  1027. y_pred = torch.cat([y_pred, model(ct_val_images, mr_val_images)], dim=0)
  1028. #print(y_pred)
  1029. y = torch.cat([y, mr_val_labels], dim=0)
  1030. acc_value = torch.eq(y_pred.argmax(dim=1), y)
  1031. acc_metric = acc_value.sum().item() / len(acc_value)
  1032. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  1033. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred,detach=True)]
  1034. y_onehot = torch.stack(y_onehot, dim=0)
  1035. y_pred_act = torch.stack(y_pred_act, dim=0)
  1036. y_onehot = y_onehot.to(device="cpu")
  1037. y_pred_act = y_pred_act.to(device="cpu")
  1038. auc_metric(y_pred_act, y_onehot)
  1039. auc_result = auc_metric.aggregate()
  1040. auc_metric.reset()
  1041. saver.finalize()
  1042. val_acc.append(round(acc_metric,3))
  1043. val_auc.append(round(auc_result,3))
  1044. print("Val Set: Accuracy: {:.4f}, AUC: {:.4f}".format(acc_metric, auc_result))
  1045. with torch.no_grad():
  1046. num_correct = 0.0
  1047. metric_count = 0
  1048. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  1049. y = torch.tensor([], dtype=torch.long, device=device)
  1050. saver = CSVSaver(output_dir="mri_classification/test_output")
  1051. y_pred_mid = []
  1052. for test_data in test_loader:
  1053. ct_test_data, mr_test_data = test_data[0], test_data[1]
  1054. ct_test_images, ct_test_labels = ct_test_data["img"].to(device), ct_test_data["label"].to(device)
  1055. mr_test_images, mr_test_labels = mr_test_data["img"].to(device), mr_test_data["label"].to(device)
  1056. test_outputs = model(ct_test_images, mr_test_images).argmax(dim=1)
  1057. y_pred_mid.append(test_outputs.cpu().numpy()[0])
  1058. #softmax_op = torch.nn.functional.softmax(model(ct_test_images, mr_test_images)).argmax(dim=1)
  1059. #y_pred_mid.append(softmax_op.cpu().numpy()[0])
  1060. y_pred = torch.cat([y_pred, model(ct_test_images, mr_test_images)], dim=0)
  1061. y = torch.cat([y, mr_test_labels], dim=0)
  1062. value = torch.eq(test_outputs, mr_test_labels)
  1063. metric_count += len(value)
  1064. num_correct += value.sum().item()
  1065. #saver.save_batch(test_outputs, test_data["img"].meta)
  1066. y_pred_test.append(y_pred_mid)
  1067. metric = num_correct / metric_count
  1068. # print("test evaluation metric:", metric)
  1069. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  1070. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred, detach=True)]
  1071. y_onehot = torch.stack(y_onehot, dim=0)
  1072. y_pred_act = torch.stack(y_pred_act, dim=0)
  1073. y_onehot = y_onehot.to(device="cpu")
  1074. y_pred_act = y_pred_act.to(device="cpu")
  1075. y_pred_auc.append(y_pred_act)
  1076. auc_metric(y_pred_act, y_onehot)
  1077. auc_result = auc_metric.aggregate()
  1078. auc_metric.reset()
  1079. # auc_score =roc_auc_score(y_onehot, y_pred_act, multi_class='ovr', average='weighted')
  1080. # print(auc_score)
  1081. saver.finalize()
  1082. test_acc.append(round(metric,3))
  1083. test_auc.append(round(auc_result,3))
  1084. print("Test Set: Accuracy: {:.4f}, AUC: {:.4f}".format(metric, auc_result))
  1085. print("Fold {} completed...Next Fold starting...".format(fold_number))
  1086. print("\n")
  1087. fold_number += 1
  1088. writer.close()
  1089. print("\n")
  1090. print("Train Acc: {:.4f} +/- {:.4f}, Train AUC: {:.4f} +/- {:.4f}".format(mean(train_acc), np.std(train_acc), mean(train_auc), np.std(train_auc)))
  1091. print("Val Acc: {:.4f} +/- {:.4f}, Val AUC: {:.4f} +/- {:.4f}".format(mean(val_acc), np.std(val_acc), mean(val_auc), np.std(val_auc)))
  1092. # # print("Val Accuracies of 5 Folds:", val_acc)
  1093. # print("Val AUCs of 5 Folds:", val_auc)
  1094. print("Test Acc: {:.4f} +/- {:.4f}, Test AUC: {:.4f} +/- {:.4f}".format(mean(test_acc), np.std(test_acc), mean(test_auc), np.std(test_auc)))
  1095. # print("Test Accuracies of 5 Folds:", test_acc)
  1096. # print("Test AUCs of 5 Folds:", test_auc)
  1097. y_pred_tr = np.transpose(y_pred_test)
  1098. final = []
  1099. for i in range(y_pred_tr.shape[0]):
  1100. final.append(mode(y_pred_tr[i]))
  1101. y_true = y.cpu().numpy()
  1102. final_np = np.array(final)
  1103. y_pred_auc = np.array(y_pred_auc)
  1104. sum = np.sum(y_pred_auc, axis=0)
  1105. sum = sum/5
  1106. target_names = ['class 0', 'class 1', 'class 2']
  1107. print("\n")
  1108. print("Test Accuracy (Hard Voting of 5 models):", round(accuracy_score(y_true, final_np), 4))
  1109. print("Test AUC (Hard Voting of 5 models):", round(roc_auc_score(y_onehot, sum, multi_class='ovr', average='weighted'), 4))
  1110. cm = confusion_matrix(y_true, final_np)
  1111. print(classification_report(y_true, final_np, target_names=target_names))
  1112. disp = ConfusionMatrixDisplay(confusion_matrix=cm)
  1113. disp.plot()
  1114. plt.show()
  1115. # %%
  1116. # %%
  1117. import logging
  1118. import os
  1119. import sys
  1120. from pathlib import Path
  1121. import pandas as pd
  1122. import monai
  1123. import numpy as np
  1124. import torch
  1125. import torch.nn as nn
  1126. from torch.utils.data import DataLoader as _TorchDataLoader
  1127. from torch.utils.data import Dataset
  1128. from torch.utils.tensorboard import SummaryWriter
  1129. from monai.data import decollate_batch, CSVSaver, DataLoader
  1130. from monai.data.utils import list_data_collate, set_rnd, worker_init_fn
  1131. from monai.metrics import ROCAUCMetric
  1132. from monai.transforms import ScaleIntensityRanged, NormalizeIntensityd, Activations, AsDiscrete, Compose, LoadImaged, RandRotated, Resized, ScaleIntensityd, RandFlipd
  1133. from monai.networks.nets import resnet10, resnet18, resnet34, resnet50
  1134. from sklearn.model_selection import StratifiedKFold
  1135. from statistics import mean, mode
  1136. import matplotlib.pyplot as plt
  1137. from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay, classification_report, accuracy_score, roc_auc_score
  1138. class CustomDataset(monai.data.Dataset):
  1139. def __init__(self, d1, d2):
  1140. self.d1 = d1
  1141. self.d2 = d2
  1142. def __getitem__(self, idx):
  1143. dict1 = self.d1.__getitem__(idx)
  1144. image1, label1 = dict1["img"], dict1["label"]
  1145. dict2 = self.d2.__getitem__(idx)
  1146. image2, label2 = dict2["img"], dict2["label"]
  1147. assert label1==label2
  1148. dict_1 = dict()
  1149. dict_1["img"] = image1
  1150. dict_1["label"] = label1
  1151. dict_2 = dict()
  1152. dict_2["img"] = image2
  1153. dict_2["label"] = label2
  1154. return dict_1, dict_2
  1155. def __len__(self):
  1156. return len(self.d1)
  1157. class MyEnsemble(nn.Module):
  1158. def __init__(self, modelA, modelB):
  1159. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  1160. super(MyEnsemble, self).__init__()
  1161. self.modelA = modelA
  1162. self.modelB = modelB
  1163. # Remove last linear layer
  1164. self.modelA.fc = nn.Identity()
  1165. self.modelB.fc = nn.Identity()
  1166. # Create new classifier
  1167. self.mlp1 = nn.Linear(1024,256).to(device)
  1168. self.mlp2 = nn.Linear(256,32).to(device)
  1169. self.classifier = nn.Linear(32,3).to(device)
  1170. def forward(self, i1, i2):
  1171. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  1172. x1 = self.modelA(i1) # clone to make sure x is not changed by inplace methods
  1173. x1 = x1.view(x1.size(0), -1).to(device)
  1174. x2 = self.modelB(i2)
  1175. x2 = x2.view(x2.size(0), -1).to(device)
  1176. x = torch.cat((x1, x2), dim=1).to(device)
  1177. x = nn.functional.relu(self.mlp1(x)).to(device)
  1178. x = nn.functional.relu(self.mlp2(x)).to(device)
  1179. x = self.classifier(x).to(device)
  1180. return x
  1181. def load_data(train_images, val_images, train_labels, val_labels, test_images, test_labels):
  1182. train_labels_array = np.array(train_labels, dtype=np.int64)
  1183. val_labels_array = np.array(val_labels, dtype=np.int64)
  1184. test_labels_array = np.array(test_labels, dtype=np.int64)
  1185. #loading the CT images
  1186. #ct_path = Path('/scratch/ajoshi83/Data/Reg_CT')
  1187. ct_path = Path('/scratch/ajoshi83/generated_ct')
  1188. #ct_path = Path('/data/amciilab/ajoshi83/generated_ct2')
  1189. ct_train_images_path = [os.path.join(ct_path, f) for f in train_images]
  1190. ct_val_images_path = [os.path.join(ct_path, f) for f in val_images]
  1191. ct_test_images_path = [os.path.join(ct_path, f) for f in test_images]
  1192. ct_train_files = [{"img": img, "label": label} for img, label in zip(ct_train_images_path, train_labels_array)]
  1193. ct_val_files = [{"img": img, "label": label} for img, label in zip(ct_val_images_path, val_labels_array)]
  1194. ct_test_files = [{"img": img, "label": label} for img, label in zip(ct_test_images_path, test_labels_array)]
  1195. #loading the MR images
  1196. mr_path = Path('/scratch/ajoshi83/Data/Reg_MR')
  1197. #mr_path = Path('/data/amciilab/ajoshi83/generated_mr2')
  1198. mr_train_images_path = [os.path.join(mr_path, f) for f in train_images]
  1199. mr_val_images_path = [os.path.join(mr_path, f) for f in val_images]
  1200. mr_test_images_path = [os.path.join(mr_path, f) for f in test_images]
  1201. mr_train_files = [{"img": img, "label": label} for img, label in zip(mr_train_images_path, train_labels_array)]
  1202. mr_val_files = [{"img": img, "label": label} for img, label in zip(mr_val_images_path, val_labels_array)]
  1203. mr_test_files = [{"img": img, "label": label} for img, label in zip(mr_test_images_path, test_labels_array)]
  1204. return ct_train_files, ct_val_files, mr_train_files, mr_val_files, ct_test_files, mr_test_files
  1205. #monai.config.print_config()
  1206. # old_stdout = sys.stdout
  1207. # log_file = open("mri_classification/logs/5fold_eval_fakemr.log","w")
  1208. # sys.stdout = log_file
  1209. print("this will be written to message.log")
  1210. # file_handler = logging.FileHandler(filename='logs/tmp.log')
  1211. # #stdout_handler = logging.StreamHandler(stream=sys.stdout)
  1212. # handlers = [file_handler]
  1213. # logging.basicConfig(
  1214. # level=logging.DEBUG,
  1215. # format='[%(asctime)s] {%(filename)s:%(lineno)d} %(levelname)s - %(message)s',
  1216. # handlers=handlers
  1217. # )
  1218. # logger = logging.getLogger('LOGGER_NAME')
  1219. # edit the path accordingly
  1220. gose = pd.read_csv('mri_classification/Pilot_169_train_val_test_split.csv')
  1221. X = []
  1222. Y = []
  1223. test_subjects = []
  1224. test_labels = []
  1225. torch.cuda.empty_cache()
  1226. for i in range(gose.shape[0]):
  1227. subj = gose['Main.GUID'][i]
  1228. subj_id = str(subj)[4:]
  1229. name = str(subj_id) + '.nii'
  1230. if gose['Set'][i]=='Train':
  1231. X.append(name)
  1232. Y.append(gose['Class'][i])
  1233. elif gose['Set'][i]=='Val':
  1234. X.append(name)
  1235. Y.append(gose['Class'][i])
  1236. elif gose['Set'][i]=='Test':
  1237. test_subjects.append(name)
  1238. test_labels.append(gose['Class'][i])
  1239. else:
  1240. print("Unknown Set: ", gose['Set'][i])
  1241. print("Total Subjects for 5-Fold (Train+Val):", len(X))
  1242. print("Test Subjects (kept separate):", len(test_subjects))
  1243. print("5-Fold Subjects Information:\n")
  1244. print("Subjects of Class 0: ", Y.count(0))
  1245. print("Subjects of Class 1: ", Y.count(1))
  1246. print("Subjects of Class 2: ", Y.count(2))
  1247. skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=1)
  1248. fold_number = 1
  1249. lr =1e-5
  1250. decay = 1e-5
  1251. X = np.array(X)
  1252. Y = np.array(Y)
  1253. train_acc, train_auc = [], []
  1254. val_acc, val_auc = [], []
  1255. test_acc, test_auc = [], []
  1256. y_pred_test = []
  1257. y_pred_auc = []
  1258. for train_index, test_index in skf.split(X, Y):
  1259. print("Starting Fold {}..!".format(fold_number))
  1260. x_train_fold, x_test_fold = X[train_index], X[test_index]
  1261. y_train_fold, y_test_fold = Y[train_index], Y[test_index]
  1262. y_train_fold = list(y_train_fold)
  1263. y_test_fold = list(y_test_fold)
  1264. x_train_fold = list(x_train_fold)
  1265. x_test_fold = list(x_test_fold)
  1266. # print("\n")
  1267. # print("Fold {} statistics:\n".format(fold_number))
  1268. # print("Train Subjects: {}".format(len(y_train_fold)))
  1269. # print("Subjects of Class 0: ", y_train_fold.count(0))
  1270. # print("Subjects of Class 1: ", y_train_fold.count(1))
  1271. # print("Subjects of Class 2: {}\n".format(y_train_fold.count(2)))
  1272. # print("Val Subjects: {}".format(len(y_test_fold)))
  1273. # print("Subjects of Class 0: ", y_test_fold.count(0))
  1274. # print("Subjects of Class 1: ", y_test_fold.count(1))
  1275. # print("Subjects of Class 2: ", y_test_fold.count(2))
  1276. ct_train_files, ct_val_files, mr_train_files, mr_val_files, ct_test_files, mr_test_files = load_data(x_train_fold, x_test_fold, y_train_fold, y_test_fold, test_subjects, test_labels)
  1277. # Define transforms for CT and MR respectively
  1278. ct_transforms = Compose(
  1279. [
  1280. LoadImaged(keys=["img"], ensure_channel_first=True),
  1281. #ScaleIntensityRanged(keys=["img"],a_min=0, a_max=85, b_min=0, b_max=1),
  1282. NormalizeIntensityd(keys=["img"]),
  1283. #Resized(keys=["img"], spatial_size=(96, 96, 96)),
  1284. ]
  1285. )
  1286. mr_transforms = Compose(
  1287. [
  1288. LoadImaged(keys=["img"], ensure_channel_first=True),
  1289. #ScaleIntensityRanged(keys=["img"],a_min=0, a_max=85, b_min=0, b_max=1),
  1290. NormalizeIntensityd(keys=["img"]),
  1291. # RandFlipd(keys=["img"], prob=1, spatial_axis=2),
  1292. # RandRotated(keys=["img"], prob=1, range_x=[0.4,0.4])
  1293. #ScaleIntensityRanged(keys=["img"],a_min=0, a_max=85, b_min=0, b_max=1),
  1294. ]
  1295. )
  1296. post_pred = Compose([Activations(softmax=True)])
  1297. post_label = Compose([AsDiscrete(to_onehot=3)])
  1298. # create a training data loader
  1299. ct_train_ds = monai.data.Dataset(data=ct_train_files, transform=ct_transforms)
  1300. mr_train_ds = monai.data.Dataset(data=mr_train_files, transform=mr_transforms)
  1301. combined_train_ds = CustomDataset(d1 = mr_train_ds, d2 = ct_train_ds)
  1302. train_loader = DataLoader(combined_train_ds, batch_size=1, shuffle=False, pin_memory=torch.cuda.is_available())
  1303. check_data1, check_data2 = monai.utils.misc.first(train_loader)
  1304. #print(check_data1["img"].shape, check_data1["label"], check_data2["img"].shape, check_data2["label"])
  1305. # create a validation data loader
  1306. ct_val_ds = monai.data.Dataset(data=ct_val_files, transform=ct_transforms)
  1307. mr_val_ds = monai.data.Dataset(data=mr_val_files, transform=mr_transforms)
  1308. combined_val_ds = CustomDataset(d1 = mr_val_ds, d2 = ct_val_ds)
  1309. val_loader = DataLoader(combined_val_ds, batch_size=1, shuffle=False, pin_memory=torch.cuda)
  1310. # create a test data loader
  1311. ct_test_ds = monai.data.Dataset(data=ct_test_files, transform=ct_transforms)
  1312. mr_test_ds = monai.data.Dataset(data=mr_test_files, transform=mr_transforms)
  1313. combined_test_ds = CustomDataset(d1 = mr_test_ds, d2 = ct_test_ds)
  1314. test_loader = DataLoader(combined_test_ds, batch_size=1, shuffle=False, pin_memory=torch.cuda)
  1315. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  1316. #print(device)
  1317. model_ct = resnet18(pretrained=False, spatial_dims=3, n_input_channels=1, num_classes=3).to(device)
  1318. #model_ct.load_state_dict(torch.load("models/best_val_resnet18.pth"))
  1319. model_mr = resnet18(pretrained=False, spatial_dims=3, n_input_channels=1, num_classes=3).to(device)
  1320. #model_mr.load_state_dict(torch.load("models/best_mr_val_resnet18.pth"))
  1321. # Freeze these models
  1322. for param in model_mr.parameters():
  1323. param.requires_grad_(True)
  1324. for param in model_ct.parameters():
  1325. param.requires_grad_(True)
  1326. # Create ensemble model
  1327. model = MyEnsemble(model_ct, model_mr)
  1328. #model.load_state_dict(torch.load("/data/amciilab/ajoshi83/models/CT_Fold{}.pth".format(fold_number)))
  1329. model.load_state_dict(torch.load("/scratch/ajoshi83/models_sbatch/Fake_CT_Fold{}_ensemble_Adam_{}.pth".format(int(fold_number), lr)))
  1330. #model.load_state_dict(torch.load("/scratch/ajoshi83/models_sbatch/Both_Real_Fold{}_ensemble_Adam_{}_{}.pth".format(int(fold_number), lr, decay)))
  1331. # loss_function = torch.nn.CrossEntropyLoss()
  1332. # optimizer = torch.optim.Adam(model.parameters(),lr= 1e-5, weight_decay=1e-3)
  1333. auc_metric = ROCAUCMetric(average="weighted")
  1334. # starting evaluation
  1335. val_interval = 1
  1336. best_metric = -1
  1337. best_metric_epoch = -1
  1338. best_val_loss = 2
  1339. writer = SummaryWriter()
  1340. with torch.no_grad():
  1341. num_correct = 0.0
  1342. metric_count = 0
  1343. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  1344. y = torch.tensor([], dtype=torch.long, device=device)
  1345. saver = CSVSaver(output_dir="mri_classification/train_output")
  1346. for batch_data in train_loader:
  1347. #step += 1
  1348. #print(batch_data["img"])
  1349. ct_batch, mr_batch = batch_data[0], batch_data[1]
  1350. #print(ct_batch)
  1351. #print(mr_batch)
  1352. ct_inputs, ct_labels = ct_batch["img"].to(device), ct_batch["label"].to(device)
  1353. mr_inputs, mr_labels = mr_batch["img"].to(device), mr_batch["label"].to(device)
  1354. train_outputs = model(ct_inputs, mr_inputs).argmax(dim=1)
  1355. y_pred = torch.cat([y_pred, model(ct_inputs, mr_inputs)], dim=0)
  1356. y = torch.cat([y, mr_labels], dim=0)
  1357. value = torch.eq(train_outputs, mr_labels)
  1358. metric_count += len(value)
  1359. num_correct += value.sum().item()
  1360. #saver.save_batch(train_outputs, train_data["img"].meta)
  1361. metric = num_correct / metric_count
  1362. # print("val evaluation metric:", metric)
  1363. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  1364. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred, detach = True)]
  1365. y_onehot = torch.stack(y_onehot, dim=0)
  1366. y_pred_act = torch.stack(y_pred_act, dim=0)
  1367. y_onehot = y_onehot.to(device="cpu")
  1368. y_pred_act = y_pred_act.to(device="cpu")
  1369. auc_metric(y_pred_act, y_onehot)
  1370. auc_result = auc_metric.aggregate()
  1371. auc_metric.reset()
  1372. saver.finalize()
  1373. train_acc.append(round(metric,3))
  1374. train_auc.append(round(auc_result,3))
  1375. print("Train Set: Accuracy: {:.4f}, AUC: {:.4f}".format(metric, auc_result))
  1376. with torch.no_grad():
  1377. num_correct = 0.0
  1378. metric_count = 0
  1379. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  1380. y = torch.tensor([], dtype=torch.long, device=device)
  1381. saver = CSVSaver(output_dir="mri_classification/val_output")
  1382. for val_data in val_loader:
  1383. ct_val_data, mr_val_data = val_data[0], val_data[1]
  1384. ct_val_images, ct_val_labels = ct_val_data["img"].to(device), ct_val_data["label"].to(device)
  1385. mr_val_images, mr_val_labels = mr_val_data["img"].to(device), mr_val_data["label"].to(device)
  1386. y_pred = torch.cat([y_pred, model(ct_val_images, mr_val_images)], dim=0)
  1387. #print(y_pred)
  1388. y = torch.cat([y, mr_val_labels], dim=0)
  1389. acc_value = torch.eq(y_pred.argmax(dim=1), y)
  1390. acc_metric = acc_value.sum().item() / len(acc_value)
  1391. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  1392. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred,detach=True)]
  1393. y_onehot = torch.stack(y_onehot, dim=0)
  1394. y_pred_act = torch.stack(y_pred_act, dim=0)
  1395. y_onehot = y_onehot.to(device="cpu")
  1396. y_pred_act = y_pred_act.to(device="cpu")
  1397. auc_metric(y_pred_act, y_onehot)
  1398. auc_result = auc_metric.aggregate()
  1399. auc_metric.reset()
  1400. saver.finalize()
  1401. val_acc.append(round(acc_metric,3))
  1402. val_auc.append(round(auc_result,3))
  1403. print("Val Set: Accuracy: {:.4f}, AUC: {:.4f}".format(acc_metric, auc_result))
  1404. with torch.no_grad():
  1405. num_correct = 0.0
  1406. metric_count = 0
  1407. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  1408. y = torch.tensor([], dtype=torch.long, device=device)
  1409. saver = CSVSaver(output_dir="mri_classification/test_output")
  1410. y_pred_mid = []
  1411. for test_data in test_loader:
  1412. ct_test_data, mr_test_data = test_data[0], test_data[1]
  1413. ct_test_images, ct_test_labels = ct_test_data["img"].to(device), ct_test_data["label"].to(device)
  1414. mr_test_images, mr_test_labels = mr_test_data["img"].to(device), mr_test_data["label"].to(device)
  1415. test_outputs = model(ct_test_images, mr_test_images).argmax(dim=1)
  1416. y_pred_mid.append(test_outputs.cpu().numpy()[0])
  1417. #softmax_op = torch.nn.functional.softmax(model(ct_test_images, mr_test_images)).argmax(dim=1)
  1418. #y_pred_mid.append(softmax_op.cpu().numpy()[0])
  1419. y_pred = torch.cat([y_pred, model(ct_test_images, mr_test_images)], dim=0)
  1420. y = torch.cat([y, mr_test_labels], dim=0)
  1421. value = torch.eq(test_outputs, mr_test_labels)
  1422. metric_count += len(value)
  1423. num_correct += value.sum().item()
  1424. #saver.save_batch(test_outputs, test_data["img"].meta)
  1425. y_pred_test.append(y_pred_mid)
  1426. metric = num_correct / metric_count
  1427. # print("test evaluation metric:", metric)
  1428. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  1429. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred, detach=True)]
  1430. y_onehot = torch.stack(y_onehot, dim=0)
  1431. y_pred_act = torch.stack(y_pred_act, dim=0)
  1432. y_onehot = y_onehot.to(device="cpu")
  1433. y_pred_act = y_pred_act.to(device="cpu")
  1434. y_pred_auc.append(y_pred_act)
  1435. auc_metric(y_pred_act, y_onehot)
  1436. auc_result = auc_metric.aggregate()
  1437. auc_metric.reset()
  1438. # auc_score =roc_auc_score(y_onehot, y_pred_act, multi_class='ovr', average='weighted')
  1439. # print(auc_score)
  1440. saver.finalize()
  1441. test_acc.append(round(metric,3))
  1442. test_auc.append(round(auc_result,3))
  1443. print("Test Set: Accuracy: {:.4f}, AUC: {:.4f}".format(metric, auc_result))
  1444. print("Fold {} completed...Next Fold starting...".format(fold_number))
  1445. print("\n")
  1446. fold_number += 1
  1447. writer.close()
  1448. print("\n")
  1449. print("Train Acc: {:.4f} +/- {:.4f}, Train AUC: {:.4f} +/- {:.4f}".format(mean(train_acc), np.std(train_acc), mean(train_auc), np.std(train_auc)))
  1450. print("Val Acc: {:.4f} +/- {:.4f}, Val AUC: {:.4f} +/- {:.4f}".format(mean(val_acc), np.std(val_acc), mean(val_auc), np.std(val_auc)))
  1451. # # print("Val Accuracies of 5 Folds:", val_acc)
  1452. # print("Val AUCs of 5 Folds:", val_auc)
  1453. print("Test Acc: {:.4f} +/- {:.4f}, Test AUC: {:.4f} +/- {:.4f}".format(mean(test_acc), np.std(test_acc), mean(test_auc), np.std(test_auc)))
  1454. # print("Test Accuracies of 5 Folds:", test_acc)
  1455. # print("Test AUCs of 5 Folds:", test_auc)
  1456. y_pred_tr = np.transpose(y_pred_test)
  1457. final = []
  1458. for i in range(y_pred_tr.shape[0]):
  1459. final.append(mode(y_pred_tr[i]))
  1460. y_true = y.cpu().numpy()
  1461. final_np = np.array(final)
  1462. y_pred_auc = np.array(y_pred_auc)
  1463. sum = np.sum(y_pred_auc, axis=0)
  1464. sum = sum/5
  1465. target_names = ['class 0', 'class 1', 'class 2']
  1466. print("\n")
  1467. print("Test Accuracy (Hard Voting of 5 models):", round(accuracy_score(y_true, final_np), 4))
  1468. print("Test AUC (Hard Voting of 5 models):", round(roc_auc_score(y_onehot, sum, multi_class='ovr', average='weighted'), 4))
  1469. cm = confusion_matrix(y_true, final_np)
  1470. print(classification_report(y_true, final_np, target_names=target_names))
  1471. disp = ConfusionMatrixDisplay(confusion_matrix=cm)
  1472. disp.plot()
  1473. plt.show()
  1474. # %%
  1475. # %%
  1476. import logging
  1477. import os
  1478. import sys
  1479. from pathlib import Path
  1480. import pandas as pd
  1481. import monai
  1482. import numpy as np
  1483. import torch
  1484. import torch.nn as nn
  1485. from torch.utils.data import DataLoader as _TorchDataLoader
  1486. from torch.utils.data import Dataset
  1487. from torch.utils.tensorboard import SummaryWriter
  1488. from monai.data import decollate_batch, CSVSaver, DataLoader
  1489. from monai.data.utils import list_data_collate, set_rnd, worker_init_fn
  1490. from monai.metrics import ROCAUCMetric
  1491. from monai.transforms import ScaleIntensityRanged, NormalizeIntensityd, Activations, AsDiscrete, Compose, LoadImaged, RandRotated, Resized, ScaleIntensityd, RandFlipd
  1492. from monai.networks.nets import resnet10, resnet18, resnet34, resnet50
  1493. from sklearn.model_selection import StratifiedKFold
  1494. from statistics import mean, mode
  1495. import matplotlib.pyplot as plt
  1496. from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay, classification_report, accuracy_score, roc_auc_score
  1497. class CustomDataset(monai.data.Dataset):
  1498. def __init__(self, d1, d2):
  1499. self.d1 = d1
  1500. self.d2 = d2
  1501. def __getitem__(self, idx):
  1502. dict1 = self.d1.__getitem__(idx)
  1503. image1, label1 = dict1["img"], dict1["label"]
  1504. dict2 = self.d2.__getitem__(idx)
  1505. image2, label2 = dict2["img"], dict2["label"]
  1506. assert label1==label2
  1507. dict_1 = dict()
  1508. dict_1["img"] = image1
  1509. dict_1["label"] = label1
  1510. dict_2 = dict()
  1511. dict_2["img"] = image2
  1512. dict_2["label"] = label2
  1513. return dict_1, dict_2
  1514. def __len__(self):
  1515. return len(self.d1)
  1516. class MyEnsemble(nn.Module):
  1517. def __init__(self, modelA, modelB):
  1518. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  1519. super(MyEnsemble, self).__init__()
  1520. self.modelA = modelA
  1521. self.modelB = modelB
  1522. # Remove last linear layer
  1523. self.modelA.fc = nn.Identity()
  1524. self.modelB.fc = nn.Identity()
  1525. # Create new classifier
  1526. self.mlp1 = nn.Linear(1024,256).to(device)
  1527. self.mlp2 = nn.Linear(256,32).to(device)
  1528. self.classifier = nn.Linear(32,3).to(device)
  1529. def forward(self, i1, i2):
  1530. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  1531. x1 = self.modelA(i1) # clone to make sure x is not changed by inplace methods
  1532. x1 = x1.view(x1.size(0), -1).to(device)
  1533. x2 = self.modelB(i2)
  1534. x2 = x2.view(x2.size(0), -1).to(device)
  1535. x = torch.cat((x1, x2), dim=1).to(device)
  1536. x = nn.functional.relu(self.mlp1(x)).to(device)
  1537. x = nn.functional.relu(self.mlp2(x)).to(device)
  1538. x = self.classifier(x).to(device)
  1539. return x
  1540. def load_data(train_images, val_images, train_labels, val_labels, test_images, test_labels):
  1541. train_labels_array = np.array(train_labels, dtype=np.int64)
  1542. val_labels_array = np.array(val_labels, dtype=np.int64)
  1543. test_labels_array = np.array(test_labels, dtype=np.int64)
  1544. #loading the CT images
  1545. ct_path = Path('/scratch/ajoshi83/Data/Reg_CT')
  1546. #ct_path = Path('/scratch/ajoshi83/generated_ct')
  1547. #ct_path = Path('/data/amciilab/ajoshi83/generated_ct2')
  1548. ct_train_images_path = [os.path.join(ct_path, f) for f in train_images]
  1549. ct_val_images_path = [os.path.join(ct_path, f) for f in val_images]
  1550. ct_test_images_path = [os.path.join(ct_path, f) for f in test_images]
  1551. ct_train_files = [{"img": img, "label": label} for img, label in zip(ct_train_images_path, train_labels_array)]
  1552. ct_val_files = [{"img": img, "label": label} for img, label in zip(ct_val_images_path, val_labels_array)]
  1553. ct_test_files = [{"img": img, "label": label} for img, label in zip(ct_test_images_path, test_labels_array)]
  1554. #loading the MR images
  1555. #mr_path = Path('/scratch/ajoshi83/Data/Reg_MR')
  1556. mr_path = Path('/scratch/ajoshi83/generated_mr')
  1557. #mr_path = Path('/data/amciilab/ajoshi83/generated_mr2')
  1558. mr_train_images_path = [os.path.join(mr_path, f) for f in train_images]
  1559. mr_val_images_path = [os.path.join(mr_path, f) for f in val_images]
  1560. mr_test_images_path = [os.path.join(mr_path, f) for f in test_images]
  1561. mr_train_files = [{"img": img, "label": label} for img, label in zip(mr_train_images_path, train_labels_array)]
  1562. mr_val_files = [{"img": img, "label": label} for img, label in zip(mr_val_images_path, val_labels_array)]
  1563. mr_test_files = [{"img": img, "label": label} for img, label in zip(mr_test_images_path, test_labels_array)]
  1564. return ct_train_files, ct_val_files, mr_train_files, mr_val_files, ct_test_files, mr_test_files
  1565. #monai.config.print_config()
  1566. # old_stdout = sys.stdout
  1567. # log_file = open("mri_classification/logs/5fold_eval_fakemr.log","w")
  1568. # sys.stdout = log_file
  1569. print("this will be written to message.log")
  1570. # file_handler = logging.FileHandler(filename='logs/tmp.log')
  1571. # #stdout_handler = logging.StreamHandler(stream=sys.stdout)
  1572. # handlers = [file_handler]
  1573. # logging.basicConfig(
  1574. # level=logging.DEBUG,
  1575. # format='[%(asctime)s] {%(filename)s:%(lineno)d} %(levelname)s - %(message)s',
  1576. # handlers=handlers
  1577. # )
  1578. # logger = logging.getLogger('LOGGER_NAME')
  1579. # edit the path accordingly
  1580. gose = pd.read_csv('mri_classification/Pilot_169_train_val_test_split.csv')
  1581. X = []
  1582. Y = []
  1583. test_subjects = []
  1584. test_labels = []
  1585. torch.cuda.empty_cache()
  1586. for i in range(gose.shape[0]):
  1587. subj = gose['Main.GUID'][i]
  1588. subj_id = str(subj)[4:]
  1589. name = str(subj_id) + '.nii'
  1590. if gose['Set'][i]=='Train':
  1591. X.append(name)
  1592. Y.append(gose['Class'][i])
  1593. elif gose['Set'][i]=='Val':
  1594. X.append(name)
  1595. Y.append(gose['Class'][i])
  1596. elif gose['Set'][i]=='Test':
  1597. test_subjects.append(name)
  1598. test_labels.append(gose['Class'][i])
  1599. else:
  1600. print("Unknown Set: ", gose['Set'][i])
  1601. print("Total Subjects for 5-Fold (Train+Val):", len(X))
  1602. print("Test Subjects (kept separate):", len(test_subjects))
  1603. print("5-Fold Subjects Information:\n")
  1604. print("Subjects of Class 0: ", Y.count(0))
  1605. print("Subjects of Class 1: ", Y.count(1))
  1606. print("Subjects of Class 2: ", Y.count(2))
  1607. skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=1)
  1608. fold_number = 1
  1609. lr =1e-5
  1610. decay = 1e-5
  1611. X = np.array(X)
  1612. Y = np.array(Y)
  1613. train_acc, train_auc = [], []
  1614. val_acc, val_auc = [], []
  1615. test_acc, test_auc = [], []
  1616. y_pred_test = []
  1617. y_pred_auc = []
  1618. for train_index, test_index in skf.split(X, Y):
  1619. print("Starting Fold {}..!".format(fold_number))
  1620. x_train_fold, x_test_fold = X[train_index], X[test_index]
  1621. y_train_fold, y_test_fold = Y[train_index], Y[test_index]
  1622. y_train_fold = list(y_train_fold)
  1623. y_test_fold = list(y_test_fold)
  1624. x_train_fold = list(x_train_fold)
  1625. x_test_fold = list(x_test_fold)
  1626. # print("\n")
  1627. # print("Fold {} statistics:\n".format(fold_number))
  1628. # print("Train Subjects: {}".format(len(y_train_fold)))
  1629. # print("Subjects of Class 0: ", y_train_fold.count(0))
  1630. # print("Subjects of Class 1: ", y_train_fold.count(1))
  1631. # print("Subjects of Class 2: {}\n".format(y_train_fold.count(2)))
  1632. # print("Val Subjects: {}".format(len(y_test_fold)))
  1633. # print("Subjects of Class 0: ", y_test_fold.count(0))
  1634. # print("Subjects of Class 1: ", y_test_fold.count(1))
  1635. # print("Subjects of Class 2: ", y_test_fold.count(2))
  1636. ct_train_files, ct_val_files, mr_train_files, mr_val_files, ct_test_files, mr_test_files = load_data(x_train_fold, x_test_fold, y_train_fold, y_test_fold, test_subjects, test_labels)
  1637. # Define transforms for CT and MR respectively
  1638. ct_transforms = Compose(
  1639. [
  1640. LoadImaged(keys=["img"], ensure_channel_first=True),
  1641. ScaleIntensityRanged(keys=["img"],a_min=0, a_max=85, b_min=0, b_max=1),
  1642. NormalizeIntensityd(keys=["img"]),
  1643. #Resized(keys=["img"], spatial_size=(96, 96, 96)),
  1644. ]
  1645. )
  1646. mr_transforms = Compose(
  1647. [
  1648. LoadImaged(keys=["img"], ensure_channel_first=True),
  1649. #ScaleIntensityRanged(keys=["img"],a_min=0, a_max=85, b_min=0, b_max=1),
  1650. NormalizeIntensityd(keys=["img"]),
  1651. # RandFlipd(keys=["img"], prob=1, spatial_axis=2),
  1652. # RandRotated(keys=["img"], prob=1, range_x=[0.4,0.4])
  1653. #ScaleIntensityRanged(keys=["img"],a_min=0, a_max=85, b_min=0, b_max=1),
  1654. ]
  1655. )
  1656. post_pred = Compose([Activations(softmax=True)])
  1657. post_label = Compose([AsDiscrete(to_onehot=3)])
  1658. # create a training data loader
  1659. ct_train_ds = monai.data.Dataset(data=ct_train_files, transform=ct_transforms)
  1660. mr_train_ds = monai.data.Dataset(data=mr_train_files, transform=mr_transforms)
  1661. combined_train_ds = CustomDataset(d1 = mr_train_ds, d2 = ct_train_ds)
  1662. train_loader = DataLoader(combined_train_ds, batch_size=1, shuffle=False, pin_memory=torch.cuda.is_available())
  1663. check_data1, check_data2 = monai.utils.misc.first(train_loader)
  1664. #print(check_data1["img"].shape, check_data1["label"], check_data2["img"].shape, check_data2["label"])
  1665. # create a validation data loader
  1666. ct_val_ds = monai.data.Dataset(data=ct_val_files, transform=ct_transforms)
  1667. mr_val_ds = monai.data.Dataset(data=mr_val_files, transform=mr_transforms)
  1668. combined_val_ds = CustomDataset(d1 = mr_val_ds, d2 = ct_val_ds)
  1669. val_loader = DataLoader(combined_val_ds, batch_size=1, shuffle=False, pin_memory=torch.cuda)
  1670. # create a test data loader
  1671. ct_test_ds = monai.data.Dataset(data=ct_test_files, transform=ct_transforms)
  1672. mr_test_ds = monai.data.Dataset(data=mr_test_files, transform=mr_transforms)
  1673. combined_test_ds = CustomDataset(d1 = mr_test_ds, d2 = ct_test_ds)
  1674. test_loader = DataLoader(combined_test_ds, batch_size=1, shuffle=False, pin_memory=torch.cuda)
  1675. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  1676. #print(device)
  1677. model_ct = resnet18(pretrained=False, spatial_dims=3, n_input_channels=1, num_classes=3).to(device)
  1678. #model_ct.load_state_dict(torch.load("models/best_val_resnet18.pth"))
  1679. model_mr = resnet18(pretrained=False, spatial_dims=3, n_input_channels=1, num_classes=3).to(device)
  1680. #model_mr.load_state_dict(torch.load("models/best_mr_val_resnet18.pth"))
  1681. # Freeze these models
  1682. for param in model_mr.parameters():
  1683. param.requires_grad_(True)
  1684. for param in model_ct.parameters():
  1685. param.requires_grad_(True)
  1686. # Create ensemble model
  1687. model = MyEnsemble(model_ct, model_mr)
  1688. #model.load_state_dict(torch.load("/data/amciilab/ajoshi83/models/CT_Fold{}.pth".format(fold_number)))
  1689. model.load_state_dict(torch.load("/scratch/ajoshi83/models_sbatch/Fake_MR_Fold{}_ensemble_AdamW_{}.pth".format(int(fold_number), lr)))
  1690. #model.load_state_dict(torch.load("/scratch/ajoshi83/models_sbatch/Both_Real_Fold{}_ensemble_Adam_{}_{}.pth".format(int(fold_number), lr, decay)))
  1691. # loss_function = torch.nn.CrossEntropyLoss()
  1692. # optimizer = torch.optim.Adam(model.parameters(),lr= 1e-5, weight_decay=1e-3)
  1693. auc_metric = ROCAUCMetric(average="weighted")
  1694. # starting evaluation
  1695. val_interval = 1
  1696. best_metric = -1
  1697. best_metric_epoch = -1
  1698. best_val_loss = 2
  1699. writer = SummaryWriter()
  1700. with torch.no_grad():
  1701. num_correct = 0.0
  1702. metric_count = 0
  1703. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  1704. y = torch.tensor([], dtype=torch.long, device=device)
  1705. saver = CSVSaver(output_dir="mri_classification/train_output")
  1706. for batch_data in train_loader:
  1707. #step += 1
  1708. #print(batch_data["img"])
  1709. ct_batch, mr_batch = batch_data[0], batch_data[1]
  1710. #print(ct_batch)
  1711. #print(mr_batch)
  1712. ct_inputs, ct_labels = ct_batch["img"].to(device), ct_batch["label"].to(device)
  1713. mr_inputs, mr_labels = mr_batch["img"].to(device), mr_batch["label"].to(device)
  1714. train_outputs = model(ct_inputs, mr_inputs).argmax(dim=1)
  1715. y_pred = torch.cat([y_pred, model(ct_inputs, mr_inputs)], dim=0)
  1716. y = torch.cat([y, mr_labels], dim=0)
  1717. value = torch.eq(train_outputs, mr_labels)
  1718. metric_count += len(value)
  1719. num_correct += value.sum().item()
  1720. #saver.save_batch(train_outputs, train_data["img"].meta)
  1721. metric = num_correct / metric_count
  1722. # print("val evaluation metric:", metric)
  1723. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  1724. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred, detach = True)]
  1725. y_onehot = torch.stack(y_onehot, dim=0)
  1726. y_pred_act = torch.stack(y_pred_act, dim=0)
  1727. y_onehot = y_onehot.to(device="cpu")
  1728. y_pred_act = y_pred_act.to(device="cpu")
  1729. auc_metric(y_pred_act, y_onehot)
  1730. auc_result = auc_metric.aggregate()
  1731. auc_metric.reset()
  1732. saver.finalize()
  1733. train_acc.append(round(metric,3))
  1734. train_auc.append(round(auc_result,3))
  1735. print("Train Set: Accuracy: {:.4f}, AUC: {:.4f}".format(metric, auc_result))
  1736. with torch.no_grad():
  1737. num_correct = 0.0
  1738. metric_count = 0
  1739. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  1740. y = torch.tensor([], dtype=torch.long, device=device)
  1741. saver = CSVSaver(output_dir="mri_classification/val_output")
  1742. for val_data in val_loader:
  1743. ct_val_data, mr_val_data = val_data[0], val_data[1]
  1744. ct_val_images, ct_val_labels = ct_val_data["img"].to(device), ct_val_data["label"].to(device)
  1745. mr_val_images, mr_val_labels = mr_val_data["img"].to(device), mr_val_data["label"].to(device)
  1746. y_pred = torch.cat([y_pred, model(ct_val_images, mr_val_images)], dim=0)
  1747. #print(y_pred)
  1748. y = torch.cat([y, mr_val_labels], dim=0)
  1749. acc_value = torch.eq(y_pred.argmax(dim=1), y)
  1750. acc_metric = acc_value.sum().item() / len(acc_value)
  1751. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  1752. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred,detach=True)]
  1753. y_onehot = torch.stack(y_onehot, dim=0)
  1754. y_pred_act = torch.stack(y_pred_act, dim=0)
  1755. y_onehot = y_onehot.to(device="cpu")
  1756. y_pred_act = y_pred_act.to(device="cpu")
  1757. auc_metric(y_pred_act, y_onehot)
  1758. auc_result = auc_metric.aggregate()
  1759. auc_metric.reset()
  1760. saver.finalize()
  1761. val_acc.append(round(acc_metric,3))
  1762. val_auc.append(round(auc_result,3))
  1763. print("Val Set: Accuracy: {:.4f}, AUC: {:.4f}".format(acc_metric, auc_result))
  1764. with torch.no_grad():
  1765. num_correct = 0.0
  1766. metric_count = 0
  1767. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  1768. y = torch.tensor([], dtype=torch.long, device=device)
  1769. saver = CSVSaver(output_dir="mri_classification/test_output")
  1770. y_pred_mid = []
  1771. for test_data in test_loader:
  1772. ct_test_data, mr_test_data = test_data[0], test_data[1]
  1773. ct_test_images, ct_test_labels = ct_test_data["img"].to(device), ct_test_data["label"].to(device)
  1774. mr_test_images, mr_test_labels = mr_test_data["img"].to(device), mr_test_data["label"].to(device)
  1775. test_outputs = model(ct_test_images, mr_test_images).argmax(dim=1)
  1776. y_pred_mid.append(test_outputs.cpu().numpy()[0])
  1777. #softmax_op = torch.nn.functional.softmax(model(ct_test_images, mr_test_images)).argmax(dim=1)
  1778. #y_pred_mid.append(softmax_op.cpu().numpy()[0])
  1779. y_pred = torch.cat([y_pred, model(ct_test_images, mr_test_images)], dim=0)
  1780. y = torch.cat([y, mr_test_labels], dim=0)
  1781. value = torch.eq(test_outputs, mr_test_labels)
  1782. metric_count += len(value)
  1783. num_correct += value.sum().item()
  1784. #saver.save_batch(test_outputs, test_data["img"].meta)
  1785. y_pred_test.append(y_pred_mid)
  1786. metric = num_correct / metric_count
  1787. # print("test evaluation metric:", metric)
  1788. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  1789. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred, detach=True)]
  1790. y_onehot = torch.stack(y_onehot, dim=0)
  1791. y_pred_act = torch.stack(y_pred_act, dim=0)
  1792. y_onehot = y_onehot.to(device="cpu")
  1793. y_pred_act = y_pred_act.to(device="cpu")
  1794. y_pred_auc.append(y_pred_act)
  1795. auc_metric(y_pred_act, y_onehot)
  1796. auc_result = auc_metric.aggregate()
  1797. auc_metric.reset()
  1798. # auc_score =roc_auc_score(y_onehot, y_pred_act, multi_class='ovr', average='weighted')
  1799. # print(auc_score)
  1800. saver.finalize()
  1801. test_acc.append(round(metric,3))
  1802. test_auc.append(round(auc_result,3))
  1803. print("Test Set: Accuracy: {:.4f}, AUC: {:.4f}".format(metric, auc_result))
  1804. print("Fold {} completed...Next Fold starting...".format(fold_number))
  1805. print("\n")
  1806. fold_number += 1
  1807. writer.close()
  1808. print("\n")
  1809. print("Train Acc: {:.4f} +/- {:.4f}, Train AUC: {:.4f} +/- {:.4f}".format(mean(train_acc), np.std(train_acc), mean(train_auc), np.std(train_auc)))
  1810. print("Val Acc: {:.4f} +/- {:.4f}, Val AUC: {:.4f} +/- {:.4f}".format(mean(val_acc), np.std(val_acc), mean(val_auc), np.std(val_auc)))
  1811. # # print("Val Accuracies of 5 Folds:", val_acc)
  1812. # print("Val AUCs of 5 Folds:", val_auc)
  1813. print("Test Acc: {:.4f} +/- {:.4f}, Test AUC: {:.4f} +/- {:.4f}".format(mean(test_acc), np.std(test_acc), mean(test_auc), np.std(test_auc)))
  1814. # print("Test Accuracies of 5 Folds:", test_acc)
  1815. # print("Test AUCs of 5 Folds:", test_auc)
  1816. y_pred_tr = np.transpose(y_pred_test)
  1817. final = []
  1818. for i in range(y_pred_tr.shape[0]):
  1819. final.append(mode(y_pred_tr[i]))
  1820. y_true = y.cpu().numpy()
  1821. final_np = np.array(final)
  1822. y_pred_auc = np.array(y_pred_auc)
  1823. sum = np.sum(y_pred_auc, axis=0)
  1824. sum = sum/5
  1825. target_names = ['class 0', 'class 1', 'class 2']
  1826. print("\n")
  1827. print("Test Accuracy (Hard Voting of 5 models):", round(accuracy_score(y_true, final_np), 4))
  1828. print("Test AUC (Hard Voting of 5 models):", round(roc_auc_score(y_onehot, sum, multi_class='ovr', average='weighted'), 4))
  1829. cm = confusion_matrix(y_true, final_np)
  1830. print(classification_report(y_true, final_np, target_names=target_names))
  1831. disp = ConfusionMatrixDisplay(confusion_matrix=cm)
  1832. disp.plot()
  1833. plt.show()
  1834. # %%
  1835. # %%
  1836. import logging
  1837. import os
  1838. import sys
  1839. from pathlib import Path
  1840. import pandas as pd
  1841. import monai
  1842. import numpy as np
  1843. import torch
  1844. import torch.nn as nn
  1845. from torch.utils.data import DataLoader as _TorchDataLoader
  1846. from torch.utils.data import Dataset
  1847. from torch.utils.tensorboard import SummaryWriter
  1848. from monai.data import decollate_batch, CSVSaver, DataLoader
  1849. from monai.data.utils import list_data_collate, set_rnd, worker_init_fn
  1850. from monai.metrics import ROCAUCMetric
  1851. from monai.transforms import ScaleIntensityRanged, NormalizeIntensityd, Activations, AsDiscrete, Compose, LoadImaged, RandRotated, Resized, ScaleIntensityd, RandFlipd
  1852. from monai.networks.nets import resnet10, resnet18, resnet34, resnet50
  1853. from sklearn.model_selection import StratifiedKFold
  1854. from statistics import mean, mode
  1855. import matplotlib.pyplot as plt
  1856. from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay, classification_report, accuracy_score, roc_auc_score
  1857. class CustomDataset(monai.data.Dataset):
  1858. def __init__(self, d1, d2):
  1859. self.d1 = d1
  1860. self.d2 = d2
  1861. def __getitem__(self, idx):
  1862. dict1 = self.d1.__getitem__(idx)
  1863. image1, label1 = dict1["img"], dict1["label"]
  1864. dict2 = self.d2.__getitem__(idx)
  1865. image2, label2 = dict2["img"], dict2["label"]
  1866. assert label1==label2
  1867. dict_1 = dict()
  1868. dict_1["img"] = image1
  1869. dict_1["label"] = label1
  1870. dict_2 = dict()
  1871. dict_2["img"] = image2
  1872. dict_2["label"] = label2
  1873. return dict_1, dict_2
  1874. def __len__(self):
  1875. return len(self.d1)
  1876. class MyEnsemble(nn.Module):
  1877. def __init__(self, modelA, modelB):
  1878. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  1879. super(MyEnsemble, self).__init__()
  1880. self.modelA = modelA
  1881. self.modelB = modelB
  1882. # Remove last linear layer
  1883. self.modelA.fc = nn.Identity()
  1884. self.modelB.fc = nn.Identity()
  1885. # Create new classifier
  1886. self.mlp1 = nn.Linear(1024,256).to(device)
  1887. self.mlp2 = nn.Linear(256,32).to(device)
  1888. self.classifier = nn.Linear(32,3).to(device)
  1889. def forward(self, i1, i2):
  1890. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  1891. x1 = self.modelA(i1) # clone to make sure x is not changed by inplace methods
  1892. x1 = x1.view(x1.size(0), -1).to(device)
  1893. x2 = self.modelB(i2)
  1894. x2 = x2.view(x2.size(0), -1).to(device)
  1895. x = torch.cat((x1, x2), dim=1).to(device)
  1896. x = nn.functional.relu(self.mlp1(x)).to(device)
  1897. x = nn.functional.relu(self.mlp2(x)).to(device)
  1898. x = self.classifier(x).to(device)
  1899. return x
  1900. def load_data(train_images, val_images, train_labels, val_labels, test_images, test_labels):
  1901. train_labels_array = np.array(train_labels, dtype=np.int64)
  1902. val_labels_array = np.array(val_labels, dtype=np.int64)
  1903. test_labels_array = np.array(test_labels, dtype=np.int64)
  1904. #loading the CT images
  1905. ct_path = Path('Data/Reg_MR')
  1906. #ct_path = Path('/data/amciilab/ajoshi83/generated_ct2')
  1907. ct_train_images_path = [os.path.join(ct_path, f) for f in train_images]
  1908. ct_val_images_path = [os.path.join(ct_path, f) for f in val_images]
  1909. ct_test_images_path = [os.path.join(ct_path, f) for f in test_images]
  1910. ct_train_files = [{"img": img, "label": label} for img, label in zip(ct_train_images_path, train_labels_array)]
  1911. ct_val_files = [{"img": img, "label": label} for img, label in zip(ct_val_images_path, val_labels_array)]
  1912. ct_test_files = [{"img": img, "label": label} for img, label in zip(ct_test_images_path, test_labels_array)]
  1913. #loading the MR images
  1914. mr_path = Path('Data/Reg_MR')
  1915. #mr_path = Path('/data/amciilab/ajoshi83/generated_mr2')
  1916. mr_train_images_path = [os.path.join(mr_path, f) for f in train_images]
  1917. mr_val_images_path = [os.path.join(mr_path, f) for f in val_images]
  1918. mr_test_images_path = [os.path.join(mr_path, f) for f in test_images]
  1919. mr_train_files = [{"img": img, "label": label} for img, label in zip(mr_train_images_path, train_labels_array)]
  1920. mr_val_files = [{"img": img, "label": label} for img, label in zip(mr_val_images_path, val_labels_array)]
  1921. mr_test_files = [{"img": img, "label": label} for img, label in zip(mr_test_images_path, test_labels_array)]
  1922. return ct_train_files, ct_val_files, mr_train_files, mr_val_files, ct_test_files, mr_test_files
  1923. #monai.config.print_config()
  1924. # old_stdout = sys.stdout
  1925. # log_file = open("mri_classification/logs/5fold_eval_fakemr.log","w")
  1926. # sys.stdout = log_file
  1927. print("this will be written to message.log")
  1928. # file_handler = logging.FileHandler(filename='logs/tmp.log')
  1929. # #stdout_handler = logging.StreamHandler(stream=sys.stdout)
  1930. # handlers = [file_handler]
  1931. # logging.basicConfig(
  1932. # level=logging.DEBUG,
  1933. # format='[%(asctime)s] {%(filename)s:%(lineno)d} %(levelname)s - %(message)s',
  1934. # handlers=handlers
  1935. # )
  1936. # logger = logging.getLogger('LOGGER_NAME')
  1937. # edit the path accordingly
  1938. gose = pd.read_csv('mri_classification/Pilot_169_train_val_test_split.csv')
  1939. X = []
  1940. Y = []
  1941. test_subjects = []
  1942. test_labels = []
  1943. torch.cuda.empty_cache()
  1944. for i in range(gose.shape[0]):
  1945. subj = gose['Main.GUID'][i]
  1946. subj_id = str(subj)[4:]
  1947. name = str(subj_id) + '.nii'
  1948. if gose['Set'][i]=='Train':
  1949. X.append(name)
  1950. Y.append(gose['Class'][i])
  1951. elif gose['Set'][i]=='Val':
  1952. X.append(name)
  1953. Y.append(gose['Class'][i])
  1954. elif gose['Set'][i]=='Test':
  1955. test_subjects.append(name)
  1956. test_labels.append(gose['Class'][i])
  1957. else:
  1958. print("Unknown Set: ", gose['Set'][i])
  1959. print("Total Subjects for 5-Fold (Train+Val):", len(X))
  1960. print("Test Subjects (kept separate):", len(test_subjects))
  1961. print("5-Fold Subjects Information:\n")
  1962. print("Subjects of Class 0: ", Y.count(0))
  1963. print("Subjects of Class 1: ", Y.count(1))
  1964. print("Subjects of Class 2: ", Y.count(2))
  1965. skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=1)
  1966. fold_number = 1
  1967. X = np.array(X)
  1968. Y = np.array(Y)
  1969. train_acc, train_auc = [], []
  1970. val_acc, val_auc = [], []
  1971. test_acc, test_auc = [], []
  1972. y_pred_test = []
  1973. y_pred_auc = []
  1974. for train_index, test_index in skf.split(X, Y):
  1975. print("Starting Fold {}..!".format(fold_number))
  1976. x_train_fold, x_test_fold = X[train_index], X[test_index]
  1977. y_train_fold, y_test_fold = Y[train_index], Y[test_index]
  1978. y_train_fold = list(y_train_fold)
  1979. y_test_fold = list(y_test_fold)
  1980. x_train_fold = list(x_train_fold)
  1981. x_test_fold = list(x_test_fold)
  1982. # print("\n")
  1983. # print("Fold {} statistics:\n".format(fold_number))
  1984. # print("Train Subjects: {}".format(len(y_train_fold)))
  1985. # print("Subjects of Class 0: ", y_train_fold.count(0))
  1986. # print("Subjects of Class 1: ", y_train_fold.count(1))
  1987. # print("Subjects of Class 2: {}\n".format(y_train_fold.count(2)))
  1988. # print("Val Subjects: {}".format(len(y_test_fold)))
  1989. # print("Subjects of Class 0: ", y_test_fold.count(0))
  1990. # print("Subjects of Class 1: ", y_test_fold.count(1))
  1991. # print("Subjects of Class 2: ", y_test_fold.count(2))
  1992. ct_train_files, ct_val_files, mr_train_files, mr_val_files, ct_test_files, mr_test_files = load_data(x_train_fold, x_test_fold, y_train_fold, y_test_fold, test_subjects, test_labels)
  1993. # Define transforms for CT and MR respectively
  1994. ct_transforms = Compose(
  1995. [
  1996. LoadImaged(keys=["img"], ensure_channel_first=True),
  1997. #ScaleIntensityRanged(keys=["img"],a_min=0, a_max=85, b_min=0, b_max=1),
  1998. NormalizeIntensityd(keys=["img"]),
  1999. #Resized(keys=["img"], spatial_size=(96, 96, 96)),
  2000. ]
  2001. )
  2002. mr_transforms = Compose(
  2003. [
  2004. LoadImaged(keys=["img"], ensure_channel_first=True),
  2005. NormalizeIntensityd(keys=["img"]),
  2006. RandFlipd(keys=["img"], prob=1, spatial_axis=2),
  2007. RandRotated(keys=["img"], prob=1, range_x=[0.4,0.4])
  2008. #ScaleIntensityRanged(keys=["img"],a_min=0, a_max=85, b_min=0, b_max=1),
  2009. ]
  2010. )
  2011. post_pred = Compose([Activations(softmax=True)])
  2012. post_label = Compose([AsDiscrete(to_onehot=3)])
  2013. # create a training data loader
  2014. ct_train_ds = monai.data.Dataset(data=ct_train_files, transform=ct_transforms)
  2015. mr_train_ds = monai.data.Dataset(data=mr_train_files, transform=mr_transforms)
  2016. combined_train_ds = CustomDataset(d1 = mr_train_ds, d2 = ct_train_ds)
  2017. train_loader = DataLoader(combined_train_ds, batch_size=1, shuffle=False, pin_memory=torch.cuda.is_available())
  2018. check_data1, check_data2 = monai.utils.misc.first(train_loader)
  2019. #print(check_data1["img"].shape, check_data1["label"], check_data2["img"].shape, check_data2["label"])
  2020. # create a validation data loader
  2021. ct_val_ds = monai.data.Dataset(data=ct_val_files, transform=ct_transforms)
  2022. mr_val_ds = monai.data.Dataset(data=mr_val_files, transform=mr_transforms)
  2023. combined_val_ds = CustomDataset(d1 = mr_val_ds, d2 = ct_val_ds)
  2024. val_loader = DataLoader(combined_val_ds, batch_size=1, shuffle=False, pin_memory=torch.cuda)
  2025. # create a test data loader
  2026. ct_test_ds = monai.data.Dataset(data=ct_test_files, transform=ct_transforms)
  2027. mr_test_ds = monai.data.Dataset(data=mr_test_files, transform=mr_transforms)
  2028. combined_test_ds = CustomDataset(d1 = mr_test_ds, d2 = ct_test_ds)
  2029. test_loader = DataLoader(combined_test_ds, batch_size=1, shuffle=False, pin_memory=torch.cuda)
  2030. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  2031. #print(device)
  2032. model_ct = resnet18(pretrained=False, spatial_dims=3, n_input_channels=1, num_classes=3).to(device)
  2033. #model_ct.load_state_dict(torch.load("models/best_val_resnet18.pth"))
  2034. model_mr = resnet18(pretrained=False, spatial_dims=3, n_input_channels=1, num_classes=3).to(device)
  2035. #model_mr.load_state_dict(torch.load("models/best_mr_val_resnet18.pth"))
  2036. # Freeze these models
  2037. for param in model_mr.parameters():
  2038. param.requires_grad_(True)
  2039. for param in model_ct.parameters():
  2040. param.requires_grad_(True)
  2041. # Create ensemble model
  2042. model = MyEnsemble(model_ct, model_mr)
  2043. #model.load_state_dict(torch.load("mri_classification/models/5-Fold/Final_Fold{}_val_resnet18.pth".format(int(fold_number))))
  2044. #model.load_state_dict(torch.load("mri_classification/models_gen/5-Fold/Fake_MR_Fold{}_ensemble.pth".format(int(fold_number))))
  2045. model.load_state_dict(torch.load("/data/amciilab/ajoshi83/models/MR_Fold{}.pth".format(fold_number)))
  2046. #print("Ensemble Model arch check:")
  2047. #print(model)
  2048. loss_function = torch.nn.CrossEntropyLoss()
  2049. optimizer = torch.optim.Adam(model.parameters(),lr= 1e-5, weight_decay=1e-3)
  2050. auc_metric = ROCAUCMetric(average="weighted")
  2051. # starting evaluation
  2052. val_interval = 1
  2053. best_metric = -1
  2054. best_metric_epoch = -1
  2055. best_val_loss = 2
  2056. writer = SummaryWriter()
  2057. with torch.no_grad():
  2058. num_correct = 0.0
  2059. metric_count = 0
  2060. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  2061. y = torch.tensor([], dtype=torch.long, device=device)
  2062. saver = CSVSaver(output_dir="mri_classification/train_output")
  2063. for batch_data in train_loader:
  2064. #step += 1
  2065. #print(batch_data["img"])
  2066. ct_batch, mr_batch = batch_data[0], batch_data[1]
  2067. #print(ct_batch)
  2068. #print(mr_batch)
  2069. ct_inputs, ct_labels = ct_batch["img"].to(device), ct_batch["label"].to(device)
  2070. mr_inputs, mr_labels = mr_batch["img"].to(device), mr_batch["label"].to(device)
  2071. train_outputs = model(ct_inputs, mr_inputs).argmax(dim=1)
  2072. y_pred = torch.cat([y_pred, model(ct_inputs, mr_inputs)], dim=0)
  2073. y = torch.cat([y, mr_labels], dim=0)
  2074. value = torch.eq(train_outputs, mr_labels)
  2075. metric_count += len(value)
  2076. num_correct += value.sum().item()
  2077. #saver.save_batch(train_outputs, train_data["img"].meta)
  2078. metric = num_correct / metric_count
  2079. # print("val evaluation metric:", metric)
  2080. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  2081. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred, detach = True)]
  2082. y_onehot = torch.stack(y_onehot, dim=0)
  2083. y_pred_act = torch.stack(y_pred_act, dim=0)
  2084. y_onehot = y_onehot.to(device="cpu")
  2085. y_pred_act = y_pred_act.to(device="cpu")
  2086. auc_metric(y_pred_act, y_onehot)
  2087. auc_result = auc_metric.aggregate()
  2088. auc_metric.reset()
  2089. saver.finalize()
  2090. train_acc.append(round(metric,3))
  2091. train_auc.append(round(auc_result,3))
  2092. print("Train Set: Accuracy: {:.4f}, AUC: {:.4f}".format(metric, auc_result))
  2093. with torch.no_grad():
  2094. num_correct = 0.0
  2095. metric_count = 0
  2096. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  2097. y = torch.tensor([], dtype=torch.long, device=device)
  2098. saver = CSVSaver(output_dir="mri_classification/val_output")
  2099. for val_data in val_loader:
  2100. ct_val_data, mr_val_data = val_data[0], val_data[1]
  2101. ct_val_images, ct_val_labels = ct_val_data["img"].to(device), ct_val_data["label"].to(device)
  2102. mr_val_images, mr_val_labels = mr_val_data["img"].to(device), mr_val_data["label"].to(device)
  2103. y_pred = torch.cat([y_pred, model(ct_val_images, mr_val_images)], dim=0)
  2104. #print(y_pred)
  2105. y = torch.cat([y, mr_val_labels], dim=0)
  2106. acc_value = torch.eq(y_pred.argmax(dim=1), y)
  2107. acc_metric = acc_value.sum().item() / len(acc_value)
  2108. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  2109. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred,detach=True)]
  2110. y_onehot = torch.stack(y_onehot, dim=0)
  2111. y_pred_act = torch.stack(y_pred_act, dim=0)
  2112. y_onehot = y_onehot.to(device="cpu")
  2113. y_pred_act = y_pred_act.to(device="cpu")
  2114. auc_metric(y_pred_act, y_onehot)
  2115. auc_result = auc_metric.aggregate()
  2116. auc_metric.reset()
  2117. saver.finalize()
  2118. val_acc.append(round(acc_metric,3))
  2119. val_auc.append(round(auc_result,3))
  2120. print("Val Set: Accuracy: {:.4f}, AUC: {:.4f}".format(acc_metric, auc_result))
  2121. with torch.no_grad():
  2122. num_correct = 0.0
  2123. metric_count = 0
  2124. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  2125. y = torch.tensor([], dtype=torch.long, device=device)
  2126. saver = CSVSaver(output_dir="mri_classification/test_output")
  2127. y_pred_mid = []
  2128. for test_data in test_loader:
  2129. ct_test_data, mr_test_data = test_data[0], test_data[1]
  2130. ct_test_images, ct_test_labels = ct_test_data["img"].to(device), ct_test_data["label"].to(device)
  2131. mr_test_images, mr_test_labels = mr_test_data["img"].to(device), mr_test_data["label"].to(device)
  2132. test_outputs = model(ct_test_images, mr_test_images).argmax(dim=1)
  2133. y_pred_mid.append(test_outputs.cpu().numpy()[0])
  2134. #softmax_op = torch.nn.functional.softmax(model(ct_test_images, mr_test_images)).argmax(dim=1)
  2135. #y_pred_mid.append(softmax_op.cpu().numpy()[0])
  2136. y_pred = torch.cat([y_pred, model(ct_test_images, mr_test_images)], dim=0)
  2137. y = torch.cat([y, mr_test_labels], dim=0)
  2138. value = torch.eq(test_outputs, mr_test_labels)
  2139. metric_count += len(value)
  2140. num_correct += value.sum().item()
  2141. #saver.save_batch(test_outputs, test_data["img"].meta)
  2142. y_pred_test.append(y_pred_mid)
  2143. metric = num_correct / metric_count
  2144. # print("test evaluation metric:", metric)
  2145. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  2146. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred, detach=True)]
  2147. y_onehot = torch.stack(y_onehot, dim=0)
  2148. y_pred_act = torch.stack(y_pred_act, dim=0)
  2149. y_onehot = y_onehot.to(device="cpu")
  2150. y_pred_act = y_pred_act.to(device="cpu")
  2151. y_pred_auc.append(y_pred_act)
  2152. auc_metric(y_pred_act, y_onehot)
  2153. auc_result = auc_metric.aggregate()
  2154. auc_metric.reset()
  2155. # auc_score =roc_auc_score(y_onehot, y_pred_act, multi_class='ovr', average='weighted')
  2156. # print(auc_score)
  2157. saver.finalize()
  2158. test_acc.append(round(metric,3))
  2159. test_auc.append(round(auc_result,3))
  2160. print("Test Set: Accuracy: {:.4f}, AUC: {:.4f}".format(metric, auc_result))
  2161. print("Fold {} completed...Next Fold starting...".format(fold_number))
  2162. print("\n")
  2163. fold_number += 1
  2164. writer.close()
  2165. print("\n")
  2166. print("Train Acc: {:.4f} +/- {:.4f}, Train AUC: {:.4f} +/- {:.4f}".format(mean(train_acc), np.std(train_acc), mean(train_auc), np.std(train_auc)))
  2167. print("Val Acc: {:.4f} +/- {:.4f}, Val AUC: {:.4f} +/- {:.4f}".format(mean(val_acc), np.std(val_acc), mean(val_auc), np.std(val_auc)))
  2168. # # print("Val Accuracies of 5 Folds:", val_acc)
  2169. # print("Val AUCs of 5 Folds:", val_auc)
  2170. print("Test Acc: {:.4f} +/- {:.4f}, Test AUC: {:.4f} +/- {:.4f}".format(mean(test_acc), np.std(test_acc), mean(test_auc), np.std(test_auc)))
  2171. # print("Test Accuracies of 5 Folds:", test_acc)
  2172. # print("Test AUCs of 5 Folds:", test_auc)
  2173. y_pred_tr = np.transpose(y_pred_test)
  2174. final = []
  2175. for i in range(y_pred_tr.shape[0]):
  2176. final.append(mode(y_pred_tr[i]))
  2177. y_true = y.cpu().numpy()
  2178. final_np = np.array(final)
  2179. y_pred_auc = np.array(y_pred_auc)
  2180. sum = np.sum(y_pred_auc, axis=0)
  2181. sum = sum/5
  2182. target_names = ['class 0', 'class 1', 'class 2']
  2183. print("\n")
  2184. print("Test Accuracy (Hard Voting of 5 models):", round(accuracy_score(y_true, final_np), 4))
  2185. print("Test AUC (Hard Voting of 5 models):", round(roc_auc_score(y_onehot, sum, multi_class='ovr', average='weighted'), 4))
  2186. cm = confusion_matrix(y_true, final_np)
  2187. print(classification_report(y_true, final_np, target_names=target_names))
  2188. disp = ConfusionMatrixDisplay(confusion_matrix=cm)
  2189. disp.plot()
  2190. plt.show()
  2191. # %%
  2192. # %%
  2193. # %%
  2194. import logging
  2195. import os
  2196. import sys
  2197. from pathlib import Path
  2198. import pandas as pd
  2199. import numpy as np
  2200. import torch
  2201. from torch.utils.tensorboard import SummaryWriter
  2202. import monai
  2203. from monai.data import decollate_batch, CSVSaver, DataLoader
  2204. from monai.metrics import ROCAUCMetric
  2205. from monai.transforms import ScaleIntensityRanged, NormalizeIntensityd, Activations, AsDiscrete, Compose, LoadImaged, RandRotate90d, Resized, ScaleIntensityd
  2206. from monai.networks.nets import resnet10, resnet18, resnet34, resnet50
  2207. import torch
  2208. import torch.nn as nn
  2209. from torch.utils.data import DataLoader as _TorchDataLoader
  2210. from torch.utils.data import Dataset
  2211. from monai.data.utils import list_data_collate, set_rnd, worker_init_fn
  2212. from sklearn.model_selection import StratifiedKFold
  2213. from statistics import mean, mode
  2214. from sklearn.metrics import accuracy_score
  2215. import matplotlib.pyplot as plt
  2216. from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay, classification_report, accuracy_score, roc_auc_score
  2217. class CustomDataset(monai.data.Dataset):
  2218. def __init__(self, d1, d2):
  2219. self.d1 = d1
  2220. self.d2 = d2
  2221. def __getitem__(self, idx):
  2222. dict1 = self.d1.__getitem__(idx)
  2223. image1, label1 = dict1["img"], dict1["label"]
  2224. dict2 = self.d2.__getitem__(idx)
  2225. image2, label2 = dict2["img"], dict2["label"]
  2226. assert label1==label2
  2227. dict_1 = dict()
  2228. dict_1["img"] = image1
  2229. dict_1["label"] = label1
  2230. dict_2 = dict()
  2231. dict_2["img"] = image2
  2232. dict_2["label"] = label2
  2233. return dict_1, dict_2
  2234. def __len__(self):
  2235. return len(self.d1)
  2236. class MyEnsemble(nn.Module):
  2237. def __init__(self, modelA, modelB):
  2238. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  2239. super(MyEnsemble, self).__init__()
  2240. self.modelA = modelA
  2241. self.modelB = modelB
  2242. # Remove last linear layer
  2243. self.modelA.fc = nn.Identity()
  2244. self.modelB.fc = nn.Identity()
  2245. # Create new classifier
  2246. self.mlp1 = nn.Linear(1024,256).to(device)
  2247. self.mlp2 = nn.Linear(256,32).to(device)
  2248. self.classifier = nn.Linear(32,3).to(device)
  2249. def forward(self, i1, i2):
  2250. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  2251. x1 = self.modelA(i1) # clone to make sure x is not changed by inplace methods
  2252. x1 = x1.view(x1.size(0), -1).to(device)
  2253. x2 = self.modelB(i2)
  2254. x2 = x2.view(x2.size(0), -1).to(device)
  2255. x = torch.cat((x1, x2), dim=1).to(device)
  2256. x = nn.functional.relu(self.mlp1(x)).to(device)
  2257. x = nn.functional.relu(self.mlp2(x)).to(device)
  2258. x = self.classifier(x).to(device)
  2259. return x
  2260. def load_data(train_images, val_images, train_labels, val_labels, test_images, test_labels):
  2261. train_labels_array = np.array(train_labels, dtype=np.int64)
  2262. val_labels_array = np.array(val_labels, dtype=np.int64)
  2263. test_labels_array = np.array(test_labels, dtype=np.int64)
  2264. #loading the MR images
  2265. mr_path = Path('Data/Reg_MR')
  2266. mr_train_images_path = [os.path.join(mr_path, f) for f in train_images]
  2267. mr_val_images_path = [os.path.join(mr_path, f) for f in val_images]
  2268. mr_test_images_path = [os.path.join(mr_path, f) for f in test_images]
  2269. mr_train_files = [{"img": img, "label": label} for img, label in zip(mr_train_images_path, train_labels_array)]
  2270. mr_val_files = [{"img": img, "label": label} for img, label in zip(mr_val_images_path, val_labels_array)]
  2271. mr_test_files = [{"img": img, "label": label} for img, label in zip(mr_test_images_path, test_labels_array)]
  2272. return mr_train_files, mr_val_files, mr_test_files
  2273. #monai.config.print_config()
  2274. # old_stdout = sys.stdout
  2275. # log_file = open("mri_classification/logs/MR_5Fold_eval.log","w")
  2276. # sys.stdout = log_file
  2277. # edit the path accordingly
  2278. gose = pd.read_csv('mri_classification/Pilot_169_train_val_test_split.csv')
  2279. X = []
  2280. Y = []
  2281. test_subjects = []
  2282. test_labels = []
  2283. torch.cuda.empty_cache()
  2284. for i in range(gose.shape[0]):
  2285. subj = gose['Main.GUID'][i]
  2286. subj_id = str(subj)[4:]
  2287. name = str(subj_id) + '.nii'
  2288. if gose['Set'][i]=='Train':
  2289. X.append(name)
  2290. Y.append(gose['Class'][i])
  2291. elif gose['Set'][i]=='Val':
  2292. X.append(name)
  2293. Y.append(gose['Class'][i])
  2294. elif gose['Set'][i]=='Test':
  2295. test_subjects.append(name)
  2296. test_labels.append(gose['Class'][i])
  2297. else:
  2298. print("Unknown Set: ", gose['Set'][i])
  2299. print("Total Subjects for 5-Fold (Train+Val):", len(X))
  2300. print("Test Subjects (kept separate):", len(test_subjects))
  2301. print("5-Fold Subjects Information:\n")
  2302. print("Subjects of Class 0: ", Y.count(0))
  2303. print("Subjects of Class 1: ", Y.count(1))
  2304. print("Subjects of Class 2: ", Y.count(2))
  2305. skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=1)
  2306. fold_number = 1
  2307. X = np.array(X)
  2308. Y = np.array(Y)
  2309. train_acc, train_auc = [], []
  2310. val_acc, val_auc = [], []
  2311. test_acc, test_auc = [], []
  2312. y_pred_test = []
  2313. y_pred_auc = []
  2314. for train_index, val_index in skf.split(X, Y):
  2315. print("Starting Fold {}..!".format(fold_number))
  2316. # if fold_number <5:
  2317. # fold_number += 1
  2318. # continue
  2319. x_train_fold, x_val_fold = X[train_index], X[val_index]
  2320. y_train_fold, y_val_fold = Y[train_index], Y[val_index]
  2321. y_train_fold = list(y_train_fold)
  2322. y_val_fold = list(y_val_fold)
  2323. x_train_fold = list(x_train_fold)
  2324. x_val_fold = list(x_val_fold)
  2325. # test_subjects = list(test_subjects)
  2326. # test_labels = list(test_labels)
  2327. print("\n")
  2328. print("Fold {} statistics:\n".format(fold_number))
  2329. print("Train Subjects: {}".format(len(y_train_fold)))
  2330. print("Subjects of Class 0: ", y_train_fold.count(0))
  2331. print("Subjects of Class 1: ", y_train_fold.count(1))
  2332. print("Subjects of Class 2: {}\n".format(y_train_fold.count(2)))
  2333. print("val Subjects: {}".format(len(y_val_fold)))
  2334. print("Subjects of Class 0: ", y_val_fold.count(0))
  2335. print("Subjects of Class 1: ", y_val_fold.count(1))
  2336. print("Subjects of Class 2: ", y_val_fold.count(2))
  2337. mr_train_files, mr_val_files, mr_test_files = load_data(x_train_fold, x_val_fold, y_train_fold, y_val_fold, test_subjects, test_labels)
  2338. mr_transforms = Compose(
  2339. [
  2340. LoadImaged(keys=["img"], ensure_channel_first=True),
  2341. NormalizeIntensityd(keys=["img"]),
  2342. ]
  2343. )
  2344. post_pred = Compose([Activations(softmax=True)])
  2345. post_label = Compose([AsDiscrete(to_onehot=3)])
  2346. # create a training data loader
  2347. mr_train_ds = monai.data.Dataset(data=mr_train_files, transform=mr_transforms)
  2348. mr_train_loader = DataLoader(mr_train_ds, batch_size=1, shuffle=False, num_workers=4, pin_memory=torch.cuda.is_available())
  2349. check_data1 = monai.utils.misc.first(mr_train_loader)
  2350. print(check_data1["img"].shape, check_data1["label"])
  2351. # create a validation data loader
  2352. mr_val_ds = monai.data.Dataset(data=mr_val_files, transform=mr_transforms)
  2353. mr_val_loader = DataLoader(mr_val_ds, batch_size=1, shuffle=False, num_workers=4, pin_memory=torch.cuda.is_available())
  2354. # create a test data loader
  2355. mr_test_ds = monai.data.Dataset(data= mr_test_files, transform=mr_transforms)
  2356. mr_test_loader = DataLoader(mr_test_ds, batch_size=1, shuffle=False, num_workers=4, pin_memory=torch.cuda.is_available())
  2357. # Create Model, CrossEntropy Loss and Adam optimizer
  2358. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  2359. print(device)
  2360. model = resnet18(pretrained=False, spatial_dims=3, n_input_channels=1, num_classes=3).to(device)
  2361. #model.load_state_dict(torch.load("mri_classification/models/5-Fold/MR_Fold{}.pth".format(int(fold_number))))
  2362. model.load_state_dict(torch.load("/data/amciilab/ajoshi83/models/MR_Fold{}.pth".format(fold_number)))
  2363. loss_function = torch.nn.CrossEntropyLoss()
  2364. optimizer = torch.optim.Adam(model.parameters(),lr= 1e-5, weight_decay=1e-3)
  2365. auc_metric = ROCAUCMetric(average="weighted")
  2366. # starting evaluation
  2367. val_interval = 1
  2368. best_metric = -1
  2369. best_metric_epoch = -1
  2370. best_val_loss = 2
  2371. writer = SummaryWriter()
  2372. # with torch.no_grad():
  2373. # num_correct = 0.0
  2374. # metric_count = 0
  2375. # y_pred = torch.tensor([], dtype=torch.float32, device=device)
  2376. # y = torch.tensor([], dtype=torch.long, device=device)
  2377. # saver = CSVSaver(output_dir="./train_output")
  2378. # for batch_data in mr_train_loader:
  2379. # inputs, labels = batch_data["img"].to(device), batch_data["label"].to(device)
  2380. # train_outputs = model(inputs).argmax(dim=1)
  2381. # y_pred = torch.cat([y_pred, model(inputs)], dim=0)
  2382. # y = torch.cat([y, labels], dim=0)
  2383. # value = torch.eq(train_outputs, labels)
  2384. # metric_count += len(value)
  2385. # num_correct += value.sum().item()
  2386. # #saver.save_batch(train_outputs, train_data["img"].meta)
  2387. # metric = num_correct / metric_count
  2388. # # print("val evaluation metric:", metric)
  2389. # y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  2390. # y_pred_act = [post_pred(i) for i in decollate_batch(y_pred, detach = True)]
  2391. # y_onehot = torch.stack(y_onehot, dim=0)
  2392. # y_pred_act = torch.stack(y_pred_act, dim=0)
  2393. # y_onehot = y_onehot.to(device="cpu")
  2394. # y_pred_act = y_pred_act.to(device="cpu")
  2395. # auc_metric(y_pred_act, y_onehot)
  2396. # auc_result = auc_metric.aggregate()
  2397. # auc_metric.reset()
  2398. # saver.finalize()
  2399. # train_acc.append(round(metric,3))
  2400. # train_auc.append(round(auc_result,3))
  2401. # print("Train Set: Accuracy: {:.4f}, AUC: {:.4f}".format(metric, auc_result))
  2402. # with torch.no_grad():
  2403. # num_correct = 0.0
  2404. # metric_count = 0
  2405. # y_pred = torch.tensor([], dtype=torch.float32, device=device)
  2406. # y = torch.tensor([], dtype=torch.long, device=device)
  2407. # saver = CSVSaver(output_dir="./val_output")
  2408. # for val_data in mr_val_loader:
  2409. # val_images, val_labels = val_data["img"].to(device), val_data["label"].to(device)
  2410. # y_pred = torch.cat([y_pred, model(val_images)], dim=0)
  2411. # #print(y_pred)
  2412. # y = torch.cat([y, val_labels], dim=0)
  2413. # acc_value = torch.eq(y_pred.argmax(dim=1), y)
  2414. # acc_metric = acc_value.sum().item() / len(acc_value)
  2415. # y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  2416. # y_pred_act = [post_pred(i) for i in decollate_batch(y_pred,detach=True)]
  2417. # y_onehot = torch.stack(y_onehot, dim=0)
  2418. # y_pred_act = torch.stack(y_pred_act, dim=0)
  2419. # y_onehot = y_onehot.to(device="cpu")
  2420. # y_pred_act = y_pred_act.to(device="cpu")
  2421. # auc_metric(y_pred_act, y_onehot)
  2422. # auc_result = auc_metric.aggregate()
  2423. # auc_metric.reset()
  2424. # saver.finalize()
  2425. # val_acc.append(round(acc_metric,3))
  2426. # val_auc.append(round(auc_result,3))
  2427. # print("Val Set: Accuracy: {:.4f}, AUC: {:.4f}".format(acc_metric, auc_result))
  2428. # with torch.no_grad():
  2429. # num_correct = 0.0
  2430. # metric_count = 0
  2431. # y_pred = torch.tensor([], dtype=torch.float32, device=device)
  2432. # y = torch.tensor([], dtype=torch.long, device=device)
  2433. # #saver = CSVSaver(output_dir="./test_output")
  2434. # for test_data in mr_test_loader:
  2435. # mr_test_images, mr_test_labels = test_data["img"].to(device), test_data["label"].to(device)
  2436. # test_outputs = model(mr_test_images).argmax(dim=1)
  2437. # y_pred = torch.cat([y_pred, model(mr_test_images)], dim=0)
  2438. # y = torch.cat([y, mr_test_labels], dim=0)
  2439. # value = torch.eq(test_outputs, mr_test_labels)
  2440. # metric_count += len(value)
  2441. # num_correct += value.sum().item()
  2442. # #saver.save_batch(test_outputs, test_data["img"].meta)
  2443. # metric = num_correct / metric_count
  2444. # # print("test evaluation metric:", metric)
  2445. # y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  2446. # y_pred_act = [post_pred(i) for i in decollate_batch(y_pred, detach=True)]
  2447. # y_onehot = torch.stack(y_onehot, dim=0)
  2448. # y_pred_act = torch.stack(y_pred_act, dim=0)
  2449. # y_onehot = y_onehot.to(device="cpu")
  2450. # y_pred_act = y_pred_act.to(device="cpu")
  2451. # auc_metric(y_pred_act, y_onehot)
  2452. # auc_result = auc_metric.aggregate()
  2453. # auc_metric.reset()
  2454. # saver.finalize()
  2455. # test_acc.append(round(metric,3))
  2456. # test_auc.append(round(auc_result,3))
  2457. # print("Test Set: Accuracy: {:.4f}, AUC: {:.4f}".format(metric, auc_result))
  2458. # print("Fold {} completed...Next Fold starting...".format(fold_number))
  2459. # fold_number += 1
  2460. # writer.close()
  2461. with torch.no_grad():
  2462. num_correct = 0.0
  2463. metric_count = 0
  2464. y_pred = torch.tensor([], dtype=torch.float32, device=device)
  2465. y = torch.tensor([], dtype=torch.long, device=device)
  2466. saver = CSVSaver(output_dir="mri_classification/test_output")
  2467. y_pred_mid = []
  2468. for test_data in mr_test_loader:
  2469. mr_test_images, mr_test_labels = test_data["img"].to(device), test_data["label"].to(device)
  2470. test_outputs = model(mr_test_images).argmax(dim=1)
  2471. #y_pred_mid.append(torch.nn.functional.softmax(model(mr_test_images), dim=1).cpu().numpy()[0])
  2472. #y_pred_mid.append(test_outputs.cpu().numpy()[0])
  2473. softmax_op = torch.nn.functional.softmax(model(mr_test_images)).argmax(dim=1)
  2474. y_pred_mid.append(softmax_op.cpu().numpy()[0])
  2475. y_pred = torch.cat([y_pred, model(mr_test_images)], dim=0)
  2476. y = torch.cat([y, mr_test_labels], dim=0)
  2477. value = torch.eq(test_outputs, mr_test_labels)
  2478. metric_count += len(value)
  2479. num_correct += value.sum().item()
  2480. #saver.save_batch(test_outputs, test_data["img"].meta)
  2481. y_pred_test.append(y_pred_mid)
  2482. # print(num_correct)
  2483. # print(metric_count)
  2484. metric = num_correct / metric_count
  2485. # print("test evaluation metric:", metric)
  2486. y_onehot = [post_label(i) for i in decollate_batch(y, detach=True)]
  2487. y_pred_act = [post_pred(i) for i in decollate_batch(y_pred, detach=True)]
  2488. y_onehot = torch.stack(y_onehot, dim=0)
  2489. y_pred_act = torch.stack(y_pred_act, dim=0)
  2490. y_onehot = y_onehot.to(device="cpu")
  2491. y_pred_act = y_pred_act.to(device="cpu")
  2492. auc_metric(y_pred_act, y_onehot)
  2493. auc_result = auc_metric.aggregate()
  2494. auc_metric.reset()
  2495. saver.finalize()
  2496. test_acc.append(round(metric,3))
  2497. test_auc.append(round(auc_result,3))
  2498. print("Test Set: Accuracy: {:.4f}, AUC: {:.4f}".format(metric, auc_result))
  2499. print("Fold {} completed...Next Fold starting...".format(fold_number))
  2500. fold_number += 1
  2501. #writer.close()
  2502. print("\n")
  2503. # print("Train Acc: {:.4f} +/- {:.4f}, Train AUC: {:.4f} +/- {:.4f}".format(mean(train_acc), np.std(train_acc), mean(train_auc), np.std(train_auc)))
  2504. # print("Val Acc: {:.4f} +/- {:.4f}, Val AUC: {:.4f} +/- {:.4f}".format(mean(val_acc), np.std(val_acc), mean(val_auc), np.std(val_auc)))
  2505. # # # print("Val Accuracies of 5 Folds:", val_acc)
  2506. # print("Val AUCs of 5 Folds:", val_auc)
  2507. print("Test Acc: {:.4f} +/- {:.4f}, Test AUC: {:.4f} +/- {:.4f}".format(mean(test_acc), np.std(test_acc), mean(test_auc), np.std(test_auc)))
  2508. # print("Test Accuracies of 5 Folds:", test_acc)
  2509. # print("Test AUCs of 5 Folds:", test_auc)
  2510. y_pred_tr = np.transpose(y_pred_test)
  2511. final = []
  2512. for i in range(y_pred_tr.shape[0]):
  2513. final.append(mode(y_pred_tr[i]))
  2514. y_true = y.cpu().numpy()
  2515. final_np = np.array(final)
  2516. y_pred_auc = np.array(y_pred_auc)
  2517. sum = np.sum(y_pred_auc, axis=0)
  2518. sum = sum/5
  2519. target_names = ['class 0', 'class 1', 'class 2']
  2520. print("\n")
  2521. print("Test Accuracy (Hard Voting of 5 models):", round(accuracy_score(y_true, final_np), 4))
  2522. #print("Test AUC (Hard Voting of 5 models):", round(roc_auc_score(y_onehot, sum, multi_class='ovr', average='weighted'), 4))
  2523. cm = confusion_matrix(y_true, final_np)
  2524. print(classification_report(y_true, final_np, target_names=target_names))
  2525. disp = ConfusionMatrixDisplay(confusion_matrix=cm)
  2526. disp.plot()
  2527. plt.show()
  2528. # %%
  2529. torch.nn.functional.softmax(model(mr_test_images), dim=1).detach().cpu().numpy()[0]
  2530. # %%
  2531. y_pred_tr.shape
  2532. # %%
  2533. disp = ConfusionMatrixDisplay(confusion_matrix=cm)
  2534. disp.plot()
  2535. plt.show()
  2536. # %%
  2537. roc_auc_score(y_onehot, sum, multi_class='ovr', average='weighted')
  2538. # %%
  2539. import matplotlib.pyplot as plt
  2540. from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay, classification_report, accuracy_score
  2541. target_names = ['class 0', 'class 1', 'class 2']
  2542. print("Test Accuracy (Hard Voting of 5 models):", accuracy_score(y.cpu().numpy(), final_np))
  2543. cm = confusion_matrix(y.cpu().numpy(), final_np)
  2544. print(classification_report(y.cpu().numpy(), final_np, target_names=target_names))
  2545. disp = ConfusionMatrixDisplay(confusion_matrix=cm)
  2546. disp.plot()
  2547. plt.show()
  2548. # %%

analysis.ipynb at commit ae8ad46, no license · at the source

Overview

Authors: Yiming Che1,2, Amogh Manoj Joshi1,2, Jay Shah1,2, Md Mahfuzur Rahman Siddiquee1,2, Catherine D Chong2,3, Simona Nikolova3, Gina Dumkrieger3, Baoxin Li1,2, Teresa Wu1,2, Todd J Schwedt2,3
  1. School of Computing and Augmented Intelligence, Arizona State University, Tempe, AZ 85281, USA
  2. ASU-Mayo Center for Innovative Imaging, Tempe, AZ 85281, USA
  3. Department of Neurology, Mayo Clinic, Phoenix, AZ 85054, USA
Institutions: Arizona State University (United States); Mayo Clinic (United States); Mayo Clinic in Arizona (United States); Mayo Clinic Hospital (United States)
Journal: Brain communications, volume 8, issue 2, article fcag123
Dates: received 1 May 2025; accepted 2 April 2026; published online 6 April 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1093/braincomms/fcag123 · PMID 42004011 · PMCID PMC13084558 · OpenAlex W7152176423
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), other (modality), human (organism), traumatic brain injury (population)
Methods: Connectivity, Machine learning, Statistics
Keywords: deep learning, traumatic brain injury, concussion, neuroimaging, data harmonization
Topic: Traumatic Brain Injury Research (Epidemiology, Medicine), according to OpenAlex
Citations: cited by 1 paper (Europe PMC); 52 references in the paper

Abstract

To enhance the prediction of traumatic brain injury (mTBI) outcomes, we propose a deep learning approach that integrates brain computed tomography (CT) scans with corresponding synthetic T1-weighted magnetic resonance imaging (T1-MRI). Our method significantly outperforms the prediction using CT scans alone. TRACK-TBI Pilot dataset, which includes imaging and clinical outcome data from patients with TBI, is studied. The hypothesis is brain CT and T1-MRI complement each other and together will improve TBI prognosis compared to using either CT or T1-MRI alone. Since CT and T1-MRI may not be available for the same individual, we employed a specialized version of a generative adversarial network (GAN), known as fixed-point GAN (FP-GAN). FP-GAN was trained using unpaired CT and T1-MRI scans to generate synthetic T1-MRIs from real CT scans. This process produced pseudo-paired CT-MRI data, which was then used to train a deep learning classifier for outcome prediction. The classifier consists of dual parallel 3D ResNet-18 models, each independently processing T1-MRI and CT scans. We used Glasgow Outcome Scale-Extended (GOSE) scores at 3 months post-TBI as the measure of patient outcomes. To avoid data leakage, the subjects used in FP-GAN and ResNet-18 model have no overlap. We further divided the paired data, allocating 69 samples for 5-fold cross-validation and 17 samples for testing. Prognostic performance was evaluated using the area under the receiver operating characteristic curve (AUC), F1-score (the harmonic mean of precision and recall), sensitivity (true positive rate) and specificity (true negative rate). For binary classification, we defined good recovery as GOSE ≥ 7 (positive) and poor recovery (negative) as 3 ≤ GOSE ≤ 6. Accordingly, our training set consists of 24 subjects with poor recovery and 45 subjects with good recovery, while the testing set includes 5 subjects with poor recovery and 12 subjects with good recovery. A DeLong test on AUC confirms that the improvement from incorporating synthetic T1-MRI (AUC = 0.76 ± 0.10) is statistically significant (P < 0.05) compared to using CT alone (AUC = 0.68 ± 0.13). The significant improvement from using the combination of real CT and synthetic T1-MRI in sensitivity (SEN = 0.95 ± 0.07) and overall performance metrics, such as F1-score (F1 = 0.84 ± 0.03), suggests that the proposed approach provides a robust and effective prognostic approach compared to using CT alone (SEN = 0.83 ± 0.18 and F1 = 0.76 ± 0.07). This pilot research demonstrates the potential of a deep learning-based harmonization model to bridge the gap between CT and T1-MRI in TBI assessment. By integrating synthetic T1-MRI with CT, prediction performance is substantially enhanced.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repository

Its files are read in the Code ↔ Paper reader above, with 2 matches between paragraphs and lines of code.

SoloChe/TBI-Recovery-Prediction-Harmonization

License: none: the authors keep all their rights
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: ae8ad4680577f45f5178d327ad6ae18f6202d7ca, 8 March 2026
Languages: Python (30), Shell (4), Jupyter (3)
Size: 39 files, 37 scripts
Software Heritage: not archived
Found in: “Data availability”
Holds: README, 2 notebooks
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: PyTorch (32 files), NumPy (30 files), scikit-learn (29 files), MONAI (28 files), pandas (28 files), Matplotlib (4 files), Pillow (3 files), SciPy (2 files), scikit-image (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
38 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 37 scripts, each with its path and the digest of its content;
  • 2 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

No dataset and no data link were found in the paper.

Data availability

All data used in this study were obtained from the Federal Interagency Traumatic Brain Injury Research (FITBIR) Informatics System. Access to the FITBIR datasets requires proper authorization and adherence to their data use agreements. Researchers interested in accessing the data can apply through the FITBIR Data Access Request process. Data sharing is not applicable to this article as no new data were created or analysed in this study. Model code, training scripts and pre/postprocessing pipelines are available at: https://github.com/SoloChe/TBI-Recovery-Prediction-Harmonization.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 28 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 10 authors, 5 keywords, 1 funder, 42 references.

Cite

This paper

Che, Y., Joshi, A. M., Shah, J., Rahman Siddiquee, M. M., Chong, C. D., Nikolova, S., Dumkrieger, G., Li, B., Wu, T., & Schwedt, T. J. (2026). Traumatic brain injury recovery prediction by harmonizing real brain CT and synthetic brain MRI: a pilot study. Brain communications, 8(2), fcag123. https://doi.org/10.1093/braincomms/fcag123

BibTeX

@article{che2026traumatic,
author = {Che, Yiming and Joshi, Amogh Manoj and Shah, Jay and Rahman Siddiquee, Md Mahfuzur and Chong, Catherine D and Nikolova, Simona and Dumkrieger, Gina and Li, Baoxin and Wu, Teresa and Schwedt, Todd J},
title = {{Traumatic brain injury recovery prediction by harmonizing real brain CT and synthetic brain MRI: a pilot study}},
journal = {Brain communications},
year = {2026},
month = apr,
volume = {8},
number = {2},
pages = {fcag123},
publisher = {Oxford University Press},
issn = {2632-1297},
doi = {10.1093/braincomms/fcag123},
url = {https://doi.org/10.1093/braincomms/fcag123},
pmid = {42004011},
pmcid = {PMC13084558}
}

RIS

TY - JOUR
AU - Che, Yiming
AU - Joshi, Amogh Manoj
AU - Shah, Jay
AU - Rahman Siddiquee, Md Mahfuzur
AU - Chong, Catherine D
AU - Nikolova, Simona
AU - Dumkrieger, Gina
AU - Li, Baoxin
AU - Wu, Teresa
AU - Schwedt, Todd J
TI - Traumatic brain injury recovery prediction by harmonizing real brain CT and synthetic brain MRI: a pilot study
T2 - Brain communications
J2 - Brain Commun
PY - 2026
DA - 2026/04/06
VL - 8
IS - 2
SP - fcag123
SN - 2632-1297
PB - Oxford University Press
DO - 10.1093/braincomms/fcag123
UR - https://doi.org/10.1093/braincomms/fcag123
LA - en
ER -

CSL-JSON

{
"id": "10.1093/braincomms/fcag123",
"type": "article-journal",
"title": "Traumatic brain injury recovery prediction by harmonizing real brain CT and synthetic brain MRI: a pilot study",
"container-title": "Brain communications",
"author": [
{
"family": "Che",
"given": "Yiming"
},
{
"family": "Joshi",
"given": "Amogh Manoj"
},
{
"family": "Shah",
"given": "Jay"
},
{
"family": "Rahman Siddiquee",
"given": "Md Mahfuzur"
},
{
"family": "Chong",
"given": "Catherine D"
},
{
"family": "Nikolova",
"given": "Simona"
},
{
"family": "Dumkrieger",
"given": "Gina"
},
{
"family": "Li",
"given": "Baoxin"
},
{
"family": "Wu",
"given": "Teresa"
},
{
"family": "Schwedt",
"given": "Todd J"
}
],
"container-title-short": "Brain Commun",
"volume": "8",
"issue": "2",
"page": "fcag123",
"DOI": "10.1093/braincomms/fcag123",
"PMID": "42004011",
"PMCID": "PMC13084558",
"ISSN": "2632-1297",
"publisher": "Oxford University Press",
"URL": "https://doi.org/10.1093/braincomms/fcag123",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
6
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1371/journal.pcbi.1014555 [code]
Body surface potential driven personalisation of electrophysiological digital twins in hypertrophic cardiomyopathy.
Journal: PLoS computational biology
In common: MONAI, scikit-image, Pillow, 6 other tools, other, structural MRI / diffusion
[2] doi:10.1371/journal.pone.0348866 [code]
Using deep learning to identify inherited retinal diseases based on wide-field retinal imaging data.
Journal: PloS one
In common: MONAI, scikit-image, Pillow, 6 other tools, other
[3] doi:10.3389/frai.2026.1771088 [code]
Few-shot deployment of pretrained MRI transformers in brain imaging tasks.
Journal: Frontiers in artificial intelligence
In common: MONAI, scikit-image, Pillow, 6 other tools, structural MRI / diffusion
[4] doi:10.1016/j.patter.2026.101538 [code]
A multi-modal foundation model for brain disease diagnosis and medical imaging.
Journal: Patterns (New York, N.Y.)
In common: MONAI, scikit-image, Pillow, 6 other tools
[5] doi:10.1186/s13244-026-02296-3 [code]
A pre-trained foundation model framework for multiplanar MRI classification of extramural vascular invasion and mesorectal fascia invasion in rectal cancer.
Journal: Insights into imaging
In common: MONAI, PyTorch, scikit-learn, 4 other tools, structural MRI / diffusion, 1 reference
[6] doi:10.1186/s13244-026-02365-7 [code]
Super-resolution MRI and 2.5D deep learning for intratumoral-peritumoral radiomics in preoperative prediction of rectal cancer perineural invasion.
Journal: Insights into imaging
In common: MONAI, Pillow, PyTorch, 5 other tools, structural MRI / diffusion
[7] doi:10.3389/fnins.2026.1870124 [code]
An end-to-end pipeline for automated fetal brain segmentation and biometry from 3D SSFP MRI.
Journal: Frontiers in neuroscience
In common: MONAI, scikit-image, PyTorch, 5 other tools, structural MRI / diffusion
[8] doi:10.1111/joa.70203 [code]
Two-step workflow integrating automatic registration and manual refinement for the accurate alignment of serial histological sections in 3D reconstruction.
Journal: Journal of anatomy
In common: MONAI, scikit-image, Pillow, 5 other tools
[9] doi:10.1016/j.crmeth.2026.101473 [code]
AmygdalaGo-BOLT for boundary-aware segmentation of the human amygdala.
Journal: Cell reports methods
In common: MONAI, scikit-image, Pillow, 4 other tools, structural MRI / diffusion
[10] doi:10.1038/s41467-026-76011-7 [code]
Human cortex organizes dynamic co-fluctuations along the sensorimotor-association axis.
Journal: Nature communications
In common: MONAI, scikit-image, Pillow, 4 other tools

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.