OSCR

Deep Learning-Based Temporal Gait Analysis Using a Smartphone IMU in Older Adults with and Without Non-Specific Low Back Pain.

Code ↔ Paper

4 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 4 matches
  1. [1] § 2. Materials and Methods › 2.4. Model Training and Evaluation ↔ notebooks/03_RNN_eval.ipynb, lines 760–817 · score 0.57 · IC FO macro, training strategy, weighted, architecture, expansion, fold
  2. [2] § 2. Materials and Methods › 2.6. Statistical Analysis ↔ notebooks/03_RNN_eval.ipynb, lines 1314–1393 · score 0.55 · square error, gait event detection, absolute error, root
  3. [3] § 2. Materials and Methods › 2.4. Model Training and Evaluation ↔ notebooks/03_RNN-Copy1.ipynb, lines 117–242 · score 0.51 · PyTorch, GPU, epochs, log, loss, batch
  4. [4] § 2. Materials and Methods › 2.4. Model Training and Evaluation ↔ notebooks/03_RNN-Copy2.ipynb, lines 120–245 · score 0.51 · PyTorch, GPU, epochs, log, loss, batch

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 · 1,600 lines · 51 KB · other · 2 matches

  1. # %%
  2. import os
  3. from pathlib import Path
  4. %matplotlib inline
  5. %load_ext autoreload
  6. %autoreload 2
  7. import seaborn as sns
  8. import pickle
  9. import time
  10. # %%
  11. from datetime import datetime
  12. from glob import glob
  13. import lightning as pl
  14. import numpy as np
  15. import pandas as pd
  16. import wandb
  17. from gait_ml.data.datamodule import GaitDataModule
  18. from gait_ml.model.litmodel import LitSeq2Seq
  19. from lightning.pytorch.callbacks import LearningRateMonitor, ModelCheckpoint
  20. from lightning.pytorch.loggers import WandbLogger
  21. from sklearn.model_selection import StratifiedKFold, train_test_split
  22. from hydra import initialize, compose
  23. import hydra
  24. from omegaconf import DictConfig
  25. import torch
  26. from gait_ml import evaluate
  27. from torchmetrics.classification import (
  28. MulticlassPrecision,
  29. MulticlassRecall,
  30. BinaryPrecision,
  31. BinaryRecall,
  32. Accuracy,
  33. ConfusionMatrix,
  34. MulticlassConfusionMatrix,
  35. MulticlassF1Score,
  36. F1Score,
  37. )
  38. import matplotlib.pyplot as plt
  39. def calculate_gait_mae(
  40. ground_truth: np.ndarray, prediction: np.ndarray, tolerance: int = 20
  41. ):
  42. """
  43. Matches gait events and computes overall and per-class MAE in one pass.
  44. Args:
  45. ground_truth: Array of ground truth labels.
  46. prediction: Array of predicted labels.
  47. tolerance: Max distance to consider a match.
  48. Returns:
  49. A tuple containing:
  50. - overall_mae (float): The MAE across all matched events.
  51. - per_class_mae (dict): A dictionary mapping class labels to their MAE.
  52. """
  53. gt_indices = np.where(ground_truth > 0)[0]
  54. available_preds = list(np.where(prediction > 0)[0])
  55. # Store tuples of (label, absolute_error) for each successful match
  56. errors_by_class = []
  57. # Greedily match each ground truth event to the closest prediction
  58. for gt_idx in gt_indices:
  59. gt_label = ground_truth[gt_idx]
  60. best_dist = float("inf")
  61. best_match_idx = -1
  62. # Find the closest available prediction of the same class within tolerance
  63. for pred_idx in available_preds:
  64. if prediction[pred_idx] == gt_label:
  65. dist = abs(gt_idx - pred_idx)
  66. if dist <= tolerance and dist < best_dist:
  67. best_dist = dist
  68. best_match_idx = pred_idx
  69. # If a match is found, record its error and remove it from the pool
  70. if best_match_idx != -1:
  71. error = abs(gt_idx - best_match_idx)
  72. errors_by_class.append((gt_label, error))
  73. available_preds.remove(best_match_idx)
  74. if not errors_by_class:
  75. return np.nan, {}
  76. # --- Calculate Final Metrics ---
  77. # Overall MAE is the mean of all collected errors
  78. all_errors = [err for lbl, err in errors_by_class]
  79. overall_mae = np.mean(all_errors)
  80. # Per-class MAE is calculated by grouping errors by label
  81. unique_labels = sorted(np.unique([lbl for lbl, err in errors_by_class]))
  82. per_class_mae = {
  83. int(label): np.mean([err for lbl, err in errors_by_class if lbl == label])
  84. for label in unique_labels
  85. }
  86. return overall_mae, per_class_mae
  87. # @hydra.main(config_path="../configs", config_name="train_config", version_base="1.3")
  88. def eval(data_set: str, model_fpath: str, fold: int, return_preds_targets=None):
  89. with initialize(config_path="../configs", job_name="train", version_base="1.3"):
  90. config = compose(config_name="train_config")
  91. pl.seed_everything(config.general.random_state)
  92. all_files = glob(config.general.data_path, recursive=True)
  93. all_files = np.sort(all_files)
  94. print(f"Processing: {len(all_files)} samples")
  95. ids = [int(i.split("/")[-1].split("_")[0]) for i in all_files]
  96. group_df = pd.read_csv(config.general.group_file, index_col="ID")
  97. group_df.columns = ["group"]
  98. group_df = group_df[group_df.group.notna()]
  99. group_df.replace("h", 0, inplace=True)
  100. group_df.replace("p", 1, inplace=True)
  101. grouping = group_df.loc[ids].group.values
  102. skf = StratifiedKFold(
  103. n_splits=config.general.n_splits,
  104. shuffle=True,
  105. random_state=config.general.random_state,
  106. )
  107. # Outer loop for K-Fold cross-validation
  108. # This loop creates the primary TEST set for each fold.
  109. for cur_fold, (train_val_index, test_index) in enumerate(
  110. skf.split(np.arange(len(ids)).reshape(-1, 1), grouping)
  111. ):
  112. if cur_fold != fold:
  113. continue
  114. print("cur_fold", cur_fold)
  115. print("input fold", fold)
  116. print(
  117. f"=============== FOLD {cur_fold + 1}/{config.general.n_splits} ================"
  118. )
  119. print("current_test set:", test_index)
  120. # Split data into a temporary training+validation set and the final test set
  121. X_train_val, X_test = (
  122. np.arange(len(ids))[train_val_index],
  123. np.arange(len(ids))[test_index],
  124. )
  125. print(f"X_train_val:", X_train_val)
  126. print(f"train_val_index:", train_val_index)
  127. print(f"X_test:", X_test)
  128. print(f"test_index:", test_index)
  129. y_train_val, y_test = grouping[train_val_index], grouping[test_index]
  130. X_train, X_val, y_train, y_val = train_test_split(
  131. X_train_val,
  132. y_train_val,
  133. test_size=0.25,
  134. stratify=y_train_val,
  135. random_state=1,
  136. )
  137. model = LitSeq2Seq(
  138. input_dim=config.model.input_dim,
  139. output_dim=config.model.output_dim,
  140. hidden_dim=config.model.hidden_dim,
  141. num_layers=config.model.num_layers,
  142. dropout_prob=config.model.dropout_prob,
  143. learning_rate=config.model.learning_rate,
  144. teacher_forcing_ratio=config.model.teacher_forcing_ratio,
  145. )
  146. current_time = datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
  147. run_name = f"GRU-expandlabel{config.data.expand_labels}_{current_time}"
  148. checkpoint_callback = ModelCheckpoint(
  149. monitor=config.training.monitor_metric,
  150. mode=config.training.monitor_mode,
  151. save_top_k=config.training.save_top_k,
  152. dirpath=f"{config.general.project_name}/{run_name}/checkpoints/",
  153. filename="model-{epoch:02d}-{val_f1score:.2f}",
  154. )
  155. checkpoint_callback.best_model_path = model_fpath
  156. print(f"Best model: {checkpoint_callback.best_model_path}")
  157. device_ = (
  158. torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
  159. )
  160. best_model = LitSeq2Seq.load_from_checkpoint(
  161. checkpoint_callback.best_model_path
  162. ).to(device_)
  163. testset_f1_scores = []
  164. testset_pc_f1_scores = []
  165. testset_pc_precision = []
  166. testset_pc_recall = []
  167. testset_mae = []
  168. testset_pc_mae = []
  169. testset_cm = []
  170. all_preds = []
  171. all_targets = []
  172. if data_set == "test":
  173. cur_set_to_test = X_test
  174. else:
  175. cur_set_to_test = X_val
  176. print("X_test:", X_test)
  177. print("X_val:", X_val)
  178. group_assigment = []
  179. counter = 0
  180. for curr_test_idx in cur_set_to_test.tolist():
  181. print(
  182. f" =============== Currently Evaluating curr_test_idx: {curr_test_idx} =============== "
  183. )
  184. if counter == 0:
  185. test_datamodule = GaitDataModule(
  186. all_files,
  187. batch_size=1024,
  188. window_size=config.data.window_size,
  189. step_size=config.data.window_size,
  190. train_idx=X_train.squeeze(),
  191. val_idx=X_val.squeeze()[
  192. :1
  193. ], # We load validation/test set via the test_idx here
  194. test_idx=[curr_test_idx],
  195. expand_labels=0,
  196. acc_sheet_name="Linear Accelerometer",
  197. num_workers=16,
  198. zscale=True,
  199. )
  200. test_datamodule.setup(data_set)
  201. trainset_stats = test_datamodule.train_dataset.zscale_stats
  202. counter += 1
  203. test_datamodule = GaitDataModule(
  204. all_files,
  205. batch_size=1024,
  206. window_size=config.data.window_size,
  207. step_size=config.data.window_size,
  208. train_idx=X_train.squeeze(),
  209. val_idx=X_val.squeeze()[
  210. :1
  211. ], # We load validation/test set via the test_idx here
  212. test_idx=[curr_test_idx],
  213. expand_labels=0,
  214. acc_sheet_name="Linear Accelerometer",
  215. num_workers=16,
  216. zscale=True,
  217. zscale_stats=trainset_stats,
  218. )
  219. test_datamodule.setup(data_set)
  220. test_dataloader = test_datamodule.test_dataloader()
  221. # if data_set == "test":
  222. # test_dataloader = test_datamodule.test_dataloader()
  223. # elif data_set == "val":
  224. # test_dataloader = test_datamodule.val_dataloader()
  225. # else:
  226. # raise ValueError(f"Not supported -> {data_set}")
  227. best_model.eval()
  228. # Benchmark single-window inference
  229. sample_x, sample_y = next(iter(test_dataloader))
  230. sample_x = sample_x[:1].to(device_) # shape: [1, 256, 6]
  231. sample_y = sample_y[:1].to(device_)
  232. # Warm-up
  233. with torch.inference_mode():
  234. for _ in range(20):
  235. _ = best_model(
  236. sample_x,
  237. sample_y,
  238. teacher_forcing_ratio=0.0,
  239. )
  240. if device_.type == "cuda":
  241. torch.cuda.synchronize()
  242. # Timed runs
  243. n_runs = 100
  244. inference_times = []
  245. with torch.inference_mode():
  246. for _ in range(n_runs):
  247. if device_.type == "cuda":
  248. torch.cuda.synchronize()
  249. start = time.perf_counter()
  250. _ = best_model(
  251. sample_x,
  252. sample_y,
  253. teacher_forcing_ratio=0.0,
  254. )
  255. if device_.type == "cuda":
  256. torch.cuda.synchronize()
  257. inference_times.append(time.perf_counter() - start)
  258. inference_times = np.array(inference_times) * 1000 # ms
  259. print(f"Mean inference time: {inference_times.mean():.3f} ms/window")
  260. print(f"Median inference time: {np.median(inference_times):.3f} ms/window")
  261. print(f"SD inference time: {inference_times.std():.3f} ms")
  262. with torch.no_grad():
  263. pred = []
  264. target = []
  265. for sample_input, sample_target in test_dataloader:
  266. sample_input = sample_input.to(device_)
  267. sample_target = sample_target.to(device_)
  268. print(sample_input.shape)
  269. predicted_output = best_model(
  270. sample_input, sample_target, teacher_forcing_ratio=0.0
  271. )
  272. pred.append(predicted_output)
  273. target.append(sample_target)
  274. cur_pred = torch.concat(pred).reshape(-1, 3)
  275. cur_pred = torch.nn.functional.softmax(cur_pred, dim=-1).argmax(1)
  276. cur_target = torch.concat(target).reshape(-1)
  277. reshaped_input = (
  278. sample_input.reshape(-1, sample_input.shape[-1]).cpu().numpy()
  279. )
  280. print("Pred", cur_pred.shape, "Target", cur_target.shape)
  281. # cm = MulticlassConfusionMatrix(num_classes=3, normalize="true").to(device_)
  282. # cm.update(cur_pred, cur_target)
  283. # fig_, ax_ = cm.plot()
  284. # plt.show()
  285. ALIGN_TOLERANCE = 3
  286. merged_preds = evaluate.merge_clustered_events(cur_pred.cpu().numpy())
  287. merged_targets = evaluate.merge_clustered_events(
  288. cur_target.cpu().numpy()
  289. )
  290. aligned_preds = evaluate.align_events(
  291. merged_targets, merged_preds, ALIGN_TOLERANCE
  292. )
  293. cm = MulticlassConfusionMatrix(num_classes=3, normalize="true")
  294. cm.update(torch.tensor(aligned_preds), torch.tensor(merged_targets))
  295. fig_, ax_ = cm.plot()
  296. plt.show()
  297. unnormalized_cm = MulticlassConfusionMatrix(
  298. num_classes=3, normalize=None
  299. )
  300. unnormalized_cm.update(
  301. torch.tensor(aligned_preds), torch.tensor(merged_targets)
  302. )
  303. ucm_tensor = unnormalized_cm.compute()
  304. testset_cm.append(ucm_tensor)
  305. print(ucm_tensor)
  306. num_classes = 3
  307. f1_macro = MulticlassF1Score(num_classes=num_classes, average="macro")
  308. f1_macro_score = f1_macro(
  309. torch.tensor(aligned_preds), torch.tensor(merged_targets)
  310. )
  311. print(f"Macro F1-Score: {f1_macro_score.item():.4f} ✨")
  312. pc_f1score = MulticlassF1Score(num_classes=num_classes, average="none")
  313. perclass_f1scores = pc_f1score(
  314. torch.tensor(aligned_preds), torch.tensor(merged_targets)
  315. )
  316. testset_pc_f1_scores.append(perclass_f1scores)
  317. print(f"perclass_f1scores:", perclass_f1scores)
  318. pc_precision = MulticlassPrecision(
  319. num_classes=num_classes, average="none"
  320. )
  321. perclass_precision = pc_precision(
  322. torch.tensor(aligned_preds), torch.tensor(merged_targets)
  323. )
  324. testset_pc_precision.append(perclass_precision)
  325. print(f"perclass_precision:", perclass_precision)
  326. pc_recall = MulticlassRecall(num_classes=num_classes, average="none")
  327. perclass_recall = pc_recall(
  328. torch.tensor(aligned_preds), torch.tensor(merged_targets)
  329. )
  330. testset_pc_recall.append(perclass_recall)
  331. print(f"perclass_recall:", perclass_recall)
  332. f1_macro_score = f1_macro(
  333. torch.tensor(aligned_preds), torch.tensor(merged_targets)
  334. )
  335. overall, per_class = calculate_gait_mae(merged_targets, merged_preds)
  336. overall = overall.item() if not np.isnan(overall) else np.nan
  337. print(f"MAE overall: {overall:.4f} ✨")
  338. print(f"MAE per_class: {per_class} ✨")
  339. testset_f1_scores.append(f1_macro_score)
  340. testset_mae.append(overall)
  341. testset_pc_mae.append(per_class)
  342. group_assigment.append(grouping[curr_test_idx])
  343. if return_preds_targets:
  344. all_preds.append(aligned_preds)
  345. all_targets.append(merged_targets)
  346. # break
  347. # break
  348. return {
  349. "testset_f1_scores": torch.tensor(testset_f1_scores).numpy(),
  350. "testset_pc_f1_scores": torch.stack(testset_pc_f1_scores).numpy(),
  351. "testset_pc_precision": torch.stack(testset_pc_precision).numpy(),
  352. "testset_pc_recall": torch.stack(testset_pc_recall).numpy(),
  353. "testset_mae": np.stack(testset_mae),
  354. "testset_pc_mae": testset_pc_mae,
  355. "testset_cm": torch.stack(testset_cm).numpy(),
  356. "group": np.stack(group_assigment),
  357. "all_preds": all_preds,
  358. "all_targets": all_targets,
  359. }
  360. # %% [markdown]
  361. # # 1. Evaluate Validation set
  362. # %%
  363. # Evaluate all best model on the validation set
  364. num_folds = 5
  365. # exp_labels = [2, 1, 4, 8]
  366. exp_labels = [2]
  367. data_set = "val" # using train here because we are loading all dataset anyway and using training set stats to normalize data
  368. for i in range(num_folds):
  369. for j in exp_labels:
  370. print(i, j)
  371. # curr_model_path = np.sort(glob(f"/home/geromevivar/projects/gait_ml/backpain/Fold{i+1}*expandlabel{j}*/*/*"))[-1]
  372. # curr_model_path = np.sort(glob(f"/home/qivy00li/projects/gait_ml/backpain/RerunExp-Fold{i+1}*expandlabel{j}*/*/*"))[-1]
  373. # curr_model_path = np.sort(glob(f"/home/qivy00li/projects/gait_ml/backpain/ZscaledRerunExp4-Fold{i+1}*expandlabel{j}*/*/*"))[-1]
  374. ##### MDPI
  375. # LSTM MDPI
  376. # curr_model_path = np.sort(glob(f"/home/qivy00li/projects/gait_ml/backpain_mdpi_runs/lstm/ZscaledRerunExp4-Fold{i+1}*expandlabel{j}*/*/*"))[-1]
  377. # GRU/LSTM weighted loss
  378. # curr_model_path = np.sort(glob(f"/home/qivy00li/projects/gait_ml/backpain_mdpi_runs/gru_lossweighted/ZscaledRerunExp4-Fold{i+1}*expandlabel{j}*/*/*"))[-1]
  379. # Rerun for MDPI
  380. curr_model_path = np.sort(
  381. glob(
  382. f"/home/qivy00li/projects/gait_ml/backpain_mdpi_runs/gru/ZscaledRerunExp4-Fold{i + 1}*expandlabel{j}*/*/*"
  383. )
  384. )[-1]
  385. print(os.path.exists(curr_model_path), curr_model_path)
  386. save_name = (
  387. f"ZscaledRerunExp4-Fold{i + 1}_explabel{j}_{Path(curr_model_path).stem}.npz"
  388. )
  389. eval_output_dir = os.path.join(
  390. "./mdpi_evals/", Path(curr_model_path).parents[1].name
  391. )
  392. save_name = os.path.join(eval_output_dir, save_name)
  393. if not os.path.exists(eval_output_dir):
  394. os.makedirs(eval_output_dir, exist_ok=True)
  395. # if data_set == "test" or data_set == "val":
  396. # save_name = f"{data_set}set_{save_name}"
  397. if not os.path.exists(save_name):
  398. print(f"=== Running model: {save_name} ===")
  399. testset_results = eval(data_set=data_set, model_fpath=curr_model_path, fold=i)
  400. # np.savez(save_name, **testset_results)
  401. # %% [markdown]
  402. # # 2. Plot
  403. # %% [markdown]
  404. # ### 2.1a - MDPI new plots
  405. # %%
  406. from pathlib import Path
  407. import re
  408. import matplotlib.pyplot as plt
  409. import numpy as np
  410. import pandas as pd
  411. from scipy.stats import t
  412. # ============================================================
  413. # SETTINGS
  414. # ============================================================
  415. ROOT = Path("./mdpi_evals")
  416. # 1 = ±10 ms
  417. # 2 = ±20 ms
  418. # 4 = ±40 ms
  419. # 8 = ±80 ms
  420. SELECTED_EXPANSION = 2
  421. METRIC_KEY = "testset_f1_scores"
  422. # ============================================================
  423. # DEFINE THE SIX EXPERIMENT GROUPS
  424. # Each pattern should find one folder per fold.
  425. # ============================================================
  426. experiments = [
  427. {
  428. "architecture": "GRU",
  429. "strategy": "Point labels",
  430. "pattern": "ZscaledRerunExp4-Fold*-GRU-expandlabel0_*",
  431. "weighted": False,
  432. },
  433. {
  434. "architecture": "GRU",
  435. "strategy": "Weighted CE",
  436. "pattern": "ZscaledRerunExp4-Fold*-gru-expandlabel0_*lossweighted=True",
  437. "weighted": True,
  438. },
  439. {
  440. "architecture": "GRU",
  441. "strategy": "Label expansion",
  442. "pattern": (f"ZscaledRerunExp4-Fold*-GRU-expandlabel{SELECTED_EXPANSION}_*"),
  443. "weighted": False,
  444. },
  445. {
  446. "architecture": "LSTM",
  447. "strategy": "Point labels",
  448. "pattern": "ZscaledRerunExp4-Fold*-lstm-expandlabel0_*",
  449. "weighted": False,
  450. },
  451. {
  452. "architecture": "LSTM",
  453. "strategy": "Weighted CE",
  454. "pattern": "ZscaledRerunExp4-Fold*-lstm-expandlabel0_*lossweighted=True",
  455. "weighted": True,
  456. },
  457. {
  458. "architecture": "LSTM",
  459. "strategy": "Label expansion",
  460. "pattern": (f"ZscaledRerunExp4-Fold*-lstm-expandlabel{SELECTED_EXPANSION}_*"),
  461. "weighted": False,
  462. },
  463. ]
  464. # ============================================================
  465. # LOAD ONE EVENT MACRO-F1 VALUE FROM EACH NPZ FILE
  466. # ============================================================
  467. def load_event_macro_f1(npz_path):
  468. """
  469. Return one combined IC–FO F1 score for one fold.
  470. Current implementation assumes METRIC_KEY contains only
  471. the IC and FO F1 scores, or values whose overall mean is
  472. the intended event macro-F1.
  473. """
  474. with np.load(npz_path, allow_pickle=True) as results:
  475. if METRIC_KEY not in results.files:
  476. raise KeyError(
  477. f"{METRIC_KEY!r} not found in {npz_path}.\n"
  478. f"Available keys: {results.files}"
  479. )
  480. scores = np.asarray(
  481. results[METRIC_KEY],
  482. dtype=float,
  483. )
  484. # Use this when the array contains only IC and FO scores.
  485. event_macro_f1 = float(np.nanmean(scores))
  486. # If the array instead contains [non-event, IC, FO],
  487. # replace the line above with:
  488. #
  489. # event_macro_f1 = float(np.nanmean(scores[..., 1:3]))
  490. return event_macro_f1
  491. # ============================================================
  492. # FIND FOLD FOLDERS AND BUILD DATAFRAME
  493. # ============================================================
  494. rows = []
  495. for experiment in experiments:
  496. candidate_folders = sorted(ROOT.glob(experiment["pattern"]))
  497. if (experiment["architecture"] == "LSTM") and (
  498. experiment["strategy"] == "Point labels"
  499. ):
  500. candidate_folders = [
  501. i for i in candidate_folders if "lossweighted=True" not in i.name
  502. ]
  503. # print(len(candidate_folders), experiment["pattern"])
  504. # Point-label and weighted-CE folders share the same base
  505. # pattern, so distinguish them using lossweighted=True.
  506. selected_folders = []
  507. for folder in candidate_folders:
  508. folder_is_weighted = "lossweighted=true" in folder.name.lower()
  509. if folder_is_weighted == experiment["weighted"]:
  510. selected_folders.append(folder)
  511. print(
  512. f"{experiment['architecture']:4s} | "
  513. f"{experiment['strategy']:16s} | "
  514. f"{len(selected_folders)} folders"
  515. )
  516. for folder in selected_folders:
  517. fold_match = re.search(
  518. r"Fold(\d+)",
  519. folder.name,
  520. flags=re.IGNORECASE,
  521. )
  522. if fold_match is None:
  523. print(f"Could not identify fold: {folder}")
  524. continue
  525. fold = int(fold_match.group(1))
  526. npz_files = sorted(folder.glob("*.npz"))
  527. if len(npz_files) == 0:
  528. print(f"No NPZ file found in: {folder}")
  529. continue
  530. if len(npz_files) > 1:
  531. print(f"Multiple NPZ files found in {folder}. Using: {npz_files[-1].name}")
  532. npz_path = npz_files[-1]
  533. macro_f1 = load_event_macro_f1(npz_path)
  534. rows.append(
  535. {
  536. "Architecture": experiment["architecture"],
  537. "Strategy": experiment["strategy"],
  538. "Fold": fold,
  539. "Macro-F1": macro_f1,
  540. "File": str(npz_path),
  541. }
  542. )
  543. df = pd.DataFrame(rows)
  544. if df.empty:
  545. raise RuntimeError("No results were loaded. Check ROOT and the folder patterns.")
  546. df = df.sort_values(["Architecture", "Strategy", "Fold"]).reset_index(drop=True)
  547. print("\nLoaded results:")
  548. print(
  549. df[
  550. [
  551. "Architecture",
  552. "Strategy",
  553. "Fold",
  554. "Macro-F1",
  555. ]
  556. ].to_string(index=False)
  557. )
  558. # ============================================================
  559. # CHECK THAT EACH CONDITION HAS FIVE FOLDS
  560. # ============================================================
  561. fold_counts = (
  562. df.groupby(["Architecture", "Strategy"])["Fold"].nunique().rename("Number of folds")
  563. )
  564. print("\nFold counts:")
  565. print(fold_counts)
  566. missing_conditions = fold_counts[fold_counts != 5]
  567. if not missing_conditions.empty:
  568. raise RuntimeError(
  569. "\nSome conditions do not contain exactly five folds:\n"
  570. f"{missing_conditions}\n\n"
  571. "Check the experiment folder patterns."
  572. )
  573. # ============================================================
  574. # CALCULATE MEAN AND 95% CONFIDENCE INTERVAL
  575. # ============================================================
  576. summary = (
  577. df.groupby(["Architecture", "Strategy"])["Macro-F1"]
  578. .agg(["mean", "std", "count"])
  579. .reset_index()
  580. )
  581. summary["standard_error"] = summary["std"] / np.sqrt(summary["count"])
  582. summary["ci95"] = (
  583. t.ppf(
  584. 0.975,
  585. df=summary["count"] - 1,
  586. )
  587. * summary["standard_error"]
  588. )
  589. print("\nSummary:")
  590. print(summary.to_string(index=False))
  591. # ============================================================
  592. # PLOT
  593. # ============================================================
  594. strategy_order = [
  595. "Point labels",
  596. "Weighted CE",
  597. "Label expansion",
  598. ]
  599. architecture_order = [
  600. "GRU",
  601. "LSTM",
  602. ]
  603. x = np.arange(len(strategy_order))
  604. architecture_offsets = {
  605. "GRU": -0.07,
  606. "LSTM": 0.07,
  607. }
  608. rng = np.random.default_rng(42)
  609. fig, ax = plt.subplots(figsize=(9, 5.5))
  610. for architecture in architecture_order:
  611. architecture_summary = (
  612. summary[summary["Architecture"] == architecture]
  613. .set_index("Strategy")
  614. .reindex(strategy_order)
  615. )
  616. x_architecture = x + architecture_offsets[architecture]
  617. plot_result = ax.errorbar(
  618. x_architecture,
  619. architecture_summary["mean"],
  620. yerr=architecture_summary["ci95"],
  621. marker="o",
  622. markersize=7,
  623. linewidth=2,
  624. capsize=5,
  625. label=architecture,
  626. )
  627. line_color = plot_result.lines[0].get_color()
  628. # Add individual fold values.
  629. for position, strategy in enumerate(strategy_order):
  630. fold_values = df[
  631. (df["Architecture"] == architecture) & (df["Strategy"] == strategy)
  632. ]["Macro-F1"].to_numpy()
  633. jitter = rng.normal(
  634. loc=0,
  635. scale=0.012,
  636. size=len(fold_values),
  637. )
  638. ax.scatter(
  639. np.full(
  640. len(fold_values),
  641. x_architecture[position],
  642. )
  643. + jitter,
  644. fold_values,
  645. s=35,
  646. alpha=0.65,
  647. color=line_color,
  648. zorder=3,
  649. )
  650. expansion_ms = SELECTED_EXPANSION * 10
  651. ax.set_xticks(x)
  652. ax.set_xticklabels(
  653. [
  654. "Point labels",
  655. "Weighted CE",
  656. f"Label expansion\n±{expansion_ms} ms",
  657. ]
  658. )
  659. ax.set_xlabel("Training strategy", fontsize=16, fontweight="bold")
  660. ax.set_ylabel("IC–FO macro-F1", fontsize=16, fontweight="bold")
  661. # ax.set_title(
  662. # "Comparison of Architecture and Sparse-Label Strategy"
  663. # )
  664. ax.set_ylim(0, 1.02)
  665. ax.legend(title="Architecture")
  666. ax.grid(axis="y", alpha=0.25)
  667. plt.xticks(fontsize=14)
  668. plt.yticks(fontsize=14)
  669. plt.tight_layout()
  670. plt.savefig(
  671. "architecture_strategy_comparison.png",
  672. dpi=300,
  673. bbox_inches="tight",
  674. )
  675. plt.show()
  676. # %% [markdown]
  677. # ## 2.1 Plot F1-score and MAE box plot validation set
  678. # %%
  679. # Evaluate all best model on the validation set
  680. num_folds = 5
  681. exp_labels = [0, 1, 2, 4, 8]
  682. all_results = []
  683. for i in range(num_folds):
  684. for j in exp_labels:
  685. # curr_model_path = np.sort(glob(f"/home/geromevivar/projects/gait_ml/backpain/Fold{i+1}*expandlabel{j}*/*/*"))[-1]
  686. curr_model_path = np.sort(
  687. glob(
  688. f"/home/qivy00li/projects/gait_ml/backpain/RerunExp-Fold{i + 1}*expandlabel{j}*/*/*"
  689. )
  690. )[-1]
  691. fname = f"valset_Fold{i + 1}_explabel{j}_{Path(curr_model_path).stem}.npz"
  692. print(f" === Processing {fname} === ")
  693. curr_results = np.load(fname, allow_pickle=True)
  694. curr_results = {**curr_results}
  695. curr_results["fold"] = i + 1
  696. curr_results["explabel"] = j
  697. all_results.append(curr_results)
  698. all_results_df = pd.DataFrame(all_results)
  699. print(all_results_df.head())
  700. # Simplified Code for MEAN OF LIST LENGTHS
  701. cur_feature_name = "testset_f1_scores"
  702. cur_metric_name = "F1-Score"
  703. f1_scores = all_results_df.groupby(["fold", "explabel"])[[cur_feature_name]].agg(
  704. lambda s: s.apply(np.nanmean)
  705. )
  706. f1_scores = f1_scores.reset_index()
  707. f1_scores.head()
  708. # f1_scores.groupby(["explabel"]).mean()
  709. # f1_scores.groupby(["explabel"]).std()
  710. plot_df = pd.melt(f1_scores[f1_scores.explabel != 0], id_vars=["fold", "explabel"])
  711. plot_df = plot_df.rename(columns={"explabel": "Expand", "value": cur_metric_name})
  712. plt.figure(figsize=(10, 6))
  713. g = sns.boxplot(data=plot_df, x="Expand", y=cur_metric_name)
  714. # plt.title("Event Detection Performance on Validation Set (5-fold CV)")
  715. g.set_xticklabels(
  716. ["\u00b1 10 [ms]", "\u00b1 20 [ms]", "\u00b1 40 [ms]", "\u00b1 80 [ms]"]
  717. )
  718. g.set_xlabel("Model's Label Expansion Setting")
  719. g.set_title
  720. # %%
  721. # JNER version
  722. import matplotlib.pyplot as plt
  723. import seaborn as sns
  724. import pandas as pd
  725. import numpy as np
  726. # --- 1. JNER Style Setup ---
  727. # Use standard sans-serif fonts (Arial/Helvetica)
  728. plt.rcParams["font.family"] = "sans-serif"
  729. plt.rcParams["font.sans-serif"] = ["Arial", "Helvetica", "DejaVu Sans"]
  730. # --- [Your Data Loading Block Stays Here] ---
  731. # (Assuming all_results_df is created as in your snippet)
  732. # --- Data Preparation ---
  733. cur_feature_name = "testset_f1_scores"
  734. cur_metric_name = "F1-Score"
  735. # Process data
  736. f1_scores = all_results_df.groupby(["fold", "explabel"])[[cur_feature_name]].agg(
  737. lambda s: s.apply(np.nanmean)
  738. )
  739. f1_scores = f1_scores.reset_index()
  740. # Filter out 0ms if desired
  741. plot_df = pd.melt(f1_scores[f1_scores.explabel != 0], id_vars=["fold", "explabel"])
  742. plot_df = plot_df.rename(columns={"explabel": "Expand", "value": cur_metric_name})
  743. # --- 2. Plotting for Publication ---
  744. # Figure Size:
  745. # Journals usually have columns ~3.5 inches wide.
  746. # A width of 6-8 inches allows it to span two columns or be scaled down nicely.
  747. plt.figure(figsize=(8, 5))
  748. # Style: White background with grid is standard for scientific comparison
  749. sns.set_style("whitegrid")
  750. # Define explicit order to ensure labels match data
  751. # (Assuming exp_labels 1, 2, 4, 8 correspond to 10, 20, 40, 80)
  752. order_list = [1, 2, 4, 8]
  753. g = sns.boxplot(
  754. data=plot_df,
  755. x="Expand",
  756. y=cur_metric_name,
  757. order=order_list, # Critical: Ensures X-axis is sorted correctly
  758. width=0.5, # Thinner boxes look cleaner
  759. linewidth=1.2, # Thicker lines for visibility in print
  760. palette="Blues", # "Blues" is aesthetically pleasing and printer-safe
  761. showfliers=False, # Optional: Hide outliers if they distract (check journal preference)
  762. )
  763. # --- 3. Formatting Axes ---
  764. # Y-Axis: Use LaTeX rendering for F1-score if possible, or consistent text
  765. # Note: Matplotlib can render simple math-like text without full LaTeX
  766. plt.ylabel(r"$\mathbf{F_1}$-score", fontsize=16, fontweight="bold")
  767. # X-Axis
  768. plt.xlabel("Label Expansion Window", fontsize=16, fontweight="bold")
  769. # Ticks: Use the ± symbol and standard units
  770. # Mapping 1->10ms, 2->20ms, etc. based on your previous snippet
  771. clean_labels = ["$\pm$10 ms", "$\pm$20 ms", "$\pm$40 ms", "$\pm$80 ms"]
  772. g.set_xticklabels(clean_labels, fontsize=14)
  773. plt.yticks(fontsize=14)
  774. # Remove the top and right spines (cleaner look)
  775. sns.despine()
  776. # --- 4. Saving ---
  777. # Remove title (It belongs in the LaTeX caption, not the image)
  778. # plt.title("...")
  779. plt.tight_layout()
  780. # Save as PDF (Vector - Best) or PNG (Raster - High DPI)
  781. plt.savefig("X_figures/f1_score_expansion_boxplot.pdf", bbox_inches="tight")
  782. plt.savefig("X_figures/f1_score_expansion_boxplot.png", dpi=600, bbox_inches="tight")
  783. plt.show()
  784. # %%
  785. # all_results_df[all_results_df.explabel==1].testset_f1_scores.apply(lambda x: x.mean()).mean()
  786. # %%
  787. import pandas as pd
  788. import numpy as np
  789. from scipy import stats
  790. # --- Simplified Aggregation ---
  791. # A single lambda function to calculate the 95% Confidence Interval bounds (CI)
  792. # t.interval returns a tuple: (lower_bound, upper_bound)
  793. def get_ci_bounds(series, confidence=0.95):
  794. """Calculates the 95% CI (lower, upper) for a Series."""
  795. if series.empty:
  796. return np.nan, np.nan
  797. # stats.sem calculates the Standard Error of the Mean (sigma / sqrt(n))
  798. sem = stats.sem(series, ddof=1)
  799. # stats.t.interval computes the confidence interval
  800. return stats.t.interval(confidence, len(series) - 1, loc=series.mean(), scale=sem)
  801. # Assuming your DataFrame 'f1_scores' and column 'cur_feature_name' are defined.
  802. decimals = 3
  803. # 1. Aggregate the data (combining the CI calculation)
  804. report_df = (
  805. f1_scores.groupby("explabel")[cur_feature_name]
  806. .agg(
  807. mean_val="mean",
  808. std_val="std",
  809. ci_bounds=get_ci_bounds, # Uses the single function to get a tuple of bounds
  810. count="size",
  811. )
  812. .reset_index()
  813. )
  814. # 2. Split the CI tuple into separate columns for easier formatting
  815. report_df[["ci_lower", "ci_upper"]] = pd.DataFrame(
  816. report_df["ci_bounds"].tolist(), index=report_df.index
  817. )
  818. report_df = report_df.drop(columns=["ci_bounds"])
  819. # --- Simplified Formatting (Using f-strings and round) ---
  820. # 3. Create the formatted columns using vectorization (not row-wise apply) where possible
  821. # Note: Using .round() before f-string formatting ensures correct precision.
  822. # Mean +/- Std Column
  823. report_df["Mean_Std"] = (
  824. report_df["mean_val"].round(decimals).astype(str)
  825. + " $\\pm$ "
  826. + report_df["std_val"].round(decimals).astype(str)
  827. )
  828. # Confidence Interval Column
  829. report_df["95% CI"] = (
  830. "["
  831. + report_df["ci_lower"].round(decimals).astype(str)
  832. + ", "
  833. + report_df["ci_upper"].round(decimals).astype(str)
  834. + "]"
  835. )
  836. # 4. Select and rename final columns for publication
  837. publication_table = report_df[["explabel", "count", "Mean_Std", "95% CI"]]
  838. publication_table.columns = [
  839. "Experiment Label (Group)",
  840. "N",
  841. "Mean $\\pm$ Std. Dev.",
  842. "95% Confidence Interval",
  843. ]
  844. # %%
  845. publication_table.columns = [
  846. "Label Expand",
  847. "CV-folds",
  848. "Mean $\pm$ Std. Dev.",
  849. "95% Confidence Interval",
  850. ]
  851. publication_table = publication_table.iloc[:, 1:3]
  852. publication_table
  853. # %%
  854. # Reporting validation set metrics
  855. publication_table.to_latex()
  856. # %%
  857. # Just get the median value
  858. f1_scores.groupby("explabel")[cur_feature_name].agg(
  859. mean_val="median",
  860. std_val="std",
  861. ci_bounds=get_ci_bounds, # Uses the single function to get a tuple of bounds
  862. count="size",
  863. ).reset_index()
  864. # %% [markdown]
  865. # ## 2.2 Test set - Plot confusion matrix with CI on the test set using best model only
  866. # %%
  867. # Evaluate all best model on the validation set
  868. num_folds = 5
  869. exp_labels = [2]
  870. data_set = "test"
  871. for i in range(num_folds):
  872. for j in exp_labels:
  873. # curr_model_path = np.sort(glob(f"/home/geromevivar/projects/gait_ml/backpain/Fold{i+1}*expandlabel{j}*/*/*"))[-1]
  874. curr_model_path = np.sort(
  875. glob(
  876. f"/home/qivy00li/projects/gait_ml/backpain/ZscaledRerunExp4-Fold{i + 1}*expandlabel{j}*/*/*"
  877. )
  878. )[-1]
  879. save_name = f"Fold{i + 1}_explabel{j}_{Path(curr_model_path).stem}.npz"
  880. if data_set == "test":
  881. save_name = f"{data_set}set_{save_name}"
  882. save_name = save_name.replace(
  883. "testset_", "rerunWithPredAndTargetSaved_testset_"
  884. )
  885. save_name = save_name.replace(".npz", ".pkl")
  886. testset_results = eval(
  887. data_set=data_set,
  888. model_fpath=curr_model_path,
  889. fold=i,
  890. return_preds_targets=return_preds_targets,
  891. )
  892. if not os.path.exists(save_name):
  893. print(f"=== Running model: {save_name} ===")
  894. return_preds_targets = False
  895. testset_results = eval(
  896. data_set=data_set,
  897. model_fpath=curr_model_path,
  898. fold=i,
  899. return_preds_targets=return_preds_targets,
  900. )
  901. if return_preds_targets:
  902. with open(save_name, "wb") as f:
  903. pickle.dump(testset_results, f)
  904. else:
  905. np.savez(
  906. save_name, **testset_results
  907. ) # used in original paper experiments
  908. # %%
  909. testset_all_results = []
  910. print(f"Currently evaluating expand label: {j}")
  911. testset_fnames = np.sort(glob(f"testset_Fold*_explabel{j}*"))
  912. for fname in testset_fnames:
  913. print(f"Loading {fname}")
  914. curr_results = np.load(fname, allow_pickle=True)
  915. curr_results = {**curr_results}
  916. curr_results["fold"] = i + 1
  917. curr_results["explabel"] = j
  918. testset_all_results.append(curr_results)
  919. # %%
  920. testset_all_results_df = pd.DataFrame(testset_all_results)
  921. testset_all_results_df
  922. # %%
  923. testset_all_results_df.testset_mae.apply(lambda x: len(x))
  924. # %%
  925. cm_per_fold = testset_all_results_df.testset_cm.apply(
  926. lambda x: np.stack(x).sum(0) / np.stack(x).sum(0).sum(1)
  927. )
  928. cm_per_fold
  929. # %%
  930. combined_cm = testset_all_results_df.testset_cm.apply(lambda x: np.stack(x).sum(0)).sum(
  931. 0
  932. )
  933. normalized_cm = (combined_cm / combined_cm.sum(1)).round(2)
  934. # %%
  935. import numpy as np
  936. import matplotlib.pyplot as plt
  937. import seaborn as sns
  938. from sklearn.metrics import confusion_matrix
  939. # ... other imports ...
  940. # def plot_confusion_matrix(cm, class_names, title, fmt, cbar_label, cmap=plt.cm.Blues, file_name='confusion_matrix.png'):
  941. # # ... (function body as executed above) ...
  942. # fig, ax = plt.subplots(figsize=(5, 5))
  943. # sns.heatmap(
  944. # cm,
  945. # annot=True,
  946. # fmt=fmt,
  947. # cmap=cmap,
  948. # linewidths=0.5,
  949. # linecolor='black',
  950. # cbar=False,
  951. # # cbar_kws={'label': cbar_label, 'orientation': 'vertical', 'pad': 0.04, 'aspect': 30},
  952. # annot_kws={"fontsize": 16, "fontweight": "bold"},
  953. # ax=ax,
  954. # square=True
  955. # )
  956. # ax.set_title(title, fontsize=12, fontweight='bold', pad=5)
  957. # ax.set_ylabel('True Label', fontsize=20, fontweight='medium')
  958. # ax.set_xlabel('Predicted Label', fontsize=20, fontweight='medium')
  959. # # Set class labels on ticks, centering them
  960. # tick_marks = np.arange(len(class_names))
  961. # ax.set_xticks(tick_marks + 0.5)
  962. # ax.set_yticks(tick_marks + 0.5)
  963. # ax.set_xticklabels(class_names, fontsize=16)
  964. # ax.set_yticklabels(class_names, fontsize=16, rotation=90, va="center")
  965. # # Fix for half-pixel issues in matplotlib 3.1.1+ (sets the limits correctly)
  966. # ax.set_ylim(len(class_names), 0)
  967. # ax.tick_params(axis='both', which='major', length=0)
  968. # plt.tight_layout()
  969. # # ... (code to set ticks and save figure) ...
  970. # # plt.savefig(file_name, dpi=300, bbox_inches='tight')
  971. # # fig.show()
  972. # # plt.close(fig)
  973. # return fig
  974. # def plot_confusion_matrix(cm, class_names, title=None, fmt='d', cbar_label='Count', cmap=plt.cm.Blues, file_name='confusion_matrix.pdf'):
  975. # """
  976. # Plots a publication-ready confusion matrix for JNER.
  977. # """
  978. # # 1. Set JNER-compliant font (Arial/Helvetica is standard)
  979. # plt.rcParams['font.family'] = 'sans-serif'
  980. # plt.rcParams['font.sans-serif'] = ['Arial', 'Helvetica', 'DejaVu Sans']
  981. # # 2. Size: 3.5 inches is standard for single-column width (85-90mm)
  982. # fig, ax = plt.subplots(figsize=(5, 5))
  983. # sns.heatmap(
  984. # cm,
  985. # annot=True,
  986. # fmt=fmt,
  987. # cmap=cmap,
  988. # linewidths=1.0, # Thicker lines for better separation in print
  989. # linecolor='black',
  990. # cbar=False, # Disable colorbar if numbers are annotated (saves space)
  991. # annot_kws={"fontsize": 18, "fontweight": "bold"}, # Large font for readability when resized
  992. # ax=ax,
  993. # square=True
  994. # )
  995. # # 3. Titles: JNER prefers titles in the caption, not the image.
  996. # # Only set if strictly necessary for internal use.
  997. # if title:
  998. # ax.set_title(title, fontsize=14, fontweight='bold', pad=10)
  999. # # 4. Axis Labels: Clear and large
  1000. # ax.set_ylabel('True Event', fontsize=18, fontweight='bold')
  1001. # ax.set_xlabel('Predicted Event', fontsize=18, fontweight='bold')
  1002. # # 5. Ticks: Center them and ensure readability
  1003. # tick_marks = np.arange(len(class_names))
  1004. # ax.set_xticks(tick_marks + 0.5)
  1005. # ax.set_yticks(tick_marks + 0.5)
  1006. # ax.set_xticklabels(class_names, fontsize=16, fontweight='medium')
  1007. # # CHANGED: Rotation 0 is better for short labels like "IC/FO"
  1008. # ax.set_yticklabels(class_names, fontsize=16, fontweight='medium', rotation=0, va="center")
  1009. # # Cleanups
  1010. # ax.tick_params(axis='both', which='major', length=0)
  1011. # plt.tight_layout()
  1012. # # 6. Saving: Use 600 DPI for raster or PDF/EPS for vector (Best for JNER)
  1013. # # If saving as PNG, use 600 dpi. If PDF, dpi is less critical but good practice.
  1014. # plt.savefig(file_name, dpi=600, bbox_inches='tight', transparent=False)
  1015. # # plt.close(fig) # Uncomment to prevent display in notebooks if generating many
  1016. # return fig
  1017. # %%
  1018. from gait_ml import utils
  1019. for k, v in testset_all_results_df.iterrows():
  1020. print(k)
  1021. # plot_confusion_matrix((np.stack(v.testset_cm).sum(0)/np.stack(v.testset_cm).sum(0).sum(1)).round(2), ["NE", "IC", "TO"], "Detection Performance on Test Set", ".2f", ["NE", "IC", "TO"])
  1022. utils.plot_confusion_matrix(
  1023. cm=(np.stack(v.testset_cm).sum(0) / np.stack(v.testset_cm).sum(0).sum(1)).round(
  1024. 2
  1025. ),
  1026. class_names=["NE", "IC", "TO"],
  1027. title=None,
  1028. fmt=".2f",
  1029. cbar_label="Count",
  1030. file_name=f"X_figures/cm-f{k + 1}.pdf",
  1031. )
  1032. # %%
  1033. utils.plot_confusion_matrix(
  1034. cm=normalized_cm,
  1035. class_names=["NE", "IC", "TO"],
  1036. title=None,
  1037. fmt=".2f",
  1038. cbar_label="Count",
  1039. file_name=f"X_figures/aggregated_cm_testset.pdf",
  1040. )
  1041. # %% [markdown]
  1042. # ## 2.3 Table detection and MAE metrics on test set mean [std]
  1043. # - F1-score, Precision, Recall, TP, FP, FN, MAE per class per group
  1044. # %%
  1045. metrics_per_fold = []
  1046. for k, v in testset_all_results_df.iterrows():
  1047. print(f"==Processing fold: {k}==")
  1048. for j in range(2):
  1049. print(f"Group: {j}")
  1050. curr_metrics = dict()
  1051. curr_group = v.group
  1052. curr_metrics["f1-score"] = v.testset_pc_f1_scores[curr_group == j].mean(0)
  1053. curr_metrics["recall"] = v.testset_pc_recall[curr_group == j].mean(0)
  1054. curr_metrics["precision"] = v.testset_pc_precision[curr_group == j].mean(0)
  1055. curr_mae = pd.Series(v.testset_pc_mae[curr_group == j]).apply(
  1056. lambda x: [x[idx + 1] for idx in range(2)]
  1057. )
  1058. curr_mae = np.stack(curr_mae).mean(0)
  1059. curr_mae = np.hstack([np.array([np.nan]), curr_mae])
  1060. curr_metrics["mae"] = curr_mae * 10.0
  1061. curr_metrics["group"] = j
  1062. curr_metrics["fold"] = k + 1
  1063. metrics_per_fold.append(curr_metrics)
  1064. results_df = pd.DataFrame(metrics_per_fold)
  1065. # results_df.drop(columns=["fold"], inplace=True)
  1066. # %%
  1067. # Prepare dataframe for plotting
  1068. plot_df = results_df.melt(["group", "fold"])
  1069. event_df = pd.DataFrame(plot_df["value"].apply(pd.Series))
  1070. event_df.columns = ["NE", "IC", "FO"]
  1071. plot_df.drop(columns="value", inplace=True)
  1072. plot_df = pd.concat([plot_df, event_df], axis=1)
  1073. df_long = plot_df.melt(
  1074. id_vars=["group", "fold", "variable"],
  1075. value_vars=["NE", "IC", "FO"],
  1076. var_name="event", # New column for the variable names
  1077. value_name="value",
  1078. )
  1079. df_long.group.replace(0, "Healthy", inplace=True)
  1080. df_long.group.replace(1, "BackPain", inplace=True)
  1081. df_long.variable.replace("mae", "MAE", inplace=True)
  1082. df_long.rename(columns={"variable": "Metric", "group": "Group"}, inplace=True)
  1083. # %%
  1084. df_long.groupby(["Metric", "Group", "event"])["value"].describe().round(3)
  1085. # %%
  1086. import matplotlib.pyplot as plt
  1087. import seaborn as sns
  1088. # Assuming df_long is defined and contains your data
  1089. sns.set_context("poster")
  1090. print("+++++ WarningL: Excluding NE events in the plot!++++")
  1091. df_long = df_long[df_long.event != "NE"]
  1092. g = sns.catplot(
  1093. data=df_long,
  1094. x="event",
  1095. y="value",
  1096. hue="Group",
  1097. kind="box",
  1098. col="Metric",
  1099. col_wrap=2,
  1100. height=6,
  1101. aspect=1.2,
  1102. sharey=False,
  1103. sharex=False,
  1104. palette="colorblind",
  1105. )
  1106. # --- Step 1: Define the custom Titles, Y-labels, and Performance Goal ---
  1107. metrics = df_long["Metric"].unique()
  1108. # # Define the custom info for each metric, now including a custom title
  1109. # custom_metrics_info = {
  1110. # metrics[0]: {"title": "F1-Score Detection Performance", "label": "F1-Score (%)", "goal": "Higher is Better (↑)", "y_pos": 0.98},
  1111. # metrics[1]: {"title": "Mean Error Results", "label": "Mean Error [ms]", "goal": "Lower is Better (↓)", "y_pos": 0.05},
  1112. # metrics[2]: {"title": "Event Precision Analysis", "label": "Detection Precision", "goal": "Higher is Better (↑)", "y_pos": 0.98},
  1113. # metrics[3]: {"title": "Root Mean Squared Error [ms]", "label": "Root Mean Squared Error", "goal": "Lower is Better (↓)", "y_pos": 0.05}
  1114. # }
  1115. custom_metrics_info = {
  1116. metrics[0]: {"title": "F1-Score (↑ better)"},
  1117. metrics[1]: {"title": "Recall (↑ better)"},
  1118. metrics[2]: {"title": "Precision(↑ better)"},
  1119. metrics[3]: {"title": "Mean Absolute Error [ms] (↓ better)"},
  1120. }
  1121. # --- Step 2: Iterate and Apply Titles, Labels, and Annotations ---
  1122. for ax_index, ax in enumerate(g.axes.flat):
  1123. current_metric = metrics[ax_index]
  1124. info = custom_metrics_info.get(current_metric)
  1125. if info:
  1126. # ⭐ Key Customization 1: Set the custom plot title
  1127. ax.set_title(info["title"], fontsize=20, fontweight="bold")
  1128. # Set the custom Y-axis label
  1129. # ax.set_ylabel(info["label"], fontsize=18)
  1130. # # Add text annotation to indicate the goal
  1131. # ax.text(
  1132. # x=0.05,
  1133. # y=info["y_pos"],
  1134. # s=info["goal"],
  1135. # transform=ax.transAxes,
  1136. # fontsize=16,
  1137. # color='red' if 'Lower' in info["goal"] else 'green',
  1138. # fontweight='bold'
  1139. # )
  1140. # --- Step 3: Clean up shared labels and titles ---
  1141. # Remove the default shared label from the grid
  1142. g.set_axis_labels("Gait Events", "Values", fontweight="bold")
  1143. # Remove the default top-level title that catplot tries to set for the column
  1144. # g.set_titles(col_template='{col_name}', row_template='{row_name}', size=0) # Set size=0 to hide
  1145. # Add a main title to the figure (applies to the entire figure, not individual plots)
  1146. # g.fig.suptitle('Gait Event Detection Performance on Testset (5-fold CV)', y=1.03, fontsize=22)
  1147. # Improve tick labels (numbers/text on axes)
  1148. g.tick_params(axis="both", which="major", labelsize=20)
  1149. plt.savefig("X_figures/overall-performance2x2.pdf", bbox_inches="tight")
  1150. plt.savefig("X_figures/overall-performance2x2.png", dpi=600, bbox_inches="tight")
  1151. plt.show()
  1152. # %%
  1153. # plt.figure(figsize=(10, 5))
  1154. # g = sns.catplot(data=df_long,
  1155. # x="event",
  1156. # y="value",
  1157. # hue="Group",
  1158. # kind="box",
  1159. # col='Metric',
  1160. # col_wrap=2,
  1161. # height=6,
  1162. # aspect=1.2,
  1163. # sharey=False,
  1164. # palette='colorblind')
  1165. # # Add a main title to the figure
  1166. # g.fig.suptitle('Gait Event Detection Performance on Testset (5-fold CV)', y=1.03, fontsize=22)
  1167. # # Improve axis labels|
  1168. # g.tick_params(axis='both', which='major', labelsize=16)
  1169. # # plt.tight_layout(rect=[0, 0, 1, 0.97])
  1170. # g.set_axis_labels("Event", "Value")
  1171. # plt.show()
  1172. # %%
  1173. # # Create the plot
  1174. # g = sns.catplot(data=df_long,
  1175. # x="event",
  1176. # y="value",
  1177. # hue="Group",
  1178. # kind="box",
  1179. # col='Metric',
  1180. # # col_wrap=2,
  1181. # height=6,
  1182. # aspect=1.2,
  1183. # sharey=False,
  1184. # legend_out=True,
  1185. # palette='colorblind')
  1186. # # 1. Main Title
  1187. # # g.fig.suptitle('Gait Event Detection Performance on Testset (5-fold CV)', y=1.03, fontsize=20)
  1188. # # 2. Subplot Titles (e.g., "Metric = Accuracy")
  1189. # g.set_titles(size=25)
  1190. # # 3. Axis Labels
  1191. # g.set_axis_labels("Event", "Value", fontsize=25)
  1192. # # 4. Tick Labels
  1193. # g.tick_params(axis='both', which='major', labelsize=25)
  1194. # # 5. Legend Title and Labels
  1195. # if g.legend:
  1196. # plt.setp(g.legend.get_texts(), fontsize='25')
  1197. # plt.setp(g.legend.get_title(), fontsize='25')
  1198. # # Adjust layout
  1199. # # plt.tight_layout(rect=[0, 0, 1, 0.97])
  1200. # plt.show()
  1201. # %% [markdown]
  1202. # ### 3. Concordance Analysis
  1203. # %%
  1204. from gait_ml import utils
  1205. import numpy as np
  1206. # %%
  1207. # Evaluate all best model on the validation set
  1208. num_folds = 5
  1209. exp_labels = [2]
  1210. data_set = "test"
  1211. for i in range(num_folds):
  1212. if i == 0:
  1213. for j in exp_labels:
  1214. # curr_model_path = np.sort(glob(f"/home/geromevivar/projects/gait_ml/backpain/Fold{i+1}*expandlabel{j}*/*/*"))[-1]
  1215. curr_model_path = np.sort(
  1216. glob(
  1217. f"/home/qivy00li/projects/gait_ml/backpain/ZscaledRerunExp4-Fold{i + 1}*expandlabel{j}*/*/*"
  1218. )
  1219. )[-1]
  1220. save_name = f"Fold{i + 1}_explabel{j}_{Path(curr_model_path).stem}.npz"
  1221. if data_set == "test":
  1222. save_name = f"{data_set}set_{save_name}"
  1223. # if not os.path.exists(save_name):
  1224. print(f"=== Running model: {save_name} ===")
  1225. testset_results = eval(
  1226. data_set=data_set,
  1227. model_fpath=curr_model_path,
  1228. fold=i,
  1229. return_preds_targets=True,
  1230. )
  1231. # np.savez(save_name, **testset_results)
  1232. break
  1233. # %%
  1234. all_preds = testset_results["all_preds"]
  1235. all_targets = testset_results["all_targets"]
  1236. # %%
  1237. # Find all patient ids
  1238. # subject_ids = [Path(i).parents[3].name for i in test_set_fpaths]
  1239. # Take average feature per subject since there are ~
  1240. stridetime_pred = [utils.calculate_stride_times(i, 100) for i in all_preds]
  1241. stridetime_target = [utils.calculate_stride_times(i, 100) for i in all_targets]
  1242. stridetime_pred = pd.DataFrame([i.mean() for i in stridetime_pred])
  1243. stridetime_target = pd.DataFrame([i.mean() for i in stridetime_target])
  1244. # stance_pred = utils.aggregate_res(stance_pred, subject_ids)
  1245. # stance_target = utils.aggregate_res(stance_target, subject_ids)
  1246. stridetime_res_df = pd.concat([stridetime_pred, stridetime_target], axis=1)
  1247. stridetime_res_df
  1248. # %%
  1249. # Find all patient ids
  1250. # subject_ids = [Path(i).parents[3].name for i in test_set_fpaths]
  1251. # Take average feature per subject since there are ~
  1252. stance_pred = [utils.calculate_gait_phases_vectorized(i) for i in all_preds]
  1253. stance_target = [utils.calculate_gait_phases_vectorized(i) for i in all_targets]
  1254. stance_pred = pd.DataFrame([i.mean() for i in stance_pred])
  1255. stance_target = pd.DataFrame([i.mean() for i in stance_target])
  1256. # stance_pred = utils.aggregate_res(stance_pred, subject_ids)
  1257. # stance_target = utils.aggregate_res(stance_target, subject_ids)
  1258. # %%
  1259. res_df = pd.concat([stance_pred, stance_target], axis=1)
  1260. res_df
  1261. # %%
  1262. swing_prop_df = 100 - res_df
  1263. # %%
  1264. # 2. Call the plotting function
  1265. fig, ax = utils.plot_bland_altman_publication(
  1266. swing_prop_df.iloc[:, 0],
  1267. swing_prop_df.iloc[:, 1],
  1268. method1_name="Pred",
  1269. method2_name="GT",
  1270. units=r"[$\%_{Gait}$]",
  1271. filename="mdpi_figures/swing_bland_altman_internal.pdf",
  1272. feature_name="Swing Phase",
  1273. ypos=0.75,
  1274. )
  1275. # %%
  1276. # 2. Call the plotting function
  1277. fig, ax = utils.plot_bland_altman_publication(
  1278. res_df.iloc[:, 0],
  1279. res_df.iloc[:, 1],
  1280. method1_name="Pred",
  1281. method2_name="GT",
  1282. units=r"[$\%_{Gait}$]",
  1283. filename="mdpi_figures/stance_bland_altman_internal.pdf",
  1284. feature_name="Stance Phase",
  1285. )
  1286. # %%
  1287. # 2. Call the plotting function
  1288. fig, ax = utils.plot_bland_altman_publicationv2(
  1289. stridetime_res_df.iloc[:, 0],
  1290. stridetime_res_df.iloc[:, 1],
  1291. method1_name="Pred",
  1292. method2_name="GT",
  1293. units=r"ms",
  1294. filename="mdpi_figures/stridetime_bland_altman_internal.pdf",
  1295. feature_name="Stride Time",
  1296. )
  1297. # %% [markdown]
  1298. # ### Confidence intervals
  1299. # %%
  1300. ccc, ccc_ci_lower, ccc_ci_upper = utils.lins_ccc_with_ci(
  1301. stridetime_res_df.iloc[:, 0],
  1302. stridetime_res_df.iloc[:, 1],
  1303. confidence_level=0.95,
  1304. n_resamples=1000,
  1305. random_seed=42,
  1306. )
  1307. # %%
  1308. print("ccc:", round(ccc, 3))
  1309. print("ccc_ci_lower:", round(ccc_ci_lower, 3))
  1310. print("ccc_ci_upper:", round(ccc_ci_upper, 3))
  1311. # %%
  1312. ccc, ccc_ci_lower, ccc_ci_upper = utils.lins_ccc_with_ci(
  1313. res_df.iloc[:, 0],
  1314. res_df.iloc[:, 1],
  1315. confidence_level=0.95,
  1316. n_resamples=1000,
  1317. random_seed=42,
  1318. )
  1319. # %%
  1320. print("ccc:", round(ccc, 3))
  1321. print("ccc_ci_lower:", round(ccc_ci_lower, 3))
  1322. print("ccc_ci_upper:", round(ccc_ci_upper, 3))
  1323. # %%

