OSCR

Byzantine robust federated learning for heterogeneous brain MRI using multisignal gradient fingerprinting and adaptive trust aggregation.

Code ↔ Paper

32 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 32 matches
  1. [1] § Results › Multi-seed validation, failure-mode analysis, and overhead comparison ↔ run_r3_experiments.py, lines 1–60 · score 0.92 · headline accuracy, failure modes, adaptive attacker, attack budgets, SignGuard, R3 experiments
  2. [2] § Results › Multi-seed validation, failure-mode analysis, and overhead comparison › Summary of R3 additional analyses ↔ run_r3_experiments.py, lines 1–60 · score 0.81 · adaptive attacker, R3 experiments, attack budgets, FLTrust, multi seed, FedAvg
  3. [3] § Results › Optimizer comparison ↔ run_r2_experiments.py, lines 1000–1010 · score 0.80 · FedDWA, FedNova, FedADMM, FedProx, FedBN, FedAvg
  4. [4] § Results › Component ablation analysis ↔ run_r2_experiments.py, lines 620–747 · score 0.80 · Acc r13, Acc r25, dynamic attack schedule, 13–25, phase transition, 1–12
  5. [5] § Results › Multi-seed validation, failure-mode analysis, and overhead comparison › Multi-seed statistical validation with paired tests (R3-C) ↔ run_all_experiments.py, lines 1–49 · score 0.78 · Wilcoxon signed rank, confidence interval, FLTrust, FedBN, FedAvg, Cohen
  6. [6] § Results › Robustness under dynamic attack schedule ↔ run_r2_experiments.py, lines 620–747 · score 0.76 · Malicious clients behave, phase transition, attack schedules, 1–12, switch, dynamic
  7. [7] § Results › Multi-seed validation, failure-mode analysis, and overhead comparison › Additional Byzantine-robust baselines (R3-F) ↔ federated_learning/training/server.py, lines 1049–1115 · score 0.76 · SignGuard, robust aggregators, FLTrust, FedBN, FedAvg, trust scoring
  8. [8] § Results › Multi-seed validation, failure-mode analysis, and overhead comparison › Computational and communication overhead (R3-B) ↔ federated_learning/config/config.py, lines 233–236 · score 0.72 · peak GPU memory, wall clock, Communication volume, MB, overhead, R3
  9. [9] § Methods › Six-dimensional gradient fingerprinting ↔ federated_learning/utils/shapley_utils.py, lines 300–368 · score 0.72 · Monte Carlo Shapley, marginal contribution, random permutations, clients, gradient
  10. [10] § Results › Multi-seed validation, failure-mode analysis, and overhead comparison › Adaptive attacker evaluation (R3-D) ↔ federated_learning/training/client.py, lines 103–160 · score 0.72 · inner loop, malicious direction, VAE aware, radius, snapshot, PGD
  11. [11] § Results › Robustness under dynamic attack schedule ↔ federated_learning/training/server.py, lines 824–874 · score 0.70 · linear combination, fixed threshold, equal weight, fingerprint features, dual attention, R3
  12. [12] § Results › Component ablation analysis ↔ run_all_experiments.py, lines 1246–1330 · score 0.70 · Removing RL adaptation, Removing VAE fingerprinting, removing Shapley, computation, component, ablation
  13. [13] § Results ↔ run_all_experiments.py, lines 1–49 · score 0.69 · confidence intervals, OASIS clinical, Scalability evaluation, baseline comparison, adversarial, std
  14. [14] § Results › Multi-seed validation, failure-mode analysis, and overhead comparison › Summary of R3 additional analyses ↔ federated_learning/config/config.py, lines 233–236 · score 0.69 · peak GPU memory, wall clock, communication volume, overhead, R3, aggregator
  15. [15] § Methods › FedBN-P optimizer: masked proximal hybrid ↔ federated_learning/utils/model_utils.py, lines 14–128 · score 0.66 · BatchNorm, proximal term, BN parameters, FedBN, optimizers
  16. [16] § Results › Component ablation analysis ↔ main.py, lines 510–566 · score 0.65 · detection accuracy, detection recall, detection metrics, Detection precision, noise, scoring
  17. [17] § Results › Multi-seed validation, failure-mode analysis, and overhead comparison › Additional Byzantine-robust baselines (R3-F) ↔ federated_learning/aggregators/byzantine_baselines.py, lines 184–269 · score 0.65 · norm filter, sign cluster, SignGuard, median, baseline, Byzantine
  18. [18] § Results › Experimental configuration ↔ run_optimized_experiments.py, lines 1–53 · score 0.65 · minority class, Class imbalance, real clinical, severity, OASIS, training
  19. [19] § Results › Component ablation analysis ↔ run_all_experiments.py, lines 1246–1330 · score 0.64 · Removing VAE, Removing RL, component ablation, real clinical, fingerprint, metrics
  20. [20] § Results › Optimizer comparison ↔ run_r2_experiments.py, lines 1000–1010 · score 0.63 · FedNova, FedProx, FedBN, FedAvg, SCAFFOLD, Optimizer
  21. [21] § Methods › FedBN-P optimizer: masked proximal hybrid ↔ federated_learning/training/server.py, lines 1049–1115 · score 0.61 · BatchNorm, proximal term, FedBN, hybrid, Byzantine, RL
  22. [22] § Results › Experimental configuration ↔ federated_learning/data/cifar_dataset.py, lines 9–64 · score 0.60 · random horizontal flips, random crops, augmented, CIFAR, training
  23. [23] § Methods › Implementation details ↔ federated_learning/models/resnet.py, lines 7–56 · score 0.60 · fully connected, ReLU, ResNet, dropout, layer, class
  24. [24] § Results › Multi-seed validation, failure-mode analysis, and overhead comparison › Additional Byzantine-robust baselines (R3-F) ↔ federated_learning/aggregators/byzantine_baselines.py, lines 184–269 · score 0.60 · sign clustering filtering, SignGuard, distance, median, baseline, Byzantine
  25. [25] § Results › Robustness under dynamic attack schedule ↔ federated_learning/config/config.py, lines 228–231 · score 0.60 · linear combination, equal weight, fingerprint features, dual attention, R3, ablation
  26. [26] § Methods › FedBN-P optimizer: masked proximal hybrid ↔ federated_learning/utils/model_utils.py, lines 14–128 · score 0.59 · BatchNorm, BN parameters, model parameters, proximal, optimizer
  27. [27] § Results › Experimental configuration ↔ federated_learning/utils/data_utils.py, lines 15–106 · score 0.57 · random horizontal flips, random crops, CIFAR, MNIST, training
  28. [28] § Results › Multi-seed validation, failure-mode analysis, and overhead comparison › Failure-mode characterization under varying attack intensity (R3-A) ↔ run_r3_experiments.py, lines 388–446 · score 0.56 · attack budgets, malicious fraction, sign flip, sweeps, intensity, failure
  29. [29] § Methods › FedBN-P optimizer: masked proximal hybrid ↔ federated_learning/training/aggregation.py, lines 209–254 · score 0.56 · proximal term, BN parameters, FedBN, masked, client
  30. [30] § Results › Component ablation analysis ↔ run_r2_experiments.py, lines 349–392 · score 0.55 · VAE fingerprinting component, Gaussian noise injection, ablation, validate, seed, rounds
  31. [31] § Results › Component ablation analysis ↔ run_revision_experiments.py, lines 164–212 · score 0.53 · Peer consensus, L2 norm, dual attention, ratio, cosine, fingerprint
  32. [32] § Results › Multi-seed validation, failure-mode analysis, and overhead comparison › Computational and communication overhead (R3-B) ↔ federated_learning/training/server.py, lines 1157–1239 · score 0.51 · wall clock, MB, cache, volume, peak, overhead

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