03_RNN_eval.ipynb at commit 3c1aed0, under other · at the source

Overview

  1. Health and Physical Activity, Otto von Guericke University Magdeburg, 39104 Magdeburg, Germany; (S.S.); (L.S.)
Journal: Bioengineering (Basel, Switzerland), volume 13, issue 8, article 924
Dates: received 16 July 2026; accepted 11 August 2026; published online 14 August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.3390/bioengineering13080924 · PMID 42649812 · PMCID PMC13509232 · OpenAlex W7202368435
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: other (modality), human (organism), pain (population), methods / tools (subfield)
Methods: Spectral & time-frequency, Connectivity, Machine learning, Preprocessing, Statistics
Keywords: deep learning, gait analysis, gait event detection, non-specific low back pain, smartphone IMU, transfer learning
Topic: Balance, Gait, and Falls Prevention (Physical Therapy, Sports Therapy and Rehabilitation, Health Professions), according to OpenAlex
Funding: European Regional Development Fund (ZS/2023/11/181945)
Citations: not cited yet (Europe PMC); 34 references in the paper

Abstract

Accurate gait event detection using inertial measurement units (IMUs) is essential for temporal gait analysis, but frame-level detection is challenged by sparse initial contact (IC) and foot-off (FO) events. This study evaluated recurrent neural network architectures and training strategies for simultaneous IC and FO detection using a single shank-mounted smartphone IMU. The internal dataset included 28 healthy older adults and 18 individuals with non-specific low back pain (NSLBP). Temporal label expansion substantially improved validation performance for gated recurrent unit (GRU) and long short-term memory models, whereas point-label and class-weighted training performed poorly. The selected label-expanded GRU (LE-GRU) achieved F1 scores above 0.95 for both events and mean absolute temporal errors below 12 ms on held-out internal test folds, with high performance in both cohorts. On an external dataset with different sensor and acquisition characteristics, high performance required full-network fine-tuning, indicating the need for adaptation across datasets. Stance phase and stride time calculated from LE-GRU-predicted events showed high agreement with reference-derived values, with Lin’s concordance correlation coefficients from 0.980 to 0.994. These findings demonstrate that temporal label expansion enables accurate GRU-based gait event detection and temporal gait analysis from data collected with a single smartphone IMU.

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 4 matches between paragraphs and lines of code.