Python · 1,331 lines · 54 KB · MIT · 5 matches

  1. #!/usr/bin/env python3
  2. """
  3. =============================================================================
  4. OptiGradTrust -- Second Revision (R2) Experiment Runner
  5. =============================================================================
  6. Implements ALL experiments required by Reviewer 1 in the second revision:
  7. Experiment 1A: Multi-attack ablation under GAUSSIAN NOISE INJECTION
  8. (validates VAE component)
  9. Experiment 1B: Multi-attack ablation under SIGN-FLIPPING
  10. (validates cosine-similarity / sign-consistency)
  11. Experiment 2: Dynamic attack schedule (RL justification) -- CRITICAL
  12. Phase transition: rounds 1-12 benign, 13-25 scaling x20
  13. Experiment 3: Trust score / Shapley visualization (qualitative figure)
  14. Experiment 4: Optimizer comparison under adversarial conditions
  15. USAGE:
  16. python run_r2_experiments.py --experiment exp1a # Noise ablation
  17. python run_r2_experiments.py --experiment exp1b # Sign-flip ablation
  18. python run_r2_experiments.py --experiment exp2 # Dynamic attack (CRITICAL)
  19. python run_r2_experiments.py --experiment exp3 # Trust visualization
  20. python run_r2_experiments.py --experiment exp4 # Optimizer adversarial
  21. python run_r2_experiments.py --experiment all # Run all
  22. python run_r2_experiments.py --experiment exp2 --dry-run # Sanity check
  23. # Multi-seed options (use for critical experiments):
  24. python run_r2_experiments.py --experiment exp2 --seeds 42 123 456
  25. python run_r2_experiments.py --experiment exp1a --seeds 42 123 456 789 1024
  26. Author: OptiGradTrust Team
  27. =============================================================================
  28. """
  29. import matplotlib
  30. matplotlib.use('Agg')
  31. import os
  32. import sys
  33. import json
  34. import time
  35. import copy
  36. import csv
  37. import argparse
  38. import traceback
  39. import warnings
  40. import numpy as np
  41. import torch
  42. import torch.nn.functional as F
  43. import matplotlib.pyplot as plt
  44. import matplotlib.patches as mpatches
  45. from contextlib import contextmanager
  46. from datetime import datetime
  47. from typing import Dict, List, Optional, Tuple
  48. warnings.filterwarnings('ignore')
  49. PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__))
  50. sys.path.insert(0, PROJECT_ROOT)
  51. from run_all_experiments import (
  52. set_all_seeds, configure_for_dataset,
  53. compute_statistics, Logger, _sync_config_to_modules
  54. )
  55. # ---------------------------------------------------------------------------
  56. # Output directory & logger
  57. # ---------------------------------------------------------------------------
  58. RESULTS_DIR = os.path.join(PROJECT_ROOT, 'results', 'r2_revision')
  59. os.makedirs(RESULTS_DIR, exist_ok=True)
  60. LOG_FILE = os.path.join(
  61. RESULTS_DIR,
  62. f'r2_log_{datetime.now().strftime("%Y%m%d_%H%M%S")}.txt'
  63. )
  64. logger = Logger(LOG_FILE)
  65. # ===========================================================================
  66. # SHARED ABLATION CONFIGURATIONS
  67. # ===========================================================================
  68. # The five standard ablation configurations (same structure as Table 9)
  69. ABLATION_CONFIGS = [
  70. {'name': 'full', 'vae': True, 'shapley': True, 'dual_attention': True, 'rl': True},
  71. {'name': 'no_vae', 'vae': False, 'shapley': True, 'dual_attention': True, 'rl': True},
  72. {'name': 'no_shapley', 'vae': True, 'shapley': False, 'dual_attention': True, 'rl': True},
  73. {'name': 'no_rl', 'vae': True, 'shapley': True, 'dual_attention': True, 'rl': False},
  74. {'name': 'fedavg_no_def', 'vae': False, 'shapley': False, 'dual_attention': False, 'rl': False,
  75. '_fedavg': True},
  76. ]
  77. # Extra sign-flipping specific configs (optional if time allows)
  78. SIGNFLIP_EXTRA_CONFIGS = [
  79. {'name': 'no_cosine_sim', 'vae': True, 'shapley': True, 'dual_attention': True, 'rl': True,
  80. '_zero_features': [1, 2]}, # zero root-cosine (1) and peer-consensus (2)
  81. {'name': 'no_sign_consist', 'vae': True, 'shapley': True, 'dual_attention': True, 'rl': True,
  82. '_zero_features': [4]}, # zero sign-consistency (4)
  83. ]
  84. # ===========================================================================
  85. # FEATURE-ZEROING PATCH (same technique as run_revision_experiments.py)
  86. # ===========================================================================
  87. @contextmanager
  88. def _patch_feature_zeroing(zero_indices):
  89. """Monkey-patch Server._compute_all_gradient_features to zero specific feature columns."""
  90. from federated_learning.training.server import Server
  91. original_fn = Server._compute_all_gradient_features
  92. def _patched(self, client_gradients):
  93. features = original_fn(self, client_gradients)
  94. if isinstance(features, torch.Tensor):
  95. for idx in zero_indices:
  96. if idx < features.size(1):
  97. features[:, idx] = 0.5 # neutral / uninformative
  98. return features
  99. Server._compute_all_gradient_features = _patched
  100. try:
  101. yield
  102. finally:
  103. Server._compute_all_gradient_features = original_fn
  104. # ===========================================================================
  105. # LOW-LEVEL EXPERIMENT RUNNER (shared by Exp 1A, 1B)
  106. # ===========================================================================
  107. def _apply_alzheimer_ablation_config(cfg_mod, ablation_cfg, num_rounds):
  108. """
  109. Set Alzheimer ablation config parameters exactly as described in the R2 report:
  110. - 10 clients, 40% malicious, Dirichlet alpha=0.1, 25 rounds
  111. - Batch=16, LR=1e-4, weight_decay=5e-5, FedBN-P (mu=0.01)
  112. """
  113. cfg_mod.GLOBAL_EPOCHS = num_rounds
  114. cfg_mod.FRACTION_MALICIOUS = 0.4
  115. cfg_mod.NUM_MALICIOUS = 4
  116. cfg_mod.BATCH_SIZE = 16
  117. cfg_mod.LR = 1e-4
  118. cfg_mod.LEARNING_RATE = 1e-4
  119. cfg_mod.WEIGHT_DECAY = 5e-5
  120. cfg_mod.FEDPROX_MU = 0.01
  121. is_fedavg_no_def = ablation_cfg.get('_fedavg', False)
  122. if is_fedavg_no_def:
  123. cfg_mod.ENABLE_VAE = False
  124. cfg_mod.ENABLE_SHAPLEY = False
  125. cfg_mod.ENABLE_DUAL_ATTENTION = False
  126. cfg_mod.RL_AGGREGATION_METHOD = 'dual_attention'
  127. cfg_mod.RL_WARMUP_ROUNDS = 9999
  128. cfg_mod.AGGREGATION_METHOD = 'fedavg'
  129. cfg_mod.GRADIENT_COMBINATION_METHOD = 'fedavg'
  130. else:
  131. cfg_mod.ENABLE_VAE = ablation_cfg.get('vae', True)
  132. cfg_mod.ENABLE_SHAPLEY = ablation_cfg.get('shapley', True)
  133. cfg_mod.ENABLE_DUAL_ATTENTION = ablation_cfg.get('dual_attention', True)
  134. if not ablation_cfg.get('rl', True):
  135. cfg_mod.RL_AGGREGATION_METHOD = 'dual_attention'
  136. cfg_mod.RL_WARMUP_ROUNDS = 9999
  137. else:
  138. cfg_mod.RL_AGGREGATION_METHOD = 'hybrid'
  139. cfg_mod.RL_WARMUP_ROUNDS = 5
  140. cfg_mod.RL_RAMP_UP_ROUNDS = 10
  141. cfg_mod.AGGREGATION_METHOD = 'fedbn_fedprox'
  142. cfg_mod.GRADIENT_COMBINATION_METHOD = 'fedbn_fedprox'
  143. def _run_ablation_single(attack_type, ablation_cfg, seed, num_rounds, sigma=15.0):
  144. """
  145. Run one ablation configuration for a given attack type and seed.
  146. Returns a result dict.
  147. """
  148. import federated_learning.config.config as cfg_mod
  149. from federated_learning.training.server import Server
  150. from federated_learning.training.client import Client
  151. from federated_learning.data.dataset_utils import load_dataset, create_client_datasets
  152. from federated_learning.utils.model_utils import set_random_seeds
  153. cfg_name = ablation_cfg.get('name', 'unknown')
  154. logger.info(f" [{cfg_name}] seed={seed} attack={attack_type}")
  155. set_all_seeds(seed)
  156. configure_for_dataset(
  157. 'ALZHEIMER', num_clients=10,
  158. non_iid_config={'enable': True, 'type': 'dirichlet', 'alpha': 0.1},
  159. aggregation_method='fedbn_fedprox'
  160. )
  161. cfg_mod.RANDOM_SEED = seed
  162. _apply_alzheimer_ablation_config(cfg_mod, ablation_cfg, num_rounds)
  163. _sync_config_to_modules()
  164. set_random_seeds(seed)
  165. start_t = time.time()
  166. try:
  167. train_dataset, test_dataset = load_dataset()
  168. root_loader = torch.utils.data.DataLoader(
  169. train_dataset, batch_size=cfg_mod.BATCH_SIZE, shuffle=True, num_workers=0)
  170. server = Server()
  171. server.set_datasets(root_loader, test_dataset)
  172. server._pretrain_global_model()
  173. _, client_datasets = create_client_datasets(
  174. train_dataset=train_dataset, num_clients=10, iid=False, alpha=0.1)
  175. malicious_indices = np.random.choice(10, 4, replace=False)
  176. clients = []
  177. for i in range(10):
  178. is_mal = int(i) in malicious_indices.tolist()
  179. c = Client(client_id=i, dataset=client_datasets[i], is_malicious=is_mal)
  180. if is_mal:
  181. c.set_attack_parameters(
  182. attack_type=attack_type,
  183. scaling_factor=20.0,
  184. sigma=sigma,
  185. partial_percent=cfg_mod.PARTIAL_SCALING_PERCENT,
  186. noise_factor=cfg_mod.NOISE_FACTOR,
  187. flip_probability=cfg_mod.FLIP_PROBABILITY,
  188. )
  189. clients.append(c)
  190. server.add_clients(clients)
  191. root_gradients = server._collect_root_gradients()
  192. server.root_gradients = root_gradients # enables root-cosine & sign-consistency features
  193. if getattr(cfg_mod, 'ENABLE_VAE', True):
  194. server.vae = server.train_vae(root_gradients, vae_epochs=cfg_mod.VAE_EPOCHS)
  195. else:
  196. server.vae = None
  197. test_loader = torch.utils.data.DataLoader(
  198. test_dataset, batch_size=cfg_mod.BATCH_SIZE, shuffle=False, num_workers=0)
  199. server.test_loader = test_loader
  200. # Run training, applying optional feature-zeroing patch
  201. zero_feats = ablation_cfg.get('_zero_features')
  202. if zero_feats:
  203. with _patch_feature_zeroing(zero_feats):
  204. _, round_metrics = server.train(num_rounds=cfg_mod.GLOBAL_EPOCHS)
  205. else:
  206. _, round_metrics = server.train(num_rounds=cfg_mod.GLOBAL_EPOCHS)
  207. final_acc = server.evaluate_model()
  208. # Aggregate detection metrics over all rounds
  209. total_tp = total_fp = total_fn = total_tn = 0
  210. for rd in round_metrics.values():
  211. det = rd.get('detection_results', {})
  212. total_tp += det.get('true_positives', 0)
  213. total_fp += det.get('false_positives', 0)
  214. total_fn += det.get('false_negatives', 0)
  215. total_tn += det.get('true_negatives', 0)
  216. precision = total_tp / (total_tp + total_fp) if (total_tp + total_fp) > 0 else 0.0
  217. recall = total_tp / (total_tp + total_fn) if (total_tp + total_fn) > 0 else 0.0
  218. f1 = (2 * precision * recall / (precision + recall)
  219. if (precision + recall) > 0 else 0.0)
  220. result = {
  221. 'config_name': cfg_name,
  222. 'seed': seed,
  223. 'attack_type': attack_type,
  224. 'final_accuracy': float(final_acc),
  225. 'precision': float(precision),
  226. 'recall': float(recall),
  227. 'f1': float(f1),
  228. 'total_time': time.time() - start_t,
  229. 'status': 'completed',
  230. }
  231. logger.success(
  232. f" → Acc={final_acc*100:.2f}% "
  233. f"Prec={precision*100:.2f} Rec={recall*100:.2f} F1={f1*100:.2f}"
  234. )
  235. return result
  236. except Exception as exc:
  237. logger.error(f" → FAILED: {str(exc)[:120]}")
  238. traceback.print_exc()
  239. return {
  240. 'config_name': cfg_name,
  241. 'seed': seed,
  242. 'attack_type': attack_type,
  243. 'status': 'failed',
  244. 'error': str(exc),
  245. 'total_time': time.time() - start_t,
  246. }
  247. def _save_ablation_table(results, filename_base, attack_label):
  248. """Save CSV + print a formatted text table for ablation results."""
  249. completed = [r for r in results if r.get('status') == 'completed']
  250. # CSV
  251. csv_path = os.path.join(RESULTS_DIR, f'{filename_base}.csv')
  252. fields = ['config_name', 'seed', 'attack_type',
  253. 'final_accuracy', 'precision', 'recall', 'f1', 'total_time', 'status']
  254. with open(csv_path, 'w', newline='') as f:
  255. w = csv.DictWriter(f, fieldnames=fields, extrasaction='ignore')
  256. w.writeheader()
  257. w.writerows(results)
  258. logger.info(f" Saved CSV → {csv_path}")
  259. # JSON
  260. json_path = os.path.join(RESULTS_DIR, f'{filename_base}.json')
  261. with open(json_path, 'w') as f:
  262. json.dump(results, f, indent=2, default=str)
  263. logger.info(f" Saved JSON → {json_path}")
  264. # Pretty table
  265. header = (f"\n{'='*80}\n"
  266. f"ABLATION TABLE — {attack_label}\n"
  267. f"{'Configuration':<26} {'Acc (%)':>9} {'ΔAcc':>7} "
  268. f"{'Recall (%)':>11} {'F1 (%)':>8}\n"
  269. f"{'-'*80}")
  270. logger.info(header)
  271. # Compute full-system accuracy as baseline for delta
  272. full_acc = None
  273. for r in completed:
  274. if r['config_name'] == 'full':
  275. full_acc = r['final_accuracy']
  276. break
  277. for r in completed:
  278. acc = r['final_accuracy'] * 100
  279. delta = ((r['final_accuracy'] - full_acc) * 100
  280. if full_acc is not None and r['config_name'] != 'full' else 0.0)
  281. delta_str = f"{delta:+.2f}" if r['config_name'] != 'full' else " --"
  282. is_fedavg = 'fedavg' in r['config_name']
  283. recall_str = f"{r['recall']*100:.2f}" if not is_fedavg else " --"
  284. f1_str = f"{r['f1']*100:.2f}" if not is_fedavg else " --"
  285. logger.info(
  286. f" {r['config_name']:<24} {acc:>9.2f} {delta_str:>7} "
  287. f"{recall_str:>11} {f1_str:>8}"
  288. )
  289. logger.info('=' * 80)
  290. return csv_path, json_path
  291. # ===========================================================================
  292. # EXPERIMENT 1A — GAUSSIAN NOISE INJECTION ABLATION
  293. # ===========================================================================
  294. def run_exp1a_noise_ablation(seeds=None, dry_run=False):
  295. """
  296. Experiment 1A: Ablation under Gaussian Noise Injection (sigma=15.0).
  297. Validates the VAE fingerprinting component.
  298. """
  299. logger.info('=' * 70)
  300. logger.info('EXPERIMENT 1A: ABLATION UNDER GAUSSIAN NOISE INJECTION')
  301. logger.info(' Dataset: Alzheimer MRI | Clients: 10 | Malicious: 40%')
  302. logger.info(' Attack: gaussian_noise_injection sigma=15.0 | alpha=0.1 | 25 rounds')
  303. logger.info('=' * 70)
  304. if seeds is None:
  305. seeds = [42]
  306. if dry_run:
  307. seeds = [42]
  308. num_rounds = 2 if dry_run else 25
  309. all_results = []
  310. configs_to_run = ABLATION_CONFIGS.copy()
  311. for cfg in configs_to_run:
  312. logger.info(f"\n Config: {cfg['name']}")
  313. seed_results = []
  314. for seed in seeds:
  315. r = _run_ablation_single(
  316. attack_type='gaussian_noise_injection',
  317. ablation_cfg=cfg,
  318. seed=seed,
  319. num_rounds=num_rounds,
  320. sigma=15.0,
  321. )
  322. all_results.append(r)
  323. if r['status'] == 'completed':
  324. seed_results.append(r)
  325. if len(seed_results) > 1:
  326. accs = [r['final_accuracy'] for r in seed_results]
  327. logger.info(f" Mean Acc: {np.mean(accs)*100:.2f}% ± {np.std(accs)*100:.2f}%")
  328. _save_ablation_table(all_results, 'exp1a_noise_ablation',
  329. 'Gaussian Noise Injection (sigma=15.0) | Alzheimer MRI')
  330. logger.info('\nExperiment 1A complete.')
  331. return all_results
  332. # ===========================================================================
  333. # EXPERIMENT 1B — SIGN-FLIPPING ABLATION
  334. # ===========================================================================
  335. def run_exp1b_signflip_ablation(seeds=None, dry_run=False, include_optional=True):
  336. """
  337. Experiment 1B: Ablation under Sign-Flipping (lambda=-1).
  338. Validates cosine-similarity and sign-consistency components.
  339. """
  340. logger.info('=' * 70)
  341. logger.info('EXPERIMENT 1B: ABLATION UNDER SIGN-FLIPPING (lambda=-1)')
  342. logger.info(' Dataset: Alzheimer MRI | Clients: 10 | Malicious: 40%')
  343. logger.info(' Attack: sign_flipping_attack | alpha=0.1 | 25 rounds')
  344. logger.info('=' * 70)
  345. if seeds is None:
  346. seeds = [42]
  347. if dry_run:
  348. seeds = [42]
  349. num_rounds = 2 if dry_run else 25
  350. all_results = []
  351. configs_to_run = ABLATION_CONFIGS.copy()
  352. if include_optional and not dry_run:
  353. configs_to_run += SIGNFLIP_EXTRA_CONFIGS
  354. for cfg in configs_to_run:
  355. logger.info(f"\n Config: {cfg['name']}")
  356. seed_results = []
  357. for seed in seeds:
  358. r = _run_ablation_single(
  359. attack_type='sign_flipping_attack',
  360. ablation_cfg=cfg,
  361. seed=seed,
  362. num_rounds=num_rounds,
  363. )
  364. all_results.append(r)
  365. if r['status'] == 'completed':
  366. seed_results.append(r)
  367. if len(seed_results) > 1:
  368. accs = [r['final_accuracy'] for r in seed_results]
  369. logger.info(f" Mean Acc: {np.mean(accs)*100:.2f}% ± {np.std(accs)*100:.2f}%")
  370. _save_ablation_table(all_results, 'exp1b_signflip_ablation',
  371. 'Sign-Flipping (lambda=-1) | Alzheimer MRI')
  372. logger.info('\nExperiment 1B complete.')
  373. return all_results
  374. # ===========================================================================
  375. # EXPERIMENT 2 — DYNAMIC ATTACK SCHEDULE (RL JUSTIFICATION)
  376. # ===========================================================================
  377. EXP2_CONFIGS = [
  378. {'name': 'full', 'vae': True, 'shapley': True, 'dual_attention': True, 'rl': True},
  379. {'name': 'no_rl', 'vae': True, 'shapley': True, 'dual_attention': True, 'rl': False},
  380. {'name': 'fedavg_no_def', 'vae': False, 'shapley': False, 'dual_attention': False, 'rl': False,
  381. '_fedavg': True},
  382. ]
  383. def _run_exp2_single(cfg, seed, num_rounds, phase_transition_round=12):
  384. """
  385. Run one configuration of Experiment 2 (dynamic attack, phase-transition schedule).
  386. Phase transition:
  387. rounds 1–12 (round_idx 0–11): malicious clients behave BENIGNLY
  388. rounds 13–25 (round_idx 12–24): malicious clients launch scaling x20
  389. Returns dict with per-round accuracy + detection, and final summary metrics.
  390. """
  391. import federated_learning.config.config as cfg_mod
  392. from federated_learning.training.server import Server
  393. from federated_learning.training.client import Client
  394. from federated_learning.data.dataset_utils import load_dataset, create_client_datasets
  395. from federated_learning.utils.model_utils import set_random_seeds
  396. cfg_name = cfg.get('name', 'unknown')
  397. logger.info(f" [Exp2 | {cfg_name}] seed={seed}")
  398. set_all_seeds(seed)
  399. configure_for_dataset(
  400. 'ALZHEIMER', num_clients=10,
  401. non_iid_config={'enable': True, 'type': 'dirichlet', 'alpha': 0.1},
  402. aggregation_method='fedbn_fedprox'
  403. )
  404. cfg_mod.RANDOM_SEED = seed
  405. _apply_alzheimer_ablation_config(cfg_mod, cfg, num_rounds)
  406. # For Exp 2, use RL in ACTIVE training mode (not frozen)
  407. # RL must be actively training during all 25 rounds
  408. if not cfg.get('_fedavg', False) and cfg.get('rl', True):
  409. cfg_mod.RL_AGGREGATION_METHOD = 'hybrid'
  410. cfg_mod.RL_WARMUP_ROUNDS = 5
  411. cfg_mod.RL_RAMP_UP_ROUNDS = 10
  412. _sync_config_to_modules()
  413. set_random_seeds(seed)
  414. start_t = time.time()
  415. try:
  416. train_dataset, test_dataset = load_dataset()
  417. root_loader = torch.utils.data.DataLoader(
  418. train_dataset, batch_size=cfg_mod.BATCH_SIZE, shuffle=True, num_workers=0)
  419. server = Server()
  420. server.set_datasets(root_loader, test_dataset)
  421. server._pretrain_global_model()
  422. _, client_datasets = create_client_datasets(
  423. train_dataset=train_dataset, num_clients=10, iid=False, alpha=0.1)
  424. malicious_indices = np.random.choice(10, 4, replace=False)
  425. clients = []
  426. for i in range(10):
  427. is_mal = int(i) in malicious_indices.tolist()
  428. c = Client(client_id=i, dataset=client_datasets[i], is_malicious=is_mal)
  429. if is_mal:
  430. c.set_attack_parameters(
  431. attack_type='scaling_attack',
  432. scaling_factor=20.0,
  433. partial_percent=cfg_mod.PARTIAL_SCALING_PERCENT,
  434. noise_factor=cfg_mod.NOISE_FACTOR,
  435. flip_probability=cfg_mod.FLIP_PROBABILITY,
  436. )
  437. # Phase-transition schedule: benign for rounds 0..(T-1), then attack
  438. c.dynamic_attack_schedule = 'phase_transition'
  439. c.phase_transition_round = phase_transition_round # 0-indexed
  440. clients.append(c)
  441. server.add_clients(clients)
  442. root_gradients = server._collect_root_gradients()
  443. server.root_gradients = root_gradients # enables root-cosine & sign-consistency features
  444. if getattr(cfg_mod, 'ENABLE_VAE', True):
  445. server.vae = server.train_vae(root_gradients, vae_epochs=cfg_mod.VAE_EPOCHS)
  446. else:
  447. server.vae = None
  448. test_loader = torch.utils.data.DataLoader(
  449. test_dataset, batch_size=cfg_mod.BATCH_SIZE, shuffle=False, num_workers=0)
  450. server.test_loader = test_loader
  451. _, round_metrics = server.train(num_rounds=cfg_mod.GLOBAL_EPOCHS)
  452. final_acc = server.evaluate_model()
  453. # Build per-round metrics (round_metrics[r]['test_accuracy'] = accuracy
  454. # measured BEFORE the update in round r+1, i.e. after completing r rounds)
  455. per_round = {}
  456. for r_idx, rd in round_metrics.items():
  457. acc_before_update = rd.get('test_accuracy', 0.0)
  458. det = rd.get('detection_results', {})
  459. tp = det.get('true_positives', 0)
  460. fp = det.get('false_positives', 0)
  461. fn = det.get('false_negatives', 0)
  462. recall_r = tp / (tp + fn) if (tp + fn) > 0 else 0.0
  463. precision_r = tp / (tp + fp) if (tp + fp) > 0 else 0.0
  464. f1_r = (2 * precision_r * recall_r / (precision_r + recall_r)
  465. if (precision_r + recall_r) > 0 else 0.0)
  466. per_round[int(r_idx)] = {
  467. 'accuracy': float(acc_before_update),
  468. 'detection_recall': float(recall_r),
  469. 'detection_f1': float(f1_r),
  470. 'is_attack_phase': int(r_idx) >= phase_transition_round,
  471. }
  472. # Summary key rounds (1-indexed in paper):
  473. # r13 = just before attack, r18 = 5 rds after, r25 = final
  474. def _get_acc(round_1indexed):
  475. r0 = round_1indexed - 1 # convert to 0-indexed round_metrics key
  476. if r0 in per_round:
  477. return per_round[r0]['accuracy']
  478. return None
  479. def _get_det_f1_range(from_round, to_round):
  480. """Mean detection F1 over a range of rounds (1-indexed)."""
  481. vals = [per_round[r-1]['detection_f1']
  482. for r in range(from_round, to_round+1)
  483. if (r-1) in per_round]
  484. return float(np.mean(vals)) if vals else 0.0
  485. acc_r13 = _get_acc(13)
  486. acc_r18 = _get_acc(18)
  487. acc_r25 = float(final_acc) # after all 25 rounds
  488. det_f1_post_attack = _get_det_f1_range(13, 25)
  489. result = {
  490. 'config_name': cfg_name,
  491. 'seed': seed,
  492. 'per_round': per_round,
  493. 'acc_r13': acc_r13,
  494. 'acc_r18': acc_r18,
  495. 'acc_r25': acc_r25,
  496. 'det_f1_r13_to_r25': det_f1_post_attack,
  497. 'final_accuracy': float(final_acc),
  498. 'total_time': time.time() - start_t,
  499. 'status': 'completed',
  500. 'malicious_indices': malicious_indices.tolist(),
  501. }
  502. def _pct(v):
  503. return f"{v*100:.2f}%" if v is not None else "--"
  504. logger.success(
  505. f" → Acc@r13={_pct(acc_r13)} "
  506. f"Acc@r18={_pct(acc_r18)} "
  507. f"Acc@r25={_pct(acc_r25)} "
  508. f"DetF1(r13-25)={det_f1_post_attack:.3f}"
  509. )
  510. return result
  511. except Exception as exc:
  512. logger.error(f" → FAILED: {str(exc)[:120]}")
  513. traceback.print_exc()
  514. return {
  515. 'config_name': cfg_name,
  516. 'seed': seed,
  517. 'status': 'failed',
  518. 'error': str(exc),
  519. 'total_time': time.time() - start_t,
  520. }
  521. def run_exp2_dynamic_attack(seeds=None, dry_run=False):
  522. """
  523. Experiment 2 (CRITICAL / NON-NEGOTIABLE):
  524. Dynamic attack schedule — phase-transition for RL justification.
  525. Phase transition:
  526. Rounds 1–12: malicious clients behave BENIGNLY (build false trust)
  527. Rounds 13–25: malicious clients switch to SCALING x20
  528. Outputs:
  529. - Summary table (CSV + JSON): acc@r13, acc@r18, acc@r25, detection F1
  530. - Per-round CSV: for convergence figure in the paper
  531. """
  532. logger.info('=' * 70)
  533. logger.info('EXPERIMENT 2 (CRITICAL): DYNAMIC ATTACK — PHASE TRANSITION')
  534. logger.info(' Dataset: Alzheimer MRI | Clients: 10 | Malicious: 40%')
  535. logger.info(' Rounds 1-12: BENIGN → Rounds 13-25: Scaling x20')
  536. logger.info(' Validates: RL temporal adaptivity advantage')
  537. logger.info('=' * 70)
  538. if seeds is None:
  539. seeds = [42, 123, 456]
  540. if dry_run:
  541. seeds = [42]
  542. num_rounds = 3 if dry_run else 25
  543. phase_transition = 1 if dry_run else 12 # 0-indexed transition round
  544. all_results = {} # cfg_name → list of per-seed results
  545. all_per_round = [] # rows for per-round CSV
  546. for cfg in EXP2_CONFIGS:
  547. cfg_name = cfg['name']
  548. all_results[cfg_name] = []
  549. logger.info(f"\n Config: {cfg_name}")
  550. for seed in seeds:
  551. r = _run_exp2_single(cfg, seed, num_rounds, phase_transition)
  552. all_results[cfg_name].append(r)
  553. if r['status'] == 'completed':
  554. for r_idx, rd in r['per_round'].items():
  555. all_per_round.append({
  556. 'round': r_idx + 1, # 1-indexed for paper
  557. 'seed': seed,
  558. 'config': cfg_name,
  559. 'accuracy': rd['accuracy'],
  560. 'detection_recall': rd['detection_recall'],
  561. 'detection_f1': rd['detection_f1'],
  562. 'is_attack_phase': int(rd['is_attack_phase']),
  563. })
  564. # ---------- Per-round CSV ----------
  565. csv_pr_path = os.path.join(RESULTS_DIR, 'exp2_dynamic_per_round.csv')
  566. if all_per_round:
  567. fields_pr = ['round', 'seed', 'config', 'accuracy',
  568. 'detection_recall', 'detection_f1', 'is_attack_phase']
  569. with open(csv_pr_path, 'w', newline='') as f:
  570. w = csv.DictWriter(f, fieldnames=fields_pr)
  571. w.writeheader()
  572. w.writerows(all_per_round)
  573. logger.info(f"\n Saved per-round CSV → {csv_pr_path}")
  574. # ---------- Summary table ----------
  575. logger.info('\n' + '=' * 80)
  576. logger.info('EXPERIMENT 2 SUMMARY TABLE')
  577. logger.info(f"{'Configuration':<22} {'Acc@r13 (%)':>12} {'Acc@r18 (%)':>12} "
  578. f"{'Acc@r25 (%)':>12} {'DetF1 r13-25':>14}")
  579. logger.info('-' * 80)
  580. summary_rows = []
  581. for cfg in EXP2_CONFIGS:
  582. cfg_name = cfg['name']
  583. completed = [r for r in all_results[cfg_name] if r.get('status') == 'completed']
  584. if not completed:
  585. logger.warning(f" {cfg_name}: No completed runs.")
  586. continue
  587. accs13 = [r['acc_r13'] for r in completed if r.get('acc_r13') is not None]
  588. accs18 = [r['acc_r18'] for r in completed if r.get('acc_r18') is not None]
  589. accs25 = [r['acc_r25'] for r in completed]
  590. f1s = [r['det_f1_r13_to_r25'] for r in completed]
  591. def _fmt(vals):
  592. if not vals:
  593. return ' --'
  594. m = np.mean(vals) * 100
  595. s = np.std(vals) * 100
  596. return f"{m:.2f}±{s:.2f}" if len(vals) > 1 else f"{m:.2f}"
  597. is_fedavg = cfg.get('_fedavg', False)
  598. f1_str = ' --' if is_fedavg else _fmt(f1s)
  599. row_str = (f" {cfg_name:<20} {_fmt(accs13):>12} {_fmt(accs18):>12} "
  600. f"{_fmt(accs25):>12} {f1_str:>14}")
  601. logger.info(row_str)
  602. summary_rows.append({
  603. 'config_name': cfg_name,
  604. 'seeds': seeds,
  605. 'acc_r13_mean': float(np.mean(accs13)) if accs13 else None,
  606. 'acc_r13_std': float(np.std(accs13)) if accs13 else None,
  607. 'acc_r18_mean': float(np.mean(accs18)) if accs18 else None,
  608. 'acc_r18_std': float(np.std(accs18)) if accs18 else None,
  609. 'acc_r25_mean': float(np.mean(accs25)) if accs25 else None,
  610. 'acc_r25_std': float(np.std(accs25)) if accs25 else None,
  611. 'det_f1_r13_25_mean': float(np.mean(f1s)) if f1s else None,
  612. 'det_f1_r13_25_std': float(np.std(f1s)) if f1s else None,
  613. })
  614. logger.info('=' * 80)
  615. # Save summary CSV + JSON
  616. csv_sum_path = os.path.join(RESULTS_DIR, 'exp2_dynamic_summary.csv')
  617. if summary_rows:
  618. with open(csv_sum_path, 'w', newline='') as f:
  619. w = csv.DictWriter(f, fieldnames=list(summary_rows[0].keys()),
  620. extrasaction='ignore')
  621. w.writeheader()
  622. w.writerows(summary_rows)
  623. logger.info(f" Saved summary CSV → {csv_sum_path}")
  624. json_path = os.path.join(RESULTS_DIR, 'exp2_dynamic_results.json')
  625. with open(json_path, 'w') as f:
  626. json.dump(all_results, f, indent=2, default=str)
  627. logger.info(f" Saved JSON → {json_path}")
  628. logger.info('\nExperiment 2 complete.')
  629. return all_results, all_per_round
  630. # ===========================================================================
  631. # EXPERIMENT 3 — TRUST SCORE VISUALIZATION
  632. # ===========================================================================
  633. def run_exp3_trust_visualization(seed=42, dry_run=False):
  634. """
  635. Experiment 3: Qualitative trust score visualization.
  636. Runs a full OptiGradTrust experiment under scaling x20 and extracts
  637. per-client, per-round trust scores to generate:
  638. Option A: Trust Score Time Series (PREFERRED — required)
  639. Option B: Trust Score Box Plot
  640. Option C: Shapley Value Distribution (if Shapley enabled)
  641. All at ≥ 300 DPI, plus a raw-data CSV.
  642. """
  643. logger.info('=' * 70)
  644. logger.info('EXPERIMENT 3: TRUST SCORE / SHAPLEY VISUALIZATION')
  645. logger.info(' Dataset: Alzheimer MRI | Attack: Scaling x20')
  646. logger.info(' Configuration: Full OptiGradTrust | Seed=42 | 25 rounds')
  647. logger.info('=' * 70)
  648. import federated_learning.config.config as cfg_mod
  649. from federated_learning.training.server import Server
  650. from federated_learning.training.client import Client
  651. from federated_learning.data.dataset_utils import load_dataset, create_client_datasets
  652. from federated_learning.utils.model_utils import set_random_seeds
  653. num_rounds = 2 if dry_run else 25
  654. set_all_seeds(seed)
  655. configure_for_dataset(
  656. 'ALZHEIMER', num_clients=10,
  657. non_iid_config={'enable': True, 'type': 'dirichlet', 'alpha': 0.1},
  658. aggregation_method='fedbn_fedprox'
  659. )
  660. cfg_mod.RANDOM_SEED = seed
  661. cfg_mod.GLOBAL_EPOCHS = num_rounds
  662. cfg_mod.FRACTION_MALICIOUS = 0.4
  663. cfg_mod.NUM_MALICIOUS = 4
  664. cfg_mod.BATCH_SIZE = 16
  665. cfg_mod.LR = 1e-4
  666. cfg_mod.WEIGHT_DECAY = 5e-5
  667. cfg_mod.ENABLE_VAE = True
  668. cfg_mod.ENABLE_SHAPLEY = True
  669. cfg_mod.ENABLE_DUAL_ATTENTION = True
  670. cfg_mod.RL_AGGREGATION_METHOD = 'hybrid'
  671. cfg_mod.RL_WARMUP_ROUNDS = 5
  672. cfg_mod.RL_RAMP_UP_ROUNDS = 10
  673. cfg_mod.AGGREGATION_METHOD = 'fedbn_fedprox'
  674. cfg_mod.GRADIENT_COMBINATION_METHOD = 'fedbn_fedprox'
  675. cfg_mod.SCALING_FACTOR = 20.0
  676. _sync_config_to_modules()
  677. set_random_seeds(seed)
  678. try:
  679. train_dataset, test_dataset = load_dataset()
  680. root_loader = torch.utils.data.DataLoader(
  681. train_dataset, batch_size=cfg_mod.BATCH_SIZE, shuffle=True, num_workers=0)
  682. server = Server()
  683. server.set_datasets(root_loader, test_dataset)
  684. server._pretrain_global_model()
  685. _, client_datasets = create_client_datasets(
  686. train_dataset=train_dataset, num_clients=10, iid=False, alpha=0.1)
  687. malicious_indices = np.random.choice(10, 4, replace=False)
  688. clients = []
  689. client_is_malicious = {}
  690. for i in range(10):
  691. is_mal = int(i) in malicious_indices.tolist()
  692. client_is_malicious[i] = is_mal
  693. c = Client(client_id=i, dataset=client_datasets[i], is_malicious=is_mal)
  694. if is_mal:
  695. c.set_attack_parameters(
  696. attack_type='scaling_attack',
  697. scaling_factor=20.0,
  698. partial_percent=cfg_mod.PARTIAL_SCALING_PERCENT,
  699. noise_factor=cfg_mod.NOISE_FACTOR,
  700. flip_probability=cfg_mod.FLIP_PROBABILITY,
  701. )
  702. clients.append(c)
  703. server.add_clients(clients)
  704. root_gradients = server._collect_root_gradients()
  705. server.root_gradients = root_gradients # enables root-cosine & sign-consistency features
  706. server.vae = server.train_vae(root_gradients, vae_epochs=cfg_mod.VAE_EPOCHS)
  707. test_loader = torch.utils.data.DataLoader(
  708. test_dataset, batch_size=cfg_mod.BATCH_SIZE, shuffle=False, num_workers=0)
  709. server.test_loader = test_loader
  710. _, round_metrics = server.train(num_rounds=cfg_mod.GLOBAL_EPOCHS)
  711. final_acc = server.evaluate_model()
  712. logger.info(f" Final accuracy: {final_acc*100:.2f}%")
  713. # ---- Extract trust scores and Shapley values ----
  714. trust_data = [] # list of dicts: round, client_id, is_malicious, trust_score, shapley_value
  715. for r_idx, rd in round_metrics.items():
  716. trust_scores = rd.get('trust_scores', {})
  717. features_dict = rd.get('features', {}) # features[client_id] = [f0..f5]
  718. for client_id in range(10):
  719. ts = trust_scores.get(client_id, None)
  720. if ts is None:
  721. continue
  722. shap = None
  723. if client_id in features_dict:
  724. feats = features_dict[client_id]
  725. if isinstance(feats, list) and len(feats) >= 6:
  726. shap = float(feats[5])
  727. trust_data.append({
  728. 'round': int(r_idx) + 1,
  729. 'client_id': client_id,
  730. 'is_malicious': int(client_is_malicious.get(client_id, False)),
  731. 'trust_score': float(ts),
  732. 'shapley_value': shap,
  733. })
  734. if not trust_data:
  735. logger.warning(" No trust scores found in round_metrics. Skipping plots.")
  736. return {}
  737. # Save raw CSV
  738. csv_path = os.path.join(RESULTS_DIR, 'exp3_trust_scores.csv')
  739. fields = ['round', 'client_id', 'is_malicious', 'trust_score', 'shapley_value']
  740. with open(csv_path, 'w', newline='') as f:
  741. w = csv.DictWriter(f, fieldnames=fields, extrasaction='ignore')
  742. w.writeheader()
  743. w.writerows(trust_data)
  744. logger.info(f" Saved trust data CSV → {csv_path}")
  745. # ---- Build per-client time series ----
  746. rounds_list = sorted(set(d['round'] for d in trust_data))
  747. client_ts = {} # client_id → {round: trust_score}
  748. client_shap = {} # client_id → [shapley values across rounds]
  749. for d in trust_data:
  750. cid = d['client_id']
  751. if cid not in client_ts:
  752. client_ts[cid] = {}
  753. client_shap[cid] = []
  754. client_ts[cid][d['round']] = d['trust_score']
  755. if d['shapley_value'] is not None:
  756. client_shap[cid].append(d['shapley_value'])
  757. # Benign/malicious split
  758. benign_ids = [cid for cid, m in client_is_malicious.items() if not m]
  759. malicious_ids = [cid for cid, m in client_is_malicious.items() if m]
  760. # ---- Option A: Trust Score Time Series ----
  761. fig_a, ax_a = plt.subplots(figsize=(10, 5))
  762. for cid in benign_ids:
  763. ts_vals = [client_ts[cid].get(r, np.nan) for r in rounds_list]
  764. ax_a.plot(rounds_list, ts_vals, color='#2166ac', linewidth=1.5,
  765. alpha=0.75, label='Benign' if cid == benign_ids[0] else '')
  766. for cid in malicious_ids:
  767. ts_vals = [client_ts[cid].get(r, np.nan) for r in rounds_list]
  768. ax_a.plot(rounds_list, ts_vals, color='#d6604d', linewidth=1.5,
  769. linestyle='--', alpha=0.85,
  770. label='Malicious' if cid == malicious_ids[0] else '')
  771. patch_b = mpatches.Patch(color='#2166ac', label='Benign clients')
  772. patch_m = mpatches.Patch(color='#d6604d', label='Malicious clients')
  773. ax_a.legend(handles=[patch_b, patch_m], fontsize=11)
  774. ax_a.set_xlabel('Round', fontsize=12)
  775. ax_a.set_ylabel('Trust Score', fontsize=12)
  776. ax_a.set_ylim(-0.05, 1.05)
  777. ax_a.set_title('Trust Score Time Series: Benign vs. Malicious Clients\n'
  778. '(Alzheimer MRI, Scaling ×20, 40% Malicious, Seed=42)', fontsize=12)
  779. ax_a.grid(True, alpha=0.3)
  780. fig_a.tight_layout()
  781. fig_a_path = os.path.join(RESULTS_DIR, 'exp3_trust_timeseries.png')
  782. fig_a.savefig(fig_a_path, dpi=300, bbox_inches='tight')
  783. plt.close(fig_a)
  784. logger.success(f" Saved Option A (time series) → {fig_a_path}")
  785. # ---- Option B: Trust Score Box Plot ----
  786. benign_scores = [client_ts[cid].get(r, np.nan)
  787. for cid in benign_ids for r in rounds_list
  788. if not np.isnan(client_ts[cid].get(r, np.nan))]
  789. malicious_scores = [client_ts[cid].get(r, np.nan)
  790. for cid in malicious_ids for r in rounds_list
  791. if not np.isnan(client_ts[cid].get(r, np.nan))]
  792. if benign_scores and malicious_scores:
  793. fig_b, ax_b = plt.subplots(figsize=(6, 5))
  794. bp = ax_b.boxplot(
  795. [benign_scores, malicious_scores],
  796. labels=['Benign Clients', 'Malicious Clients'],
  797. patch_artist=True,
  798. medianprops=dict(color='black', linewidth=2),
  799. )
  800. bp['boxes'][0].set_facecolor('#a6cee3')
  801. if len(bp['boxes']) > 1:
  802. bp['boxes'][1].set_facecolor('#fb9a99')
  803. ax_b.set_ylabel('Average Trust Score', fontsize=12)
  804. ax_b.set_title('Trust Score Distribution\n'
  805. '(Benign vs. Malicious Clients)', fontsize=12)
  806. ax_b.grid(True, axis='y', alpha=0.3)
  807. fig_b.tight_layout()
  808. fig_b_path = os.path.join(RESULTS_DIR, 'exp3_trust_boxplot.png')
  809. fig_b.savefig(fig_b_path, dpi=300, bbox_inches='tight')
  810. plt.close(fig_b)
  811. logger.success(f" Saved Option B (box plot) → {fig_b_path}")
  812. # ---- Option C: Shapley Value Distribution (violin) ----
  813. benign_shap = [v for cid in benign_ids for v in client_shap.get(cid, [])]
  814. malicious_shap = [v for cid in malicious_ids for v in client_shap.get(cid, [])]
  815. if benign_shap and malicious_shap:
  816. fig_c, ax_c = plt.subplots(figsize=(6, 5))
  817. parts = ax_c.violinplot(
  818. [benign_shap, malicious_shap],
  819. positions=[1, 2],
  820. showmedians=True,
  821. )
  822. for pc, color in zip(parts['bodies'], ['#a6cee3', '#fb9a99']):
  823. pc.set_facecolor(color)
  824. pc.set_alpha(0.75)
  825. ax_c.set_xticks([1, 2])
  826. ax_c.set_xticklabels(['Benign Clients', 'Malicious Clients'])
  827. ax_c.set_ylabel('Shapley Value', fontsize=12)
  828. ax_c.set_title('Shapley Value Distribution\n'
  829. '(Benign vs. Malicious Clients)', fontsize=12)
  830. ax_c.grid(True, axis='y', alpha=0.3)
  831. fig_c.tight_layout()
  832. fig_c_path = os.path.join(RESULTS_DIR, 'exp3_shapley_violin.png')
  833. fig_c.savefig(fig_c_path, dpi=300, bbox_inches='tight')
  834. plt.close(fig_c)
  835. logger.success(f" Saved Option C (Shapley violin) → {fig_c_path}")
  836. logger.info('\nExperiment 3 complete.')
  837. return {'trust_data': trust_data, 'final_accuracy': float(final_acc)}
  838. except Exception as exc:
  839. logger.error(f" Experiment 3 FAILED: {str(exc)[:200]}")
  840. traceback.print_exc()
  841. return {'status': 'failed', 'error': str(exc)}
  842. # ===========================================================================
  843. # EXPERIMENT 4 — OPTIMIZER COMPARISON UNDER ADVERSARIAL CONDITIONS
  844. # ===========================================================================
  845. # Optimizers to compare, in priority order (run as many as time allows)
  846. EXP4_OPTIMIZERS = [
  847. {'name': 'FedAvg', 'method': 'fedavg', 'optigradtrust': False},
  848. {'name': 'FedBN', 'method': 'fedbn', 'optigradtrust': False},
  849. {'name': 'FedBN-P', 'method': 'fedbn_fedprox', 'optigradtrust': True}, # ours
  850. {'name': 'FedProx', 'method': 'fedprox', 'optigradtrust': False},
  851. {'name': 'FedNova', 'method': 'fednova', 'optigradtrust': False},
  852. {'name': 'SCAFFOLD', 'method': 'scaffold', 'optigradtrust': False},
  853. {'name': 'FedDWA', 'method': 'feddwa', 'optigradtrust': False},
  854. {'name': 'FedADMM', 'method': 'fedadmm', 'optigradtrust': False},
  855. ]
  856. def _run_exp4_single_optimizer(opt_cfg, seed, num_rounds):
  857. """Run one optimizer for Experiment 4 (adversarial conditions)."""
  858. import federated_learning.config.config as cfg_mod
  859. from federated_learning.training.server import Server
  860. from federated_learning.training.client import Client
  861. from federated_learning.data.dataset_utils import load_dataset, create_client_datasets
  862. from federated_learning.utils.model_utils import set_random_seeds
  863. opt_name = opt_cfg['name']
  864. opt_method = opt_cfg['method']
  865. use_trust = opt_cfg['optigradtrust']
  866. logger.info(f" [{opt_name}] seed={seed}")
  867. set_all_seeds(seed)
  868. # Configure
  869. configure_for_dataset(
  870. 'ALZHEIMER', num_clients=10,
  871. non_iid_config={'enable': False, 'type': 'iid', 'alpha': None},
  872. aggregation_method=opt_method
  873. )
  874. cfg_mod.RANDOM_SEED = seed
  875. cfg_mod.GLOBAL_EPOCHS = num_rounds
  876. cfg_mod.FRACTION_MALICIOUS = 0.3
  877. cfg_mod.NUM_MALICIOUS = 3
  878. cfg_mod.BATCH_SIZE = 16
  879. cfg_mod.LR = 1e-4
  880. cfg_mod.WEIGHT_DECAY = 5e-5
  881. cfg_mod.SCALING_FACTOR = 10.0
  882. # For FedBN-P (our method), enable full OptiGradTrust trust mechanism
  883. # For others, use their plain aggregation (baseline mode — no trust)
  884. if use_trust:
  885. cfg_mod.ENABLE_VAE = True
  886. cfg_mod.ENABLE_SHAPLEY = True
  887. cfg_mod.ENABLE_DUAL_ATTENTION = True
  888. cfg_mod.RL_AGGREGATION_METHOD = 'hybrid'
  889. cfg_mod.RL_WARMUP_ROUNDS = 5
  890. cfg_mod.RL_RAMP_UP_ROUNDS = 10
  891. cfg_mod.AGGREGATION_METHOD = opt_method
  892. cfg_mod.GRADIENT_COMBINATION_METHOD = opt_method
  893. # else: configure_for_dataset already set baseline mode (VAE/Shapley/DA disabled)
  894. _sync_config_to_modules()
  895. set_random_seeds(seed)
  896. start_t = time.time()
  897. try:
  898. train_dataset, test_dataset = load_dataset()
  899. root_loader = torch.utils.data.DataLoader(
  900. train_dataset, batch_size=cfg_mod.BATCH_SIZE, shuffle=True, num_workers=0)
  901. server = Server()
  902. server.set_datasets(root_loader, test_dataset)
  903. server._pretrain_global_model()
  904. _, client_datasets = create_client_datasets(
  905. train_dataset=train_dataset, num_clients=10, iid=True, alpha=None)
  906. malicious_indices = np.random.choice(10, 3, replace=False)
  907. clients = []
  908. for i in range(10):
  909. is_mal = int(i) in malicious_indices.tolist()
  910. c = Client(client_id=i, dataset=client_datasets[i], is_malicious=is_mal)
  911. if is_mal:
  912. c.set_attack_parameters(
  913. attack_type='scaling_attack',
  914. scaling_factor=10.0,
  915. partial_percent=cfg_mod.PARTIAL_SCALING_PERCENT,
  916. noise_factor=cfg_mod.NOISE_FACTOR,
  917. flip_probability=cfg_mod.FLIP_PROBABILITY,
  918. )
  919. clients.append(c)
  920. server.add_clients(clients)
  921. root_gradients = server._collect_root_gradients()
  922. server.root_gradients = root_gradients # enables root-cosine & sign-consistency features
  923. if getattr(cfg_mod, 'ENABLE_VAE', False):
  924. server.vae = server.train_vae(root_gradients, vae_epochs=cfg_mod.VAE_EPOCHS)
  925. else:
  926. server.vae = None
  927. test_loader = torch.utils.data.DataLoader(
  928. test_dataset, batch_size=cfg_mod.BATCH_SIZE, shuffle=False, num_workers=0)
  929. server.test_loader = test_loader
  930. _, round_metrics = server.train(num_rounds=cfg_mod.GLOBAL_EPOCHS)
  931. final_acc = server.evaluate_model()
  932. # Per-round accuracy
  933. per_round_acc = {}
  934. for r_idx, rd in round_metrics.items():
  935. per_round_acc[int(r_idx) + 1] = float(rd.get('test_accuracy', 0.0))
  936. result = {
  937. 'optimizer_name': opt_name,
  938. 'method': opt_method,
  939. 'seed': seed,
  940. 'final_accuracy': float(final_acc),
  941. 'per_round_acc': per_round_acc,
  942. 'total_time': time.time() - start_t,
  943. 'status': 'completed',
  944. }
  945. logger.success(f" → Final Acc={final_acc*100:.2f}%")
  946. return result
  947. except Exception as exc:
  948. logger.error(f" → FAILED: {str(exc)[:120]}")
  949. traceback.print_exc()
  950. return {
  951. 'optimizer_name': opt_name,
  952. 'method': opt_method,
  953. 'seed': seed,
  954. 'status': 'failed',
  955. 'error': str(exc),
  956. 'total_time': time.time() - start_t,
  957. }
  958. def run_exp4_optimizer_adversarial(seed=42, dry_run=False, n_optimizers=4):
  959. """
  960. Experiment 4: Optimizer comparison under adversarial conditions.
  961. IID, 30% malicious, scaling x10, 25 rounds, seed=42.
  962. Generates:
  963. - Summary table (CSV): final accuracy per optimizer
  964. - Per-round CSV: accuracy vs round (for convergence figure)
  965. - Convergence figure (PNG)
  966. """
  967. logger.info('=' * 70)
  968. logger.info('EXPERIMENT 4: OPTIMIZER COMPARISON UNDER ADVERSARIAL CONDITIONS')
  969. logger.info(' Dataset: Alzheimer MRI | IID | Malicious: 30% | Scaling x10')
  970. logger.info(' Compares: FedAvg, FedBN, FedBN-P, FedProx [+ FedNova, SCAFFOLD, FedDWA, FedADMM]')
  971. logger.info('=' * 70)
  972. num_rounds = 2 if dry_run else 25
  973. # Limit to n_optimizers (priority order); if dry_run, just top 3
  974. opts_to_run = EXP4_OPTIMIZERS[:3] if dry_run else EXP4_OPTIMIZERS[:n_optimizers]
  975. all_results = []
  976. per_round_rows = []
  977. for opt_cfg in opts_to_run:
  978. logger.info(f"\n Optimizer: {opt_cfg['name']}")
  979. r = _run_exp4_single_optimizer(opt_cfg, seed, num_rounds)
  980. all_results.append(r)
  981. if r['status'] == 'completed':
  982. for round_num, acc in r['per_round_acc'].items():
  983. per_round_rows.append({
  984. 'round': round_num,
  985. 'optimizer': opt_cfg['name'],
  986. 'accuracy': acc,
  987. 'is_our_method': int(opt_cfg['optigradtrust']),
  988. })
  989. # ---- Per-round CSV ----
  990. csv_pr_path = os.path.join(RESULTS_DIR, 'exp4_optimizer_adversarial_per_round.csv')
  991. if per_round_rows:
  992. with open(csv_pr_path, 'w', newline='') as f:
  993. w = csv.DictWriter(f, fieldnames=['round', 'optimizer', 'accuracy', 'is_our_method'])
  994. w.writeheader()
  995. w.writerows(per_round_rows)
  996. logger.info(f"\n Saved per-round CSV → {csv_pr_path}")
  997. # ---- Summary table ----
  998. logger.info('\n' + '=' * 70)
  999. logger.info('EXPERIMENT 4 SUMMARY: Optimizer Final Accuracy under Adversarial Conditions')
  1000. logger.info(f"{'Optimizer':<12} {'Final Acc (%)':>14} {'Defense':>10}")
  1001. logger.info('-' * 50)
  1002. summary_rows = []
  1003. for r in all_results:
  1004. if r['status'] != 'completed':
  1005. continue
  1006. has_trust = next((o['optigradtrust'] for o in opts_to_run
  1007. if o['name'] == r['optimizer_name']), False)
  1008. defense_str = 'OptiGradTrust' if has_trust else 'None'
  1009. logger.info(f" {r['optimizer_name']:<10} {r['final_accuracy']*100:>14.2f} {defense_str:>10}")
  1010. summary_rows.append({
  1011. 'optimizer': r['optimizer_name'],
  1012. 'method': r['method'],
  1013. 'final_accuracy': r['final_accuracy'],
  1014. 'has_defense': int(has_trust),
  1015. })
  1016. logger.info('=' * 70)
  1017. csv_sum_path = os.path.join(RESULTS_DIR, 'exp4_optimizer_adversarial_summary.csv')
  1018. if summary_rows:
  1019. with open(csv_sum_path, 'w', newline='') as f:
  1020. w = csv.DictWriter(f, fieldnames=list(summary_rows[0].keys()))
  1021. w.writeheader()
  1022. w.writerows(summary_rows)
  1023. logger.info(f" Saved summary CSV → {csv_sum_path}")
  1024. json_path = os.path.join(RESULTS_DIR, 'exp4_optimizer_adversarial_results.json')
  1025. with open(json_path, 'w') as f:
  1026. json.dump(all_results, f, indent=2, default=str)
  1027. logger.info(f" Saved JSON → {json_path}")
  1028. # ---- Convergence figure ----
  1029. if per_round_rows:
  1030. try:
  1031. opt_names = list(dict.fromkeys(r['optimizer'] for r in per_round_rows))
  1032. colors = plt.cm.tab10(np.linspace(0, 1, len(opt_names)))
  1033. fig, ax = plt.subplots(figsize=(10, 5))
  1034. for opt_name, color in zip(opt_names, colors):
  1035. rows = [(r['round'], r['accuracy'])
  1036. for r in per_round_rows if r['optimizer'] == opt_name]
  1037. if not rows:
  1038. continue
  1039. rows.sort(key=lambda x: x[0])
  1040. xs = [r[0] for r in rows]
  1041. ys = [r[1] * 100 for r in rows]
  1042. lw = 2.5 if opt_name == 'FedBN-P' else 1.5
  1043. ls = '-' if opt_name == 'FedBN-P' else '--'
  1044. ax.plot(xs, ys, label=opt_name, color=color, linewidth=lw, linestyle=ls)
  1045. ax.set_xlabel('Round', fontsize=12)
  1046. ax.set_ylabel('Test Accuracy (%)', fontsize=12)
  1047. ax.set_title('Optimizer Comparison under Adversarial Conditions\n'
  1048. '(Alzheimer MRI, IID, 30% Malicious, Scaling ×10)', fontsize=12)
  1049. ax.legend(fontsize=10)
  1050. ax.grid(True, alpha=0.3)
  1051. fig.tight_layout()
  1052. fig_path = os.path.join(RESULTS_DIR, 'exp4_optimizer_adversarial_convergence.png')
  1053. fig.savefig(fig_path, dpi=300, bbox_inches='tight')
  1054. plt.close(fig)
  1055. logger.success(f" Saved convergence figure → {fig_path}")
  1056. except Exception as exc:
  1057. logger.warning(f" Could not generate convergence figure: {exc}")
  1058. logger.info('\nExperiment 4 complete.')
  1059. return all_results
  1060. # ===========================================================================
  1061. # MAIN
  1062. # ===========================================================================
  1063. def main():
  1064. parser = argparse.ArgumentParser(
  1065. description='OptiGradTrust Round-2 Revision Experiments',
  1066. formatter_class=argparse.RawDescriptionHelpFormatter,
  1067. epilog="""
  1068. Priority order (run in this order if time is limited):
  1069. 1. exp2 — Dynamic attack / RL justification [NON-NEGOTIABLE]
  1070. 2. exp1a — Noise ablation / VAE validation [CRITICAL]
  1071. 3. exp3 — Trust visualization [REQUIRED]
  1072. 4. exp1b — Sign-flip ablation [CRITICAL]
  1073. 5. exp4 — Optimizer adversarial comparison [RECOMMENDED]
  1074. Examples:
  1075. python run_r2_experiments.py --experiment exp2 --seeds 42 123 456
  1076. python run_r2_experiments.py --experiment exp1a --seeds 42 123 456 789 1024
  1077. python run_r2_experiments.py --experiment all --seeds 42
  1078. python run_r2_experiments.py --experiment exp2 --dry-run
  1079. """
  1080. )
  1081. parser.add_argument(
  1082. '--experiment',
  1083. choices=['exp1a', 'exp1b', 'exp2', 'exp3', 'exp4', 'all'],
  1084. default='all',
  1085. help='Which experiment to run'
  1086. )
  1087. parser.add_argument(
  1088. '--seeds', type=int, nargs='+',
  1089. default=None,
  1090. help='Random seeds (default: single seed=42 for most; 42 123 456 for exp2)'
  1091. )
  1092. parser.add_argument(
  1093. '--dry-run', action='store_true',
  1094. help='Quick 2-3 round sanity check (do not use for real results)'
  1095. )
  1096. parser.add_argument(
  1097. '--n-optimizers', type=int, default=4,
  1098. help='Number of optimizers to run in Experiment 4 (1-8, default 4)'
  1099. )
  1100. args = parser.parse_args()
  1101. logger.info('=' * 70)
  1102. logger.info('OptiGradTrust — Round 2 Revision Experiments')
  1103. logger.info(f'Experiment: {args.experiment} | Dry-run: {args.dry_run}')
  1104. logger.info(f'Seeds: {args.seeds}')
  1105. logger.info(f'Output dir: {RESULTS_DIR}')
  1106. logger.info('=' * 70)
  1107. exp = args.experiment
  1108. dry = args.dry_run
  1109. if exp in ('exp1a', 'all'):
  1110. seeds = args.seeds or ([42] if dry else [42])
  1111. run_exp1a_noise_ablation(seeds=seeds, dry_run=dry)
  1112. if exp in ('exp1b', 'all'):
  1113. seeds = args.seeds or ([42] if dry else [42])
  1114. run_exp1b_signflip_ablation(seeds=seeds, dry_run=dry)
  1115. if exp in ('exp2', 'all'):
  1116. seeds = args.seeds or ([42] if dry else [42, 123, 456])
  1117. run_exp2_dynamic_attack(seeds=seeds, dry_run=dry)
  1118. if exp in ('exp3', 'all'):
  1119. seed = (args.seeds[0] if args.seeds else 42)
  1120. run_exp3_trust_visualization(seed=seed, dry_run=dry)
  1121. if exp in ('exp4', 'all'):
  1122. seed = (args.seeds[0] if args.seeds else 42)
  1123. run_exp4_optimizer_adversarial(seed=seed, dry_run=dry,
  1124. n_optimizers=args.n_optimizers)
  1125. logger.info('\n' + '=' * 70)
  1126. logger.info('All requested R2 experiments finished.')
  1127. logger.info(f'Results saved in: {RESULTS_DIR}')
  1128. logger.info('=' * 70)
  1129. if __name__ == '__main__':
  1130. main()

run_r2_experiments.py at commit 92e325e, under MIT · at the source

Overview

Authors: Mohammad Karami1, Hamed Kebriaei1, Fatemeh Ghassemi1, Hamid Azadegan2
ORCID iDs: Mohammad Karami
  1. School of Electrical and Computer Engineering, University of Tehran, Tehran, 1439957131 Iran
  2. School of Computer Engineering, Iran University of Science and Technology (IUST), Tehran, 1684613114 Iran
Journal: Scientific reports, volume 16, issue 1, article 25604
Dates: received 25 February 2026; accepted 27 May 2026; published online 4 June 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1038/s41598-026-55855-5 · PMID 42243446 · PMCID PMC13478148 · OpenAlex W4414921340
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), human (organism), Alzheimer's / dementia (population), cellular / molecular (subfield)
Methods: Statistics, Machine learning
Keywords: Federated learning, Byzantine robustness, Trust-aware aggregation, Non-IID data, Brain MRI, Alzheimer’s disease, Mathematics and computing, Neuroscience
MeSH: Brain*, Image Processing, Computer-Assisted*, Magnetic Resonance Imaging*, Algorithms, Federated Learning, Humans, Trust (* major topic)
Topic: Privacy-Preserving Technologies in Data (Artificial Intelligence, Computer Science), according to OpenAlex
Citations: not cited yet (Europe PMC); 58 references in the paper