MS-AI-OVGU/gait_ml

License: other
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 3c1aed01345c63d395565a5ae21465c877d9591e, 27 August 2026
Languages: Python (13), Jupyter (9)
Size: 38 files, 22 scripts
Software Heritage: not archived
Found in: “Data Availability Statement”
Holds: README, license file, environment (pyproject.toml), 9 notebooks
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (8 files), pandas (8 files), PyTorch Lightning (7 files), PyTorch (7 files), Matplotlib (6 files), scikit-learn (6 files), SciPy (6 files), Plotly (4 files), seaborn (3 files)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
10 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;
  • 8 scripts, each with its path and the digest of its content;
  • 4 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

Datasets cited

Data Availability Statement

The internal data supporting the findings of this study is deposited in Zenodo and available at https://doi.org/10.5281/zenodo.17899477 (accessed 10 August 2026). The source code is available at https://github.com/MS-AI-OVGU/gait_ml (accessed 10 August 2026). The analyses were performed using Python 3.10.18, and the repository includes the package dependencies and version information required to reproduce the analyses. The source code is released under the PolyForm Noncommercial License 1.0.0 for noncommercial research use. The external dataset is publicly available as described in [28].

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, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 6 authors, 6 keywords, 1 funder, 30 references.

Cite

This paper