Abstract

Federated learning enables collaborative training across institutions without centralizing patient data, but remains vulnerable to malicious clients and severe non-IID data heterogeneity. We propose a trust-aware federated learning framework for brain MRI that combines multi-signal gradient fingerprinting with adaptive aggregation to achieve Byzantine robustness. Each client update is characterized by a six-dimensional fingerprint (variational-autoencoder reconstruction error, cosine similarity to a server reference, peer similarity, gradient norm, sign consistency, and Monte Carlo Shapley contribution). A dual-attention module and a reinforcement-learning controller map these signals into trust weights and integrate with FedBN-P (Federated Batch Normalization with Proximal regularization), an optimizer co-designed for stability under heterogeneous and adversarial conditions. We evaluate on MNIST, CIFAR-10, Alzheimer’s MRI, and the OASIS brain-MRI cohort (approximately 87 test samples, used strictly as proof-of-concept) under both standard and strengthened threat models (up to 40% malicious clients). Attack-specific ablation confirms a defense-in-depth design: VAE fingerprinting is the primary noise-attack defense (3.60 pp accuracy drop upon removal), Shapley values safeguard accuracy under scaling (10.16 pp drop), and reinforcement learning improves detection consistency under dynamic attack schedules. Three-seed paired-test validation further shows detection F1 outperforms FLTrust by up to 44 pp on Non-IID Gaussian noise; a white-box adaptive attacker degrades the detector but not model accuracy, confirming the layered design. End-to-end wall-clock overhead is + 8.8% over FedAvg with identical communication volume. The framework achieves F1 above 0.98 for gradient-scaling attacks while preserving accuracy under magnitude-preserving attacks where explicit detection remains limited. Multi-site validation on larger federated cohorts (e.g., ADNI, UK Biobank, FeTS) is required before any clinical-deployment claim can be made.

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