Vivar, G., Singh, S., Bea, T., Saal, C., Munoz-Martel, V., & Schega, L. (2026). Deep Learning-Based Temporal Gait Analysis Using a Smartphone IMU in Older Adults with and Without Non-Specific Low Back Pain. Bioengineering (Basel, Switzerland), 13(8), 924. https://doi.org/10.3390/bioengineering13080924

BibTeX

@article{vivar2026deep,
author = {Vivar, Gerome and Singh, Shivam and Bea, Tobias and Saal, Christian and Munoz-Martel, Victor and Schega, Lutz},
title = {{Deep Learning-Based Temporal Gait Analysis Using a Smartphone IMU in Older Adults with and Without Non-Specific Low Back Pain}},
journal = {Bioengineering (Basel, Switzerland)},
year = {2026},
month = aug,
volume = {13},
number = {8},
pages = {924},
publisher = {Multidisciplinary Digital Publishing Institute (MDPI)},
issn = {2306-5354},
doi = {10.3390/bioengineering13080924},
url = {https://doi.org/10.3390/bioengineering13080924},
pmid = {42649812},
pmcid = {PMC13509232}
}

RIS

TY - JOUR
AU - Vivar, Gerome
AU - Singh, Shivam
AU - Bea, Tobias
AU - Saal, Christian
AU - Munoz-Martel, Victor
AU - Schega, Lutz
TI - Deep Learning-Based Temporal Gait Analysis Using a Smartphone IMU in Older Adults with and Without Non-Specific Low Back Pain
T2 - Bioengineering (Basel, Switzerland)
J2 - Bioengineering (Basel)
PY - 2026
DA - 2026/08/14
VL - 13
IS - 8
SP - 924
SN - 2306-5354
PB - Multidisciplinary Digital Publishing Institute (MDPI)
DO - 10.3390/bioengineering13080924
UR - https://doi.org/10.3390/bioengineering13080924
LA - en
ER -

CSL-JSON

{
"id": "10.3390/bioengineering13080924",
"type": "article-journal",
"title": "Deep Learning-Based Temporal Gait Analysis Using a Smartphone IMU in Older Adults with and Without Non-Specific Low Back Pain",
"container-title": "Bioengineering (Basel, Switzerland)",
"author": [
{
"family": "Vivar",
"given": "Gerome"
},
{
"family": "Singh",
"given": "Shivam"
},
{
"family": "Bea",
"given": "Tobias"
},
{
"family": "Saal",
"given": "Christian"
},
{
"family": "Munoz-Martel",
"given": "Victor"
},
{
"family": "Schega",
"given": "Lutz"
}
],
"container-title-short": "Bioengineering (Basel)",
"volume": "13",
"issue": "8",
"page": "924",
"DOI": "10.3390/bioengineering13080924",
"PMID": "42649812",
"PMCID": "PMC13509232",
"ISSN": "2306-5354",
"publisher": "Multidisciplinary Digital Publishing Institute (MDPI)",
"URL": "https://doi.org/10.3390/bioengineering13080924",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
14
]
]
}
}

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.1038/s41597-026-07077-7 [code]
Everyday Activity Science and Engineering Table Setting Dataset.
Journal: Scientific data
In common: PyTorch Lightning, PyTorch, seaborn, 5 other tools, other, methods / tools, 1 reference
[2] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: PyTorch Lightning, Plotly, PyTorch, 6 other tools
[3] doi:10.1038/s41467-026-75455-1 [code]
Shared latent representations of speech production for cross-patient speech decoding.
Journal: Nature communications
In common: PyTorch Lightning, Plotly, PyTorch, 6 other tools
[4] doi:10.1038/s41467-026-73996-z [code]
Genetic architecture of white matter microstructure captured by unsupervised deep representation learning of fractional anisotropy maps.
Journal: Nature communications
In common: PyTorch Lightning, Plotly, PyTorch, 6 other tools
[5] doi:10.1523/eneuro.0023-26.2026 [code]
Real-Time Segmentation and Classification of Birdsong Syllables for Learning Experiments.
Journal: eNeuro
In common: PyTorch Lightning, Plotly, PyTorch, 6 other tools
[6] doi:10.1093/bioinformatics/btag169 [code]
Bidirectional Mamba-2 boosts EEG super-resolution via regression and diffusion.
Journal: Bioinformatics (Oxford, England)
In common: PyTorch Lightning, Plotly, PyTorch, 5 other tools, methods / tools
[7] doi:10.1038/s41467-026-72253-7 [code]
Spurious alignment between large language models and brains can emerge from non-robust methods and overlooked confounds.
Journal: Nature communications
In common: Plotly, PyTorch, seaborn, 5 other tools, methods / tools, 1 reference
[8] doi:10.1007/s12021-026-09817-x [code]
Circle of Willis-Guided Localization for Simultaneous Detection and Classification of Large Vessel Occlusions in Brain CTA.
Journal: Neuroinformatics
In common: PyTorch Lightning, PyTorch, seaborn, 5 other tools, other
[9] doi:10.1038/s41598-026-68186-2 [code]
NeuroStream: spectral-spatio-temporal deep learning for visual stimulus classification from EEG.
Journal: Scientific reports
In common: PyTorch, seaborn, scikit-learn, 4 other tools, methods / tools, 2 references
[10] doi:10.1038/s41467-026-76939-w [code]
HIPPIE: a generative model for electrophysiological analysis across species, technologies, and modalities.
Journal: Nature communications
In common: PyTorch Lightning, PyTorch, seaborn, 5 other tools, methods / 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.