mohammadkarami79/OptiGradTrust

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 92e325ee5db6f8228f8e7daccd2abcf74472b94c, 20 April 2026
Languages: Python (52)
Size: 78 files, 52 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, environment (requirements.txt, setup.py, federated_learning/requirements.txt)
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: PyTorch (40 files), NumPy (32 files), Matplotlib (8 files), pandas (4 files), Pillow (3 files), SciPy (3 files), scikit-learn (2 files), seaborn (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
54 files

Code availability

The code implementing OptiGradTrust (training loops, six-dimensional fingerprinting, dual-attention, DDQN trust controller, FedBN-P optimizer, and all evaluation scripts) is publicly available at https://github.com/mohammadkarami79/OptiGradTrust. The implementation uses Python 3.9, PyTorch 1.12 (and tested up to PyTorch 2.7 on the R3 experiments), and standard scientific Python libraries (NumPy, SciPy, Matplotlib). Reproducibility commitment. The repository includes: (i) comprehensive documentation and configuration files for every experiment reported in the paper (original-configuration runs, strengthened-configuration runs, and the six R3 additional analyses described in Section 2.14); (ii) deterministic seed lists ( for R3, plus for the strengthened-configuration multi-seed runs) so that every reported number can be reproduced; (iii) the exact data splits used for OASIS, Alzheimer’s MRI, MNIST, and CIFAR-10; (iv) trained VAE checkpoints for the Alzheimer’s MRI benchmark; (v) per-experiment raw CSV / JSON artifacts for the R3 analyses; and (vi) an environment.yml pinning all package versions. Upon acceptance we will additionally release a Dockerfile encapsulating the full software stack.

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

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;
  • 52 scripts, each with its path and the digest of its content;
  • 32 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

This study used only publicly available, de-identified datasets. The MNIST dataset was accessed via the TensorFlow/Keras canonical release (https://storage.googleapis.com/tensorflow/tf-keras-datasets/mnist.npz). CIFAR-10 is available at https://www.cs.toronto.edu/~kriz/cifar.html. The Alzheimer’s MRI dataset is a public Kaggle resource (https://www.kaggle.com/datasets/lukechugh/best-alzheimer-mri-dataset-99-accuracy). The OASIS (Open Access Series of Imaging Studies) cross-sectional MRI dataset is available at https://www.oasis-brains.org/. All datasets were used in accordance with their respective licenses. No new human data were collected; OASIS data are de‑identified and released under open-access terms.

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, 4 authors, 8 keywords, 7 MeSH terms, 8 references.

Cite

This paper

Karami, M., Kebriaei, H., Ghassemi, F., & Azadegan, H. (2026). Byzantine robust federated learning for heterogeneous brain MRI using multisignal gradient fingerprinting and adaptive trust aggregation. Scientific reports, 16(1), 25604. https://doi.org/10.1038/s41598-026-55855-5

BibTeX

@article{karami2026byzantine,
author = {Karami, Mohammad and Kebriaei, Hamed and Ghassemi, Fatemeh and Azadegan, Hamid},
title = {{Byzantine robust federated learning for heterogeneous brain MRI using multisignal gradient fingerprinting and adaptive trust aggregation}},
journal = {Scientific reports},
year = {2026},
month = jun,
volume = {16},
number = {1},
pages = {25604},
publisher = {Nature Publishing Group},
issn = {2045-2322},
doi = {10.1038/s41598-026-55855-5},
url = {https://doi.org/10.1038/s41598-026-55855-5},
pmid = {42243446},
pmcid = {PMC13478148}
}

RIS

TY - JOUR
AU - Karami, Mohammad
AU - Kebriaei, Hamed
AU - Ghassemi, Fatemeh
AU - Azadegan, Hamid
TI - Byzantine robust federated learning for heterogeneous brain MRI using multisignal gradient fingerprinting and adaptive trust aggregation
T2 - Scientific reports
J2 - Sci Rep
PY - 2026
DA - 2026/06/04
VL - 16
IS - 1
SP - 25604
SN - 2045-2322
PB - Nature Publishing Group
DO - 10.1038/s41598-026-55855-5
UR - https://doi.org/10.1038/s41598-026-55855-5
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41598-026-55855-5",
"type": "article-journal",
"title": "Byzantine robust federated learning for heterogeneous brain MRI using multisignal gradient fingerprinting and adaptive trust aggregation",
"container-title": "Scientific reports",
"author": [
{
"family": "Karami",
"given": "Mohammad"
},
{
"family": "Kebriaei",
"given": "Hamed"
},
{
"family": "Ghassemi",
"given": "Fatemeh"
},
{
"family": "Azadegan",
"given": "Hamid"
}
],
"container-title-short": "Sci Rep",
"volume": "16",
"issue": "1",
"page": "25604",
"DOI": "10.1038/s41598-026-55855-5",
"PMID": "42243446",
"PMCID": "PMC13478148",
"ISSN": "2045-2322",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41598-026-55855-5",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
4
]
]
}
}

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.1002/alz.71649 [code]
Postmortem brain MRI reveals differential associations of subcortical and limbic volumes with cortical thinning and neurodegenerative pathologies.
Journal: Alzheimer's & dementia : the journal of the Alzheimer's Association
In common: Pillow, PyTorch, seaborn, 5 other tools, Alzheimer's / dementia, structural MRI / diffusion, cellular / molecular
[2] doi:10.64898/2026.08.18.26360725 [code]
Temporal pole blurring in hippocampal sclerosis reflects seizure-disrupted myelination
Journal: medRxiv (preprint)
In common: Pillow, PyTorch, seaborn, 5 other tools, structural MRI / diffusion, cellular / molecular
[3] doi:10.3390/jimaging12070276 [code]
Hyperelastic Regularization for Near-Diffeomorphic Transformer-Based Brain MRI Registration.
Journal: Journal of imaging
In common: Pillow, PyTorch, scikit-learn, 4 other tools, structural MRI / diffusion, 1 reference
[4] doi:10.1038/s41467-026-76837-1 [code]
Drug screen and machine learning predict neuroprotective agents in a preclinical human model of childhood dementia.
Journal: Nature communications
In common: Pillow, PyTorch, seaborn, 5 other tools, Alzheimer's / dementia
[5] doi:10.1093/bioinformatics/btag540 [code]
Deciphering spatial heterogeneity by multimodal spatial transcriptomics modelling with SpatialModal.
Journal: Bioinformatics (Oxford, England)
In common: Pillow, PyTorch, seaborn, 5 other tools, Alzheimer's / dementia
[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: Pillow, PyTorch, seaborn, 5 other tools, structural MRI / diffusion
[7] 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: Pillow, PyTorch, seaborn, 5 other tools, structural MRI / diffusion
[8] doi:10.1007/s00234-026-04103-8 [code]
Enhanced detection of subtle cortical abnormalities in focal epilepsy using 7 T MRI surface-based models and graph neural networks.
Journal: Neuroradiology
In common: Pillow, PyTorch, seaborn, 5 other tools, structural MRI / diffusion
[9] doi:10.1093/braincomms/fcag253 [code]
Disease detection and classification in temporal lobe epilepsy: step-wise versus simultaneous AI decision models in a multisite neuroimaging study.
Journal: Brain communications
In common: Pillow, PyTorch, seaborn, 5 other tools, structural MRI / diffusion
[10] doi:10.1002/ana.78203 [code]
AI-Driven Mapping of Seizure Spread Patterns.
Journal: Annals of neurology
In common: Pillow, PyTorch, seaborn, 5 other tools, structural MRI / diffusion

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.