OSCR

A generalized logistic-logit function and its application to multi-layer perceptron and neuron segmentation.

Code ↔ Paper

17 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 17 matches · 4 of them tie a paragraph to a whole file, not to given lines: weak matches, whose lines are not tinted
  1. [1] § Results and discussion › Cannistraci–Muscoloni–Gu ↔ gllf-function/CMG_GLLF.m, the whole file · a weak match · score 0.75 · generalized logistic function, generalized logit function, inflection parameter, inflection point, linear function, Gu
  2. [2] § Results and discussion › Cannistraci–Muscoloni–Gu ↔ gllf-function/CMG_GLLF.py, the whole file · a weak match · score 0.75 · generalized logistic function, generalized logit function, inflection parameter, inflection point, linear function, Gu
  3. [3] § Results and discussion › Improving affinity-graph-based neuron image segmentation algorithm with CMG ↔ neuron-seg/stats_visualization.ipynb, lines 340–450 · score 0.72 · CREMI score, inflection rate, CMG enhanced, inflection point, space, SOTA
  4. [4] § Materials and methods › Image modulation and contrast analysis ↔ ifm/plot_results.ipynb, lines 1423–1464 · score 0.65 · BT.601, CMG modulation, grayscale, RGB, luma, weights
  5. [5] § Results and discussion › CMG as MLP neural network input feature modulator (IFM) ↔ rebuttal-plot/common.py, lines 61–144 · score 0.62 · AdamW, finite loss, vanilla MLP, Grad, Param, Muon
  6. [6] § Materials and methods › IFM training setups ↔ rebuttal-plot/rules.py, the whole file · a weak match · score 0.59 · hyper parameter, simple CNN, channels, Muon, CIFAR100, MLP
  7. [7] § Results and discussion › CMG as MLP neural network input feature modulator (IFM) ↔ rebuttal-plot/rules.py, the whole file · a weak match · score 0.58 · AdamW, numerical stability, vanilla MLP, Muon, loss, CIFAR100
  8. [8] § Results and discussion › Improving affinity-graph-based neuron image segmentation algorithm with CMG ↔ neuron-seg/stats_visualization.ipynb, lines 340–450 · score 0.58 · CREMI score, CMG enhanced, inflection point, SOTA, axis, neuron
  9. [9] § Materials and methods › IFM training setups ↔ cnn/hanming_trainer.py, lines 371–425 · score 0.57 · AdamW, weight decay, momentum, SGD, Muon, scheduler
  10. [10] § Materials and methods › Quantitative evaluation metrics ↔ neuron-seg/stats_visualization.ipynb, lines 37–109 · score 0.56 · CREMI score, ARAND, VOI, segmentation, neuron, error
  11. [11] § Materials and methods › Image modulation and contrast analysis ↔ ifm/plot_results.ipynb, lines 1466–1525 · score 0.55 · pixel intensities, RMS contrast, luma
  12. [12] § Materials and methods › Quantitative evaluation metrics ↔ rebuttal-plot/common.py, lines 61–144 · score 0.54 · peak GPU memory, finite loss, fair, metrics, model, training
  13. [13] § Results and discussion › CMG as MLP neural network input feature modulator (IFM) ↔ ifm/plot_results.ipynb, lines 713–841 · score 0.54 · supra linear scheduler, sub linear scheduler, ANN, network, AAE, IFM
  14. [14] § Results and discussion › Improving affinity-graph-based neuron image segmentation algorithm with CMG ↔ neuron-seg/main.py, lines 51–116 · score 0.54 · affinity graph, augmented, thresholding, segments, SOTA, CREMI
  15. [15] § Results and discussion › Improving affinity-graph-based neuron image segmentation algorithm with CMG ↔ neuron-seg/main.py, lines 51–116 · score 0.52 · affinity graph, thresholds, segmentation, SOTA, CREMI, neuron
  16. [16] § Results and discussion › Preliminary results of CMG as input element modulator in CNN ↔ pinn/network.py, lines 60–90 · score 0.52 · fully connected, numerical stability, behavior, layers, linear, IFM
  17. [17] § Materials and methods › Learning rate schedulers ↔ ifm/plot_results.ipynb, lines 664–711 · score 0.50 · supra linear scheduler, sub linear scheduler, training

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 · 2,184 lines · 88 KB · no license · 4 matches

  1. # %% [markdown]
  2. # # plot results
  3. # %%
  4. # Cell 1 - Imports and basic functions
  5. from scipy.io import loadmat
  6. import numpy as np
  7. import os
  8. from openpyxl import Workbook
  9. from train_tuner import get_fname as get_fname_tuner
  10. import matplotlib.pyplot as plt
  11. # %%
  12. def recover_dict(data):
  13. """
  14. Recovers a dictionary loaded from a .mat file, converting arrays back to
  15. their original Python structures (lists, scalars, booleans, strings, etc.).
  16. """
  17. def recover_value(value):
  18. """Converts individual values to their appropriate Python type."""
  19. if isinstance(value, np.ndarray):
  20. # Handle scalar values
  21. if value.size == 1:
  22. item = value.item()
  23. return item if not (isinstance(item, float) and np.isnan(item)) else np.nan
  24. # Handle string arrays
  25. if value.dtype.type is np.str_ or value.dtype.type is np.object_:
  26. return value.tolist() if value.ndim > 0 else str(value.item())
  27. # Handle 2D row vectors
  28. if value.ndim == 2 and value.shape[0] == 1:
  29. return value.flatten().tolist()
  30. return value # Return other arrays as is
  31. return value
  32. return {key: recover_value(value) for key, value in data.items() if not key.startswith('__')}
  33. def get_acc(mat_fname):
  34. '''
  35. Get best accuracy
  36. '''
  37. res = recover_dict(loadmat(mat_fname))
  38. return np.max(res["top1"])
  39. # %% [markdown]
  40. # # get training curves
  41. # %%
  42. def find_best_tuned_result(dataset, network, mode, scheduler, epochs, seeds, init_method,
  43. single_param=False, rectify=False,train_2=True,
  44. weight_decay=5e-4,lrs=[0.025, 0.01, 0.001],
  45. batch_sizes=[32, 64, 128], use_tuner=True,):
  46. """
  47. Find the best hyperparameter combination for a given tuned mode based on highest test accuracy.
  48. Returns the best result data, parameters, and averaged metrics across seeds.
  49. """
  50. best_result = None
  51. best_params = None
  52. # First find the best hyperparameter combination by averaging across seeds
  53. hyperparam_results = {}
  54. for lr in lrs:
  55. for bs in batch_sizes:
  56. combo_key = (lr, lr, bs)
  57. combo_accs = []
  58. combo_aaes = []
  59. combo_results = []
  60. for seed in seeds:
  61. try:
  62. # Get file path
  63. result_path, _ = get_fname_tuner(
  64. dataset, network, bs, lr, lr, epochs, seed,
  65. use_tuner, mode, train_2, init_method, weight_decay,
  66. rectify, single_param, scheduler_w=scheduler
  67. )
  68. # Load test and train results
  69. test_file = os.path.join(result_path, 'res_test.mat')
  70. train_file = os.path.join(result_path, 'res_train.mat')
  71. if os.path.exists(test_file) and os.path.exists(train_file):
  72. test_data = recover_dict(loadmat(test_file))
  73. train_data = recover_dict(loadmat(train_file))
  74. max_test_acc = np.max(test_data["top1"])
  75. test_aae = test_data.get("aae", float(np.trapz(test_data["top1"])/(len(test_data["top1"])-1)))
  76. combo_accs.append(max_test_acc)
  77. combo_aaes.append(test_aae)
  78. combo_results.append({
  79. 'test': test_data,
  80. 'train': train_data,
  81. 'seed': seed,
  82. 'acc': max_test_acc
  83. })
  84. else:
  85. print(f"Files not found for {dataset}-{network}-{mode}-bs{bs}-lr_w{lr}-lr_s{lr}-seed{seed}:")
  86. if not os.path.exists(test_file):
  87. print(f" Missing: {test_file}")
  88. if not os.path.exists(train_file):
  89. print(f" Missing: {train_file}")
  90. except Exception as e:
  91. print(f"Error processing {dataset}-{network}-{mode}-bs{bs}-lr_w{lr}-lr_s{lr}-seed{seed}: {e}")
  92. continue
  93. if combo_accs:
  94. avg_acc = np.mean(combo_accs)
  95. avg_aae = np.mean(combo_aaes)
  96. hyperparam_results[combo_key] = {
  97. 'avg_acc': avg_acc,
  98. 'avg_aae': avg_aae,
  99. 'results': combo_results
  100. }
  101. # Find the best hyperparameter combination based on average accuracy
  102. if hyperparam_results:
  103. best_combo = max(hyperparam_results.keys(), key=lambda k: hyperparam_results[k]['avg_acc'])
  104. best_combo_data = hyperparam_results[best_combo]
  105. # Find the seed with highest accuracy for this best combination
  106. best_seed_result = max(best_combo_data['results'], key=lambda r: r['acc'])
  107. best_result = {
  108. 'test': best_seed_result['test'],
  109. 'train': best_seed_result['train']
  110. }
  111. best_params = {
  112. 'lr_w': best_combo[0],
  113. 'lr_special': best_combo[1],
  114. 'batch_size': best_combo[2],
  115. 'seed': best_seed_result['seed'],
  116. 'mode': mode,
  117. 'rectify': rectify
  118. }
  119. return best_result, best_params, best_combo_data['avg_acc'], best_combo_data['avg_aae']
  120. return None, None, 0, np.nan
  121. # %%
  122. def find_best_untuned_result(dataset, network, scheduler, epochs,seeds,
  123. weight_decay=5e-4,
  124. lrs_w=[0.025, 1e-2, 1e-3], batch_sizes=[32, 64, 128]):
  125. """
  126. Find the best hyperparameter combination for untuned baseline based on highest test accuracy.
  127. Returns the best result data, parameters, and averaged metrics across seeds.
  128. """
  129. best_result = None
  130. best_params = None
  131. # First find the best hyperparameter combination by averaging across seeds
  132. hyperparam_results = {}
  133. for lr_w in lrs_w:
  134. for bs in batch_sizes:
  135. combo_key = (lr_w, bs)
  136. combo_accs = []
  137. combo_aaes = []
  138. combo_results = []
  139. for seed in seeds:
  140. try:
  141. # Get file path for no tuner
  142. result_path, _ = get_fname_tuner(
  143. dataset, network, bs, lr_w, None, epochs, seed,
  144. False, None, False, None, weight_decay,
  145. False, False, scheduler_w=scheduler
  146. )
  147. # Load test and train results
  148. test_file = os.path.join(result_path, 'res_test.mat')
  149. train_file = os.path.join(result_path, 'res_train.mat')
  150. if os.path.exists(test_file) and os.path.exists(train_file):
  151. test_data = recover_dict(loadmat(test_file))
  152. train_data = recover_dict(loadmat(train_file))
  153. max_test_acc = np.max(test_data["top1"])
  154. test_aae = test_data.get("aae", float(np.trapz(test_data["top1"])/(len(test_data["top1"])-1)))
  155. combo_accs.append(max_test_acc)
  156. combo_aaes.append(test_aae)
  157. combo_results.append({
  158. 'test': test_data,
  159. 'train': train_data,
  160. 'seed': seed,
  161. 'acc': max_test_acc
  162. })
  163. else:
  164. print(f"Files not found for {dataset}-{network}-no_tuner-bs{bs}-lr_w{lr_w}-seed{seed}:")
  165. if not os.path.exists(test_file):
  166. print(f" Missing: {test_file}")
  167. if not os.path.exists(train_file):
  168. print(f" Missing: {train_file}")
  169. except Exception as e:
  170. print(f"Error processing {dataset}-{network}-no_tuner-bs{bs}-lr_w{lr_w}-seed{seed}: {e}")
  171. continue
  172. if combo_accs:
  173. avg_acc = np.mean(combo_accs)
  174. avg_aae = np.mean(combo_aaes)
  175. hyperparam_results[combo_key] = {
  176. 'avg_acc': avg_acc,
  177. 'avg_aae': avg_aae,
  178. 'results': combo_results
  179. }
  180. # Find the best hyperparameter combination based on average accuracy
  181. if hyperparam_results:
  182. best_combo = max(hyperparam_results.keys(), key=lambda k: hyperparam_results[k]['avg_acc'])
  183. best_combo_data = hyperparam_results[best_combo]
  184. # Find the seed with highest accuracy for this best combination
  185. best_seed_result = max(best_combo_data['results'], key=lambda r: r['acc'])
  186. best_result = {
  187. 'test': best_seed_result['test'],
  188. 'train': best_seed_result['train']
  189. }
  190. best_params = {
  191. 'lr_w': best_combo[0],
  192. 'batch_size': best_combo[1],
  193. 'seed': best_seed_result['seed'],
  194. 'mode': 'no_tuner'
  195. }
  196. return best_result, best_params, best_combo_data['avg_acc'], best_combo_data['avg_aae']
  197. return None, None, 0, np.nan
  198. # %%
  199. def plot_comparison_figure(dataset, network, modes, init_methods, epochs, seeds, scheduler,
  200. legend_fs=14,
  201. single_param=False, rectify=False, save_path=None, colors=None, **kwargs):
  202. """
  203. Plot comparison figure with 4 sub-panels showing train/test acc/loss curves.
  204. Each panel shows curves for different modes and rectify options.
  205. Parameters:
  206. - colors: list of colors to use for each curve. If None, uses default colors.
  207. """
  208. fig, axes = plt.subplots(1, 1, figsize=(8, 6))
  209. # Define colors for different modes
  210. if colors is None:
  211. colors = ['blue', 'red', 'green', 'orange', 'purple', 'brown', 'pink', 'gray', 'cyan', 'magenta', 'gold']
  212. # Collect all results
  213. all_results = []
  214. # Get best results for each mode and rectify combination
  215. for mode,init_method in zip(modes,init_methods):
  216. # for rectify in rectify_options:
  217. result, params, avg_acc, avg_aae = find_best_tuned_result(
  218. dataset, network, mode, scheduler, epochs, seeds,init_method, **kwargs
  219. )
  220. if result is not None:
  221. all_results.append({
  222. 'result': result,
  223. 'params': params,
  224. 'avg_acc': avg_acc,
  225. 'avg_aae': avg_aae,
  226. 'label': f'{mode}',
  227. 'mode': mode,
  228. 'rectify': rectify
  229. })
  230. # Get best untuned result
  231. untuned_result, untuned_params, untuned_avg_acc, untuned_avg_aae = find_best_untuned_result(
  232. dataset, network, scheduler, epochs, seeds, **kwargs
  233. )
  234. if untuned_result is not None:
  235. all_results.append({
  236. 'result': untuned_result,
  237. 'params': untuned_params,
  238. 'avg_acc': untuned_avg_acc,
  239. 'avg_aae': untuned_avg_aae,
  240. 'label': 'no_tuner',
  241. 'mode': 'no_tuner',
  242. 'rectify': False
  243. })
  244. # Plot curves
  245. for i, res_data in enumerate(all_results):
  246. color = colors[i % len(colors)]
  247. result = res_data['result']
  248. label = res_data['label']
  249. # change label names
  250. if label == 'CM-GLLF-all-valuewise-minmax-mapping':
  251. label = 'CMG'
  252. elif label == 'linear-valuewise-positive':
  253. label = 'Linear'
  254. elif label == 'cubic-valuewise-positive':
  255. label = 'Cubic'
  256. elif label == 'adaptive-sigmoid-valuewise':
  257. label = 'Adaptive Sigmoid'
  258. elif label == 'adaptive-tanh-valuewise':
  259. label = 'Adaptive Tanh'
  260. elif label == 'no_tuner':
  261. label = 'Original MLP'
  262. elif label == 'srelu_valuewise_positive':
  263. label = 'SReLU'
  264. avg_acc = res_data['avg_acc']
  265. avg_aae = res_data['avg_aae']
  266. # Cut data at max_epochs if specified
  267. test_top1_plot = result['test']['top1']
  268. # Test accuracy
  269. axes.plot(test_top1_plot, color=color,
  270. label=f'{label}\n (ACC: {avg_acc:.2f} AAE: {avg_aae:.2f})')
  271. # Set titles and labels
  272. axes.set_title(dataset,fontsize=18)
  273. axes.set_xlabel('Epoch',fontsize=14)
  274. axes.set_ylabel('Accuracy',fontsize=14)
  275. # Improve x-axis ticks to show meaningful intervals including endpoint
  276. # Set x-ticks to show 0, quarter points, half, three-quarters, and endpoint
  277. tick_positions = [pos -1 for pos in [1,25,50,75,100,125,150]]
  278. tick_labels = [str(pos+1) for pos in tick_positions]
  279. # Adjust the last label to show the actual epoch number (1-indexed)
  280. axes.set_xticks(tick_positions)
  281. axes.set_xticklabels(tick_labels)
  282. # axes.legend(bbox_to_anchor=(1.05, 1), loc='upper left')
  283. axes.legend(fontsize=legend_fs)
  284. axes.grid(True)
  285. plt.tight_layout()
  286. if save_path:
  287. plt.savefig(save_path, dpi=300, bbox_inches='tight')
  288. print(f"Figure saved as: {save_path}")
  289. plt.show()
  290. # Print summary of best results
  291. print(f"\n=== Best Results Summary for {dataset} - single_param={single_param} ===")
  292. for res_data in all_results:
  293. params = res_data['params']
  294. avg_acc = res_data['avg_acc']
  295. label = res_data['label']
  296. print(f"{label}: {avg_acc:.4f} - {params}")
  297. # %%
  298. def plot_comparison_figure_with_annotations(dataset, network, modes, init_methods, epochs, seeds, scheduler, draw_mode,
  299. legend_fs=14, single_param=False, rectify=False, save_path=None, colors=None, **kwargs):
  300. """
  301. Plot comparison figure with accuracy annotations at the end of curves.
  302. Similar to plot_comparison_figure but adds accuracy values next to curve endpoints.
  303. """
  304. fig, axes = plt.subplots(1, 1, figsize=(8, 6))
  305. # Define colors for different modes
  306. if colors is None:
  307. colors = ['blue', 'red', 'green', 'orange', 'purple', 'brown', 'pink', 'gray', 'cyan', 'magenta', 'gold']
  308. # Collect all results and annotations
  309. all_results = []
  310. annotations = []
  311. # Get best results for each mode
  312. for mode, init_method in zip(modes, init_methods):
  313. result, params, avg_acc, avg_aae = find_best_tuned_result(
  314. dataset, network, mode, scheduler, epochs, seeds, init_method, **kwargs
  315. )
  316. if result is not None:
  317. all_results.append({
  318. 'result': result,
  319. 'params': params,
  320. 'avg_acc': avg_acc,
  321. 'avg_aae': avg_aae,
  322. 'label': f'{mode}',
  323. 'mode': mode,
  324. 'rectify': rectify
  325. })
  326. # Get best untuned result
  327. untuned_result, untuned_params, untuned_avg_acc, untuned_avg_aae = find_best_untuned_result(
  328. dataset, network, scheduler, epochs, seeds, **kwargs
  329. )
  330. if untuned_result is not None:
  331. all_results.append({
  332. 'result': untuned_result,
  333. 'params': untuned_params,
  334. 'avg_acc': untuned_avg_acc,
  335. 'avg_aae': untuned_avg_aae,
  336. 'label': 'no_tuner',
  337. 'mode': 'no_tuner',
  338. 'rectify': False
  339. })
  340. # Plot curves and collect annotation data
  341. for i, res_data in enumerate(all_results):
  342. color = colors[i % len(colors)]
  343. result = res_data['result']
  344. label = res_data['label']
  345. # Change label names
  346. if label == 'CM-GLLF-all-valuewise-minmax-mapping':
  347. label = 'CMG'
  348. elif label == 'linear-valuewise-positive':
  349. label = 'Linear'
  350. elif label == 'cubic-valuewise-positive':
  351. label = 'Cubic'
  352. elif label == 'adaptive-sigmoid-valuewise':
  353. label = 'Adaptive Sigmoid'
  354. elif label == 'adaptive-tanh-valuewise':
  355. label = 'Adaptive Tanh'
  356. elif label == 'no_tuner':
  357. label = 'Original MLP'
  358. elif label == 'srelu_valuewise_positive':
  359. label = 'SReLU'
  360. avg_acc = res_data['avg_acc']
  361. avg_aae = res_data['avg_aae']
  362. # Get accuracy data
  363. test_top1_plot = result['test']['top1']
  364. # Plot curve
  365. axes.plot(test_top1_plot, color=color,
  366. label=f'{label}\n (ACC: {avg_acc:.2f} AAE: {avg_aae:.2f})')
  367. # Store annotation data
  368. annotations.append({
  369. 'x': len(test_top1_plot) - 1,
  370. 'y': test_top1_plot[-1],
  371. 'acc': avg_acc,
  372. 'color': color,
  373. 'text': f'{avg_acc:.2f}',
  374. 'label': label
  375. })
  376. # Extend x-axis to provide space for annotations
  377. axes.set_xlim(left=0, right=epochs + 20)
  378. if draw_mode == 'cmg+linear+adaptive' and dataset == 'CIFAR10':
  379. axes.set_ylim(top=68.5)
  380. elif draw_mode == 'cmg+linear+adaptive' and dataset == 'CIFAR100':
  381. axes.set_ylim(top=43)
  382. # Sort annotations by accuracy (highest first) for positioning
  383. annotations.sort(key=lambda x: x['acc'], reverse=True)
  384. # Add annotations with non-overlapping positioning
  385. y_range = axes.get_ylim()[1] - axes.get_ylim()[0]
  386. base_spacing = 0.008 * y_range
  387. for i, ann in enumerate(annotations):
  388. # Position text with vertical spacing to avoid overlap
  389. if draw_mode == 'cmg+linear':
  390. if dataset=='CIFAR100' and ann['label'] == 'CMG':
  391. adjusted_y = ann['y'] + (i+2.5) * base_spacing
  392. else:
  393. adjusted_y = ann['y']
  394. elif draw_mode == 'cmg+linear+adaptive':
  395. if dataset == 'CIFAR10' and ann['label'] == 'CMG':
  396. adjusted_y = ann['y'] + (i+4) * base_spacing
  397. elif dataset == 'CIFAR10' and ann['label'] == 'Linear':
  398. adjusted_y = ann['y'] + (i+3) * base_spacing
  399. elif dataset == 'CIFAR100' and ann['label'] == 'CMG':
  400. adjusted_y = ann['y'] + (i+6.5) * base_spacing
  401. elif dataset == 'CIFAR100' and ann['label'] == 'Linear':
  402. adjusted_y = ann['y'] + (i+3) * base_spacing
  403. elif dataset == 'CIFAR100' and ann['label'] == 'Original MLP':
  404. adjusted_y = ann['y'] + (i+1) * base_spacing
  405. else:
  406. adjusted_y = ann['y']
  407. elif draw_mode == 'cmg+srelu':
  408. if dataset=='CIFAR100' and ann['label'] == 'SReLU':
  409. adjusted_y = ann['y'] + (i+2) * base_spacing
  410. else:
  411. adjusted_y = ann['y']
  412. text_x = ann['x'] + 3
  413. axes.annotate(ann['text'],
  414. xy=(ann['x'], ann['y']),
  415. xytext=(text_x, adjusted_y),
  416. color=ann['color'],
  417. fontsize=14,
  418. fontweight='bold',
  419. verticalalignment='center')
  420. # Set titles and labels
  421. axes.set_title(dataset, fontsize=18)
  422. axes.set_xlabel('Epoch', fontsize=14)
  423. axes.set_ylabel('Accuracy', fontsize=14)
  424. # Set x-ticks
  425. tick_positions = [pos - 1 for pos in [1, 25, 50, 75, 100, 125, 150]]
  426. tick_labels = [str(pos + 1) for pos in tick_positions]
  427. axes.set_xticks(tick_positions)
  428. axes.set_xticklabels(tick_labels)
  429. axes.legend(fontsize=legend_fs)
  430. axes.grid(True)
  431. plt.tight_layout()
  432. if save_path:
  433. plt.savefig(save_path, dpi=300, bbox_inches='tight')
  434. print(f"Figure saved as: {save_path}")
  435. plt.show()
  436. # %% [markdown]
  437. # # get test acc curves
  438. # %%
  439. # Generate comparison plots for each dataset and single_param option
  440. datasets = ['CIFAR10','CIFAR100']
  441. networks = ['MLP_CIFAR10_A','MLP_CIFAR100_A']
  442. modes = ['CM-GLLF-all-valuewise-minmax-mapping',
  443. 'linear-valuewise-positive',
  444. ]
  445. init_methods = ['uniform'] * len(modes)
  446. # seeds = [0,1,2]
  447. seeds = [0,1,2]
  448. scheduler='power_to_linear'
  449. for dataset, network in zip(datasets, networks):
  450. for single_param in [False]:
  451. plot_comparison_figure_with_annotations(
  452. dataset=dataset,
  453. network=network,
  454. modes=modes,
  455. init_methods=init_methods,
  456. seeds=seeds,
  457. epochs=150,
  458. scheduler=scheduler,
  459. colors=["#2bc62b", 'grey', "#060606"], # Custom colors
  460. draw_mode='cmg+linear',
  461. )
  462. # %%
  463. # Generate comparison plots for each dataset and single_param option
  464. datasets = ['CIFAR10','CIFAR100']
  465. networks = ['MLP_CIFAR10_A','MLP_CIFAR100_A']
  466. modes = ['CM-GLLF-all-valuewise-minmax-mapping',
  467. 'linear-valuewise-positive',
  468. 'cubic-valuewise-positive',
  469. 'adaptive-sigmoid-valuewise',
  470. 'adaptive-tanh-valuewise',
  471. # 'srelu_valuewise_positive',
  472. ]
  473. init_methods = ['uniform'] * len(modes)
  474. # seeds = [0,1,2]
  475. seeds = [0,1,2]
  476. scheduler='power_to_linear'
  477. for dataset, network in zip(datasets, networks):
  478. for single_param in [False]:
  479. plot_comparison_figure_with_annotations(
  480. dataset=dataset,
  481. network=network,
  482. modes=modes,
  483. init_methods=init_methods,
  484. seeds=seeds,
  485. epochs=150,
  486. scheduler=scheduler,
  487. colors=["#2bc62b", 'grey', '#ff7f0e', "#11cfe8","#7e08ec", "#060606"], # Custom colors
  488. legend_fs=11,
  489. draw_mode='cmg+linear+adaptive',
  490. )
  491. # %%
  492. # Generate comparison plots for each dataset and single_param option
  493. datasets = ['CIFAR10','CIFAR100']
  494. networks = ['MLP_CIFAR10_A','MLP_CIFAR100_A']
  495. modes = ['CM-GLLF-all-valuewise-minmax-mapping',
  496. # 'linear-valuewise-positive',
  497. # 'cubic-valuewise-positive',
  498. # 'adaptive-sigmoid-valuewise',
  499. # 'adaptive-tanh-valuewise',
  500. 'srelu_valuewise_positive',
  501. ]
  502. init_methods = ['uniform'] * len(modes)
  503. # seeds = [0,1,2]
  504. seeds = [0,1,2]
  505. scheduler='power_to_linear'
  506. for dataset, network in zip(datasets, networks):
  507. for single_param in [False]:
  508. plot_comparison_figure_with_annotations(
  509. dataset=dataset,
  510. network=network,
  511. modes=modes,
  512. init_methods=init_methods,
  513. seeds=seeds,
  514. epochs=150,
  515. scheduler=scheduler,
  516. colors=["#2bc62b","#fb0eff", "#060606"], # Custom colors
  517. draw_mode='cmg+srelu',
  518. )
  519. # %% [markdown]
  520. # # compare schedulers
  521. # %%
  522. # first plot the scheduler shapes
  523. # Scheduler factory mirroring train_tuner.py
  524. import torch
  525. def create_scheduler(optimizer, scheduler_type, epochs):
  526. if scheduler_type == 'linear':
  527. return torch.optim.lr_scheduler.LinearLR(
  528. optimizer, start_factor=1.0, end_factor=0.01, total_iters=int(epochs * 0.9)
  529. )
  530. elif scheduler_type == 'exponential_95':
  531. stop_epoch = int(epochs * 0.9)
  532. return torch.optim.lr_scheduler.LambdaLR(
  533. optimizer, lambda epoch: 0.95 ** min(epoch, stop_epoch)
  534. )
  535. elif scheduler_type == 'exponential_97':
  536. stop_epoch = int(epochs * 0.9)
  537. return torch.optim.lr_scheduler.LambdaLR(
  538. optimizer, lambda epoch: 0.97 ** min(epoch, stop_epoch)
  539. )
  540. elif scheduler_type == 'exponential_99':
  541. stop_epoch = int(epochs * 0.9)
  542. return torch.optim.lr_scheduler.LambdaLR(
  543. optimizer, lambda epoch: 0.99 ** min(epoch, stop_epoch)
  544. )
  545. elif scheduler_type == 'slow_to_fast':
  546. # Piecewise multiplicative decay: slow early, faster later, capped at 90% of epochs
  547. stop_epoch = int(epochs * 0.9)
  548. half = max(1, stop_epoch // 2)
  549. def factor(epoch):
  550. e = min(epoch, stop_epoch)
  551. if e <= half:
  552. return (0.995) ** e
  553. else:
  554. # First half at 0.995, remaining at 0.95
  555. return (0.995 ** half) * (0.95 ** (e - half))
  556. return torch.optim.lr_scheduler.LambdaLR(optimizer, factor)
  557. elif scheduler_type == 'exp_symmetric_95':
  558. # Smooth slow-to-fast curve symmetric (in shape) to exponential decay, capped at 90% of epochs
  559. stop_epoch = int(epochs * 0.9)
  560. q = 2.0 # convex warp power; q>1 => slow start then faster end
  561. gamma = 0.95
  562. def factor(epoch):
  563. e = min(epoch, stop_epoch)
  564. if stop_epoch <= 0:
  565. return 1.0
  566. t = e / float(stop_epoch)
  567. g = t ** q # convex time-warp
  568. return gamma ** (stop_epoch * g)
  569. return torch.optim.lr_scheduler.LambdaLR(optimizer, factor)
  570. elif scheduler_type == 'exponential_to_linear':
  571. # Exponential decay that reaches 0.01 at 90% of epochs (same as linear)
  572. stop_epoch = int(epochs * 0.9)
  573. # Calculate gamma such that gamma^stop_epoch = 0.01
  574. gamma = 0.01 ** (1.0 / stop_epoch)
  575. return torch.optim.lr_scheduler.LambdaLR(
  576. optimizer, lambda epoch: gamma ** min(epoch, stop_epoch)
  577. )
  578. elif scheduler_type == 'power_to_linear':
  579. stop_epoch = int(epochs * 0.9)
  580. power = 0.5 # You can adjust this power as needed
  581. def power_decay(epoch):
  582. if epoch >= stop_epoch:
  583. return 0.01
  584. t = epoch / stop_epoch
  585. return 0.01 + (1.0 - 0.01) * ((1 - t) ** power)
  586. return torch.optim.lr_scheduler.LambdaLR(optimizer, power_decay)
  587. else:
  588. raise ValueError(f'Unknown scheduler type: {scheduler_type}')
  589. # %%
  590. # Simulate and plot learning rates
  591. epochs = 150
  592. base_lr = 0.01 # choose a representative base learning rate
  593. scheduler_types = ['linear', 'exponential_to_linear', 'power_to_linear']
  594. lr_histories = {}
  595. for sched_type in scheduler_types:
  596. # fresh model/optimizer per schedule to avoid cross-contamination
  597. model = torch.nn.Linear(1, 1)
  598. optimizer = torch.optim.SGD(model.parameters(), lr=base_lr)
  599. scheduler = create_scheduler(optimizer, sched_type, epochs)
  600. lrs = []
  601. for epoch in range(epochs):
  602. # record LR used this epoch (scheduler is stepped at the end of each epoch in train_tuner.py)
  603. lrs.append(optimizer.param_groups[0]['lr'])
  604. scheduler.step()
  605. lr_histories[sched_type] = lrs
  606. # Plot
  607. plt.figure(figsize=(8, 6))
  608. for name, lrs in lr_histories.items():
  609. if name=='power_to_linear':
  610. name='Supra-linear Scheduler'
  611. color="#2bc62b"
  612. elif name=='exponential_to_linear':
  613. name='Sub-linear Scheduler'
  614. color='red'
  615. elif name=='linear':
  616. name='Linear Scheduler'
  617. color='blue'
  618. plt.plot(range(epochs), lrs, label=name,color=color,linewidth=2)
  619. plt.xlabel('Epoch', fontsize=14)
  620. plt.ylabel('Learning Rate', fontsize=14)
  621. plt.title('LR Schedulers', fontsize=18)
  622. plt.grid(True, alpha=0.3)
  623. plt.legend(fontsize=14)
  624. # Set x-ticks to match other figures
  625. tick_positions = [pos -1 for pos in [1,25,50,75,100,125,150]]
  626. tick_labels = [str(pos+1) for pos in tick_positions]
  627. plt.xticks(tick_positions, tick_labels)
  628. os.makedirs('plots', exist_ok=True)
  629. plt.tight_layout()
  630. plt.savefig('plots/lr_schedules_200epochs.png', dpi=180)
  631. plt.show()
  632. # %%
  633. def plot_combined_scheduler_comparison(dataset, network, mode, init_method, epochs, seeds, schedulers,
  634. single_param=False, rectify=False, save_path=None, colors=None, **kwargs):
  635. """
  636. Plot comparison figure showing both tuned (solid lines) and no-tune (dashed lines) results
  637. for different schedulers in a single plot.
  638. """
  639. fig, ax = plt.subplots(1, 1, figsize=(8, 6))
  640. # Define colors for different schedulers
  641. if colors is None:
  642. colors = ['blue', 'red', 'green', 'orange', 'purple']
  643. # Collect results for both tuned and no-tune
  644. annotations = [] # Store annotation data for positioning
  645. for i, scheduler in enumerate(schedulers):
  646. color = colors[i % len(colors)]
  647. # Get tuned results
  648. tuned_result, _, tuned_avg_acc, tuned_avg_aae = find_best_tuned_result(
  649. dataset, network, mode, scheduler, epochs, seeds, init_method, **kwargs
  650. )
  651. # Get no-tune results
  652. notune_result, _, notune_avg_acc, notune_avg_aae = find_best_untuned_result(
  653. dataset, network, scheduler, epochs, seeds, **kwargs
  654. )
  655. # Clean scheduler name
  656. scheduler_label = scheduler
  657. if scheduler == 'linear':
  658. scheduler_label = 'Linear Scheduler'
  659. elif scheduler == 'exponential_to_linear':
  660. scheduler_label = 'Sub-linear Scheduler'
  661. elif scheduler == 'power_to_linear':
  662. scheduler_label = 'Supra-linear Scheduler'
  663. # Plot tuned results (solid line)
  664. if tuned_result is not None:
  665. test_top1_tuned = tuned_result['test']['top1']
  666. ax.plot(test_top1_tuned, color=color, linestyle='-', linewidth=2,
  667. label=f'{scheduler_label} (CMG)\n (ACC:{tuned_avg_acc:.2f} AAE:{tuned_avg_aae:.2f})')
  668. # Store annotation data for tuned
  669. annotations.append({
  670. 'x': len(test_top1_tuned) - 1,
  671. 'y': test_top1_tuned[-1],
  672. 'acc': tuned_avg_acc,
  673. 'color': color,
  674. 'text': f'{tuned_avg_acc:.2f}',
  675. 'scheduler': scheduler,
  676. 'type': 'tuned'
  677. })
  678. # Plot no-tune results (dashed line)
  679. if notune_result is not None:
  680. test_top1_notune = notune_result['test']['top1']
  681. ax.plot(test_top1_notune, color=color, linestyle=(0, (2, 2)), linewidth=2,
  682. label=f'{scheduler_label} (Original MLP)\n (ACC:{notune_avg_acc:.2f} AAE:{notune_avg_aae:.2f})')
  683. # Store annotation data for no-tune
  684. annotations.append({
  685. 'x': len(test_top1_notune) - 1,
  686. 'y': test_top1_notune[-1],
  687. 'acc': notune_avg_acc,
  688. 'color': color,
  689. 'text': f'{notune_avg_acc:.2f}',
  690. 'scheduler': scheduler,
  691. 'type': 'notune'
  692. })
  693. # Extend x-axis to provide more space for annotations
  694. ax.set_xlim(left=0, right=epochs + 20) # Add 15 units of padding to the right
  695. # Sort annotations by accuracy (highest first) for proper vertical positioning
  696. annotations.sort(key=lambda x: x['acc'], reverse=True)
  697. # Add annotations with improved positioning logic
  698. y_range = ax.get_ylim()[1] - ax.get_ylim()[0]
  699. base_spacing = 0.008 * y_range # Base vertical spacing between annotations
  700. for i, ann in enumerate(annotations):
  701. # Calculate text position with special handling for supra-linear scheduler
  702. base_y = ann['y']
  703. # Give supra-linear (power_to_linear) extra height to avoid overlap
  704. if ann['scheduler'] == 'power_to_linear' and ann['type'] == 'tuned':
  705. # Position supra-linear tuned result higher
  706. adjusted_y = base_y + (i + 3) * base_spacing
  707. elif ann['scheduler'] == 'power_to_linear' and ann['type'] == 'notune':
  708. # Position supra-linear tuned result higher
  709. adjusted_y = base_y + (i) * base_spacing
  710. else:
  711. # adjusted_y = base_y + i * base_spacing
  712. adjusted_y = base_y
  713. # Position text further to the right with more padding
  714. text_x = ann['x'] + 4
  715. ax.annotate(ann['text'],
  716. xy=(ann['x'], ann['y']),
  717. xytext=(text_x, adjusted_y),
  718. color=ann['color'],
  719. fontsize=14,
  720. fontweight='bold',
  721. verticalalignment='center')
  722. # Set titles and labels
  723. ax.set_title(f'Scheduler Accuracy Comparison', fontsize=18)
  724. ax.set_xlabel('Epoch', fontsize=14)
  725. ax.set_ylabel('Accuracy', fontsize=14)
  726. # Set x-ticks to match other figures
  727. tick_positions = [pos -1 for pos in [1,25,50,75,100,125,150]]
  728. tick_labels = [str(pos+1) for pos in tick_positions]
  729. ax.set_xticks(tick_positions)
  730. ax.set_xticklabels(tick_labels)
  731. ax.legend(fontsize=11)
  732. ax.grid(True)
  733. plt.tight_layout()
  734. if save_path:
  735. plt.savefig(save_path, dpi=300, bbox_inches='tight')
  736. print(f"Figure saved as: {save_path}")
  737. plt.show()
  738. # %%
  739. # Generate combined scheduler comparison plot
  740. plot_combined_scheduler_comparison(
  741. dataset="CIFAR100",
  742. network="MLP_CIFAR100_A",
  743. mode='CM-GLLF-all-valuewise-minmax-mapping',
  744. init_method='uniform',
  745. epochs=150,
  746. seeds=[0,1,2],
  747. schedulers=['linear', 'exponential_to_linear', 'power_to_linear'],
  748. colors=['blue', 'red', "#2bc62b"]
  749. )
  750. # %% [markdown]
  751. # # Qualitative test
  752. # %%
  753. from models_tuner import MLP_CIFAR10_A, MLP_CIFAR100_A
  754. import torch
  755. def get_best_model(dataset, network, mode, scheduler, epochs, seeds, init_method):
  756. _, best_params, _, _ = find_best_tuned_result(dataset, network, mode, scheduler, epochs, seeds, init_method)
  757. # get the path
  758. result_path, _ = get_fname_tuner(
  759. dataset, network, best_params['batch_size'], best_params['lr_w'],
  760. best_params['lr_special'], epochs, best_params['seed'],
  761. True, mode, True, init_method, 5e-4,
  762. False, False, scheduler_w=scheduler
  763. )
  764. model_path = os.path.join(result_path, 'best_model.pth')
  765. if os.path.exists(model_path):
  766. # Create the model architecture
  767. if network == 'MLP_CIFAR10_A':
  768. from models_tuner import MLP_CIFAR10_A
  769. model = MLP_CIFAR10_A(3072, 10, True, mode, True, init_method, dataset, single_param=False)
  770. elif network == 'MLP_CIFAR100_A':
  771. from models_tuner import MLP_CIFAR100_A
  772. model = MLP_CIFAR100_A(3072, 100, True, mode, True, init_method, dataset, single_param=False)
  773. else:
  774. raise FileNotFoundError(f"Model file not found: {model_path}")
  775. # Load the trained parameters
  776. state_dict = torch.load(model_path, map_location='cpu')
  777. model.load_state_dict(state_dict)
  778. results = {}
  779. # Extract tuner parameters
  780. if hasattr(model, 'tuner'):
  781. tuner = model.tuner
  782. tuner_class_name = tuner.__class__.__name__
  783. print(f"Detected tuner type: {tuner_class_name}")
  784. results['tuner_type'] = tuner_class_name
  785. results['tuner'] = tuner # Store the tuner itself for direct use
  786. if tuner_class_name == 'CM_GLLF_tuner_flexible':
  787. if hasattr(tuner, 'mu_raw'):
  788. mu_raw = tuner.mu_raw.detach().cpu()
  789. # Transform based on tuner mode
  790. if tuner.mode == 'all':
  791. mu = torch.sigmoid(mu_raw)
  792. elif tuner.mode == 'logistic':
  793. mu = 0.5 * torch.sigmoid(mu_raw)
  794. elif tuner.mode == 'logit':
  795. mu = 0.5 + 0.5 * torch.sigmoid(mu_raw)
  796. else:
  797. raise ValueError(f"Unknown mode: {tuner.mode}")
  798. results['mu_raw'] = mu_raw
  799. results['mu'] = mu
  800. print(f"Parameter 'mu_raw' shape: {results['mu_raw'].shape}")
  801. print(f"Parameter 'mu' shape: {results['mu'].shape}")
  802. print(f"mu range: [{results['mu'].min():.4f}, {results['mu'].max():.4f}]")
  803. if hasattr(tuner, 'I_raw'):
  804. I_raw = tuner.I_raw.detach().cpu()
  805. # Transform: I = sigmoid(I_raw)
  806. I = torch.sigmoid(I_raw)
  807. results['I_raw'] = I_raw
  808. results['I'] = I
  809. print(f"Parameter 'I_raw' shape: {results['I_raw'].shape}")
  810. print(f"Parameter 'I' shape: {results['I'].shape}")
  811. print(f"I range: [{results['I'].min():.4f}, {results['I'].max():.4f}]")
  812. # Handle tuners with 'a' and 'b' parameters
  813. elif tuner_class_name in ['Linear_tuner_valuewise', 'Cubic_tuner_valuewise', 'AdaptiveClassic_tuner_valuewise']:
  814. # Handle 'a' parameter
  815. if hasattr(tuner, 'a_raw'):
  816. a_raw = tuner.a_raw.detach().cpu()
  817. if tuner.positive_a:
  818. a = torch.exp(a_raw) # For positive_a=True, a = exp(a_raw)
  819. else:
  820. a = a_raw # For positive_a=False, a is used directly
  821. results['a_raw'] = a_raw
  822. results['a'] = a
  823. print(f"Parameter 'a_raw' shape: {results['a_raw'].shape}")
  824. print(f"Parameter 'a' shape: {results['a'].shape}")
  825. print(f"a range: [{results['a'].min():.4f}, {results['a'].max():.4f}]")
  826. elif hasattr(tuner, 'a'):
  827. a = tuner.a.detach().cpu()
  828. results['a'] = a
  829. print(f"Parameter 'a' shape: {results['a'].shape}")
  830. print(f"a range: [{results['a'].min():.4f}, {results['a'].max():.4f}]")
  831. # Handle 'b' parameter
  832. if hasattr(tuner, 'b'):
  833. b = tuner.b.detach().cpu()
  834. results['b'] = b
  835. print(f"Parameter 'b' shape: {results['b'].shape}")
  836. print(f"b range: [{results['b'].min():.4f}, {results['b'].max():.4f}]")
  837. # Handle SReLU tuners with different parameter names
  838. elif 'SReLU' in tuner_class_name:
  839. if hasattr(tuner, 't_left'):
  840. results['t_left'] = tuner.t_left.detach().cpu()
  841. print(f"Parameter 't_left' shape: {results['t_left'].shape}")
  842. print(f"t_left range: [{results['t_left'].min():.4f}, {results['t_left'].max():.4f}]")
  843. if hasattr(tuner, 't_right'):
  844. results['t_right'] = tuner.t_right.detach().cpu()
  845. print(f"Parameter 't_right' shape: {results['t_right'].shape}")
  846. print(f"t_right range: [{results['t_right'].min():.4f}, {results['t_right'].max():.4f}]")
  847. # Handle slope parameters (a_left, a_right)
  848. if hasattr(tuner, 'a_left_raw'):
  849. a_left_raw = tuner.a_left_raw.detach().cpu()
  850. a_left = torch.exp(a_left_raw) if tuner.positive_slopes else a_left_raw
  851. results['a_left_raw'] = a_left_raw
  852. results['a_left'] = a_left
  853. print(f"Parameter 'a_left' range: [{results['a_left'].min():.4f}, {results['a_left'].max():.4f}]")
  854. elif hasattr(tuner, 'a_left'):
  855. results['a_left'] = tuner.a_left.detach().cpu()
  856. print(f"Parameter 'a_left' range: [{results['a_left'].min():.4f}, {results['a_left'].max():.4f}]")
  857. if hasattr(tuner, 'a_right_raw'):
  858. a_right_raw = tuner.a_right_raw.detach().cpu()
  859. a_right = torch.exp(a_right_raw) if tuner.positive_slopes else a_right_raw
  860. results['a_right_raw'] = a_right_raw
  861. results['a_right'] = a_right
  862. print(f"Parameter 'a_right' range: [{results['a_right'].min():.4f}, {results['a_right'].max():.4f}]")
  863. elif hasattr(tuner, 'a_right'):
  864. results['a_right'] = tuner.a_right.detach().cpu()
  865. print(f"Parameter 'a_right' range: [{results['a_right'].min():.4f}, {results['a_right'].max():.4f}]")
  866. else:
  867. print(f"Warning: Unknown tuner type '{tuner_class_name}'. Only basic info extracted.")
  868. else:
  869. raise ValueError("The model to inspect should have attr tuner")
  870. return results
  871. # %%
  872. def plot_tuner_heatmaps(mu_tensor, I_tensor, height=32, width=32, channels=3,
  873. channel_names=None, cmap=None, channel_cmaps=None,
  874. title=None, figsize=(12, 8)):
  875. """Visualize tuner parameters as 2xN heatmaps (mu row, I row, channel columns).
  876. Parameters
  877. ----------
  878. mu_tensor : torch.Tensor or array-like
  879. Tuner parameter, expected as a flat vector of length ``height*width*channels``
  880. or already reshaped to (channels, height, width).
  881. I_tensor : torch.Tensor or array-like
  882. Tuner parameter I with the same layout constraints as ``mu_tensor``.
  883. height, width : int
  884. Spatial dimensions for the resulting heatmaps (default: 32x32 for CIFAR).
  885. channels : int
  886. Number of color channels to visualize (default: 3 for RGB).
  887. channel_names : list[str] | None
  888. Optional list of labels per channel; defaults to ['R', 'G', 'B'] when channels==3,
  889. otherwise ['Channel 1', ...].
  890. cmap : str | matplotlib.colors.Colormap | Sequence | None
  891. Shared colormap applied to every channel if provided. When ``None`` and ``channel_cmaps``
  892. is also ``None``, RGB channels default to Reds/Greens/Blues so higher values render darker.
  893. channel_cmaps : Sequence[str | matplotlib.colors.Colormap] | None
  894. Optional per-channel colormap overrides. Must match ``channels`` in length.
  895. title : str | None
  896. Optional figure title.
  897. figsize : tuple
  898. Figure size forwarded to ``plt.subplots``.
  899. """
  900. import torch
  901. import numpy as np
  902. import matplotlib.pyplot as plt
  903. from matplotlib.colors import LinearSegmentedColormap
  904. # Create custom colormap: black -> white (at 0.5) -> red
  905. colors = ['black', 'white', 'red']
  906. n_bins = 256
  907. custom_cmap = LinearSegmentedColormap.from_list('custom_bwr', colors, N=n_bins)
  908. def _to_channel_maps(tensor, name):
  909. if tensor is None:
  910. raise ValueError(f"{name} must not be None")
  911. if isinstance(tensor, torch.Tensor):
  912. data = tensor.detach().cpu().numpy()
  913. else:
  914. data = np.asarray(tensor)
  915. if data.ndim == 1:
  916. expected = height * width * channels
  917. if data.size != expected:
  918. raise ValueError(
  919. f"Flat {name} length {data.size} cannot be reshaped to ({channels}, {height}, {width})"
  920. )
  921. data = data.reshape(channels, height, width)
  922. else:
  923. raise ValueError(
  924. f"{name} must be 1D flat, 2D flattened, or 3D channel-first arrays; got ndim={data.ndim}."
  925. )
  926. return data
  927. def _resolve_cmap(entry):
  928. if isinstance(entry, str):
  929. return plt.get_cmap(entry)
  930. if entry is None:
  931. raise ValueError("Colormap entries must not be None")
  932. return entry
  933. mu_maps = _to_channel_maps(mu_tensor, "mu")
  934. I_maps = _to_channel_maps(I_tensor, "I")
  935. if mu_maps.shape != I_maps.shape:
  936. raise ValueError(
  937. f"mu and I must share the same shape after reshaping; got {mu_maps.shape} vs {I_maps.shape}."
  938. )
  939. if channel_names is None:
  940. if channels == 3:
  941. channel_names = ["R", "G", "B"]
  942. else:
  943. channel_names = [f"Channel {idx + 1}" for idx in range(channels)]
  944. if len(channel_names) != channels:
  945. raise ValueError("channel_names length must match the number of channels")
  946. # Use the custom colormap for all channels
  947. resolved_cmaps = [custom_cmap for _ in range(channels)]
  948. # Create figure with RGB composite column
  949. fig, axes = plt.subplots(2, channels + 1, figsize=(figsize[0] + 4, figsize[1]))
  950. if title:
  951. fig.suptitle(title, fontsize=16)
  952. for col in range(channels):
  953. cmap_to_use = resolved_cmaps[col]
  954. mu_im = axes[0, col].imshow(mu_maps[col], cmap=cmap_to_use, vmin=0, vmax=1)
  955. axes[0, col].axis('off')
  956. mu_cbar = fig.colorbar(mu_im, ax=axes[0, col], fraction=0.046, pad=0.04)
  957. mu_cbar.set_ticks([0, 0.5, 1])
  958. I_im = axes[1, col].imshow(I_maps[col], cmap=cmap_to_use, vmin=0, vmax=1)
  959. axes[1, col].axis('off')
  960. I_cbar = fig.colorbar(I_im, ax=axes[1, col], fraction=0.046, pad=0.04)
  961. I_cbar.set_ticks([0, 0.5, 1])
  962. # Create RGB composite images
  963. if channels == 3:
  964. # Create RGB composite for μ
  965. mu_rgb = np.stack([mu_maps[0], mu_maps[1], mu_maps[2]], axis=-1)
  966. axes[0, channels].imshow(mu_rgb, vmin=0, vmax=1)
  967. axes[0, channels].axis('off')
  968. # Create RGB composite for I
  969. I_rgb = np.stack([I_maps[0], I_maps[1], I_maps[2]], axis=-1)
  970. axes[1, channels].imshow(I_rgb, vmin=0, vmax=1)
  971. axes[1, channels].axis('off')
  972. else:
  973. # For non-RGB cases, create grayscale composite
  974. mu_composite = np.mean(mu_maps, axis=0)
  975. axes[0, channels].imshow(mu_composite, cmap='gray', vmin=0, vmax=1)
  976. axes[0, channels].axis('off')
  977. I_composite = np.mean(I_maps, axis=0)
  978. axes[1, channels].imshow(I_composite, cmap='gray', vmin=0, vmax=1)
  979. axes[1, channels].axis('off')
  980. # Add column headers (R, G, B, RGB Composite) at the top
  981. for col in range(channels):
  982. axes[0, col].set_title(channel_names[col], fontsize=20, pad=10)
  983. # Add header for RGB composite column
  984. composite_title = "RGB Composite" if channels == 3 else "Composite"
  985. axes[0, channels].set_title(composite_title, fontsize=20, pad=10)
  986. # Add row labels (μ, I) on the left - temporarily turn on axis for labels
  987. for row in range(2):
  988. axes[row, 0].axis('on')
  989. axes[row, 0].set_xticks([])
  990. axes[row, 0].set_yticks([])
  991. # Remove spines to keep clean look
  992. for spine in axes[row, 0].spines.values():
  993. spine.set_visible(False)
  994. axes[0, 0].set_ylabel('μ', fontsize=20, rotation=0, labelpad=30, ha='center', va='center')
  995. axes[1, 0].set_ylabel('I', fontsize=20, rotation=0, labelpad=30, ha='center', va='center')
  996. fig.tight_layout(rect=(0, 0, 1, 0.97) if title else None)
  997. plt.subplots_adjust(left=0.1) # Add more space on the left for row labels
  998. # return fig
  999. # %%
  1000. def create_cifar10_transformation_grid(tuner_a,
  1001. tuner_b,
  1002. show_grayscale=False,
  1003. samples_per_class=1,
  1004. dataset_root='./data',
  1005. title=None,
  1006. labels=None,
  1007. figsize=(24, 8),
  1008. verbose=True):
  1009. """Create a 3-row grid comparing original images with two tuner outputs.
  1010. Parameters:
  1011. -----------
  1012. tuner_a, tuner_b : torch.nn.Module or dict
  1013. Either callable tuner modules or dict specs with 'mu'/'I' parameters
  1014. show_grayscale : bool
  1015. Whether to display images in grayscale
  1016. samples_per_class : int
  1017. Number of samples per CIFAR-10 class to display
  1018. dataset_root : str
  1019. Path to CIFAR-10 dataset
  1020. title : str, optional
  1021. Figure title
  1022. labels : list of str, optional
  1023. Labels for the two tuners. If None, uses default labels.
  1024. figsize : tuple
  1025. Figure size
  1026. verbose : bool
  1027. Whether to print min/max values during processing
  1028. """
  1029. import torch
  1030. import numpy as np
  1031. import matplotlib.pyplot as plt
  1032. from torchvision import datasets, transforms
  1033. from tuners import CM_GLLF_tuner_flexible
  1034. tuner_a.eval()
  1035. tuner_b.eval()
  1036. if samples_per_class < 1:
  1037. raise ValueError("samples_per_class must be >= 1")
  1038. device = torch.device('cpu')
  1039. # Validate labels
  1040. if labels is not None and len(labels) != 2:
  1041. raise ValueError("labels must contain exactly two entries when provided")
  1042. # CIFAR-10 normalization constants
  1043. mean = torch.tensor([0.4914, 0.4822, 0.4465], dtype=torch.float32, device=device).view(3, 1, 1)
  1044. std = torch.tensor([0.2470, 0.2435, 0.2616], dtype=torch.float32, device=device).view(3, 1, 1)
  1045. def _apply_transform(module, normalized_tensor):
  1046. """Apply tuner transformation to normalized image tensor"""
  1047. original_shape = normalized_tensor.shape
  1048. # batch_size = normalized_tensor.shape[0]
  1049. flat = normalized_tensor.view(1, -1)
  1050. with torch.no_grad():
  1051. output = module(flat)
  1052. return output.view(original_shape)
  1053. def _prepare_for_display(tensor):
  1054. """Convert normalized tensor back to displayable format"""
  1055. unnormalized = tensor * std + mean
  1056. if verbose:
  1057. min_val = float(unnormalized.min())
  1058. max_val = float(unnormalized.max())
  1059. print(f"Display tensor range: [{min_val:.4f}, {max_val:.4f}]")
  1060. return unnormalized.clamp(0.0, 1.0)
  1061. # Load CIFAR-10 dataset
  1062. dataset = datasets.CIFAR10(root=dataset_root, train=False, download=True,
  1063. transform=transforms.ToTensor())
  1064. class_names = dataset.classes
  1065. # Collect samples by class
  1066. samples_by_class = {cls: [] for cls in range(10)}
  1067. for image, label in dataset:
  1068. if len(samples_by_class[label]) < samples_per_class:
  1069. samples_by_class[label].append(image.to(device))
  1070. if all(len(v) >= samples_per_class for v in samples_by_class.values()):
  1071. break
  1072. # Validate we have enough samples
  1073. if any(len(v) < samples_per_class for v in samples_by_class.values()):
  1074. raise RuntimeError(
  1075. "Unable to gather the requested number of samples for every class; "
  1076. "try reducing samples_per_class or ensuring the dataset is available."
  1077. )
  1078. # Setup plotting
  1079. total_cols = 10 * samples_per_class
  1080. fig, axes = plt.subplots(3, total_cols, figsize=figsize)
  1081. if total_cols == 1:
  1082. axes = axes.reshape(3, 1)
  1083. def _show_image(ax, tensor, title_text=None):
  1084. """Display image tensor on given axis"""
  1085. if show_grayscale:
  1086. from torchvision.transforms.functional import rgb_to_grayscale
  1087. # Convert RGB to grayscale using standard weights
  1088. grayscale_tensor = rgb_to_grayscale(tensor, num_output_channels=1)
  1089. # Convert to NumPy for displaying with matplotlib
  1090. gray_img = grayscale_tensor.squeeze().cpu().numpy()
  1091. ax.imshow(gray_img, cmap='gray', vmin=0.0, vmax=1.0)
  1092. else:
  1093. # Display color image (convert CHW to HWC)
  1094. img_array = tensor.cpu().permute(1, 2, 0).numpy()
  1095. ax.imshow(np.clip(img_array, 0, 1)) # Ensure valid range
  1096. ax.axis('off')
  1097. if title_text:
  1098. ax.set_title(title_text, fontsize=10, pad=2)
  1099. # Generate the grid
  1100. tuners = (tuner_a, tuner_b)
  1101. col_idx = 0
  1102. for cls in range(10):
  1103. for sample_idx, original in enumerate(samples_by_class[cls], start=1):
  1104. # Normalize input image
  1105. normalized = (original - mean) / std
  1106. # Show original image
  1107. original_disp = _prepare_for_display(normalized)
  1108. title_text = class_names[cls] if samples_per_class == 1 else f"{class_names[cls]} #{sample_idx}"
  1109. _show_image(axes[0, col_idx], original_disp, title_text)
  1110. # Apply each tuner transformation
  1111. for row_offset, module in enumerate(tuners, start=1):
  1112. try:
  1113. transformed = _apply_transform(module, normalized)
  1114. transformed_disp = _prepare_for_display(transformed)
  1115. _show_image(axes[row_offset, col_idx], transformed_disp)
  1116. except Exception as e:
  1117. print(f"Error applying tuner {row_offset} to image {cls}-{sample_idx}: {e}")
  1118. # Show a black image on error
  1119. _show_image(axes[row_offset, col_idx], torch.zeros_like(original_disp))
  1120. col_idx += 1
  1121. # Set row labels
  1122. axes[0, 0].set_ylabel('Original', fontsize=12, rotation=90, labelpad=10)
  1123. # for row_offset, label in enumerate(labels):
  1124. # axes[row_offset, 0].set_ylabel(label, fontsize=12, rotation=90, labelpad=10)
  1125. # Set title
  1126. if title:
  1127. fig.suptitle(title, fontsize=16, y=0.98)
  1128. # Adjust layout
  1129. fig.tight_layout(rect=(0, 0, 1, 0.96) if title else None)
  1130. plt.subplots_adjust(hspace=0.1, wspace=0.05)
  1131. # %%
  1132. def create_detailed_comparison_figure(tuner, class_name, show_gray=False, normalize_way='clamp', dataset_root='./data', figsize=(15, 5)):
  1133. """
  1134. Create a detailed comparison figure for a single image showing:
  1135. 1. Original image
  1136. 2. CM-GLLF transformed image
  1137. 3. Global contrast comparison (bar plot) - using grayscale (luma)
  1138. 4. Pixel intensity distribution comparison - using grayscale (luma)
  1139. Parameters:
  1140. -----------
  1141. tuner : torch.nn.Module
  1142. The CM-GLLF tuner model
  1143. class_name : str
  1144. CIFAR-10 class name (default: 'dog')
  1145. dataset_root : str
  1146. Path to CIFAR-10 dataset
  1147. figsize : tuple
  1148. Figure size
  1149. """
  1150. import torch
  1151. import numpy as np
  1152. import matplotlib.pyplot as plt
  1153. from torchvision import datasets, transforms
  1154. from torchvision.transforms.functional import rgb_to_grayscale
  1155. # Load CIFAR-10 dataset
  1156. dataset = datasets.CIFAR10(root=dataset_root, train=False, download=True,
  1157. transform=transforms.ToTensor())
  1158. class_names = dataset.classes
  1159. # Find the class index
  1160. if class_name not in class_names:
  1161. raise ValueError(f"Class '{class_name}' not found. Available classes: {class_names}")
  1162. class_idx = class_names.index(class_name)
  1163. # Get first image from the specified class
  1164. original_image = None
  1165. for image, label in dataset:
  1166. if label == class_idx:
  1167. original_image = image
  1168. break
  1169. if original_image is None:
  1170. raise ValueError(f"No image found for class '{class_name}'")
  1171. # CIFAR-10 normalization constants
  1172. device = torch.device('cpu')
  1173. mean = torch.tensor([0.4914, 0.4822, 0.4465], dtype=torch.float32, device=device).view(3, 1, 1)
  1174. std = torch.tensor([0.2470, 0.2435, 0.2616], dtype=torch.float32, device=device).view(3, 1, 1)
  1175. # Normalize and apply transformation
  1176. normalized = (original_image - mean) / std
  1177. # Apply CM-GLLF transformation
  1178. tuner.eval()
  1179. original_shape = normalized.shape
  1180. flat = normalized.view(1, -1)
  1181. with torch.no_grad():
  1182. transformed_flat = tuner(flat)
  1183. transformed = transformed_flat.view(original_shape)
  1184. # Convert back to unnormalized format
  1185. def unnormalize_tensor(tensor):
  1186. """Convert normalized tensor back to original range without clamping"""
  1187. return tensor * std + mean
  1188. # Normalize to [0,1] range using min-max normalization instead of clamping
  1189. def normalize_rescale(tensor):
  1190. """Normalize tensor to [0,1] range for display.
  1191. If all values in the unnormalized tensor already lie strictly within (0,1)
  1192. (allowing a small epsilon tolerance), skip min-max normalization and just
  1193. clamp to [0,1]. Otherwise perform min-max normalization as before.
  1194. """
  1195. unnormalized = unnormalize_tensor(tensor)
  1196. # Small tolerance to treat values extremely close to 0/1 as inside (0,1)
  1197. min_val = float(unnormalized.min())
  1198. max_val = float(unnormalized.max())
  1199. # Check if all values are strictly within (0,1) up to tolerance
  1200. if (min_val > 0.0) and (max_val < 1.0):
  1201. # Already in (0,1) — no normalization needed, just clamp for safety
  1202. print(f"Unnormalized range: [{min_val:.6f}, {max_val:.6f}] -> Values already in (0,1), skipping min-max normalization.")
  1203. return unnormalized.clamp(0.0, 1.0)
  1204. # Otherwise perform min-max normalization to [0,1]
  1205. if max_val > min_val:
  1206. normalized = (unnormalized - min_val) / (max_val - min_val)
  1207. print(f"Unnormalized range: [{min_val:.6f}, {max_val:.6f}] -> Min-max normalized to [0,1].")
  1208. return normalized
  1209. else:
  1210. # Edge case: nearly constant tensor — return zeros to avoid NaNs
  1211. print(f"Unnormalized range: [{min_val:.6f}, {max_val:.6f}] -> Nearly constant, returning zeros.")
  1212. return torch.zeros_like(unnormalized)
  1213. def normalize_rescale_all(tensor):
  1214. """Normalize tensor to [0,1] range for display.
  1215. If all values in the unnormalized tensor already lie strictly within (0,1)
  1216. (allowing a small epsilon tolerance), skip min-max normalization and just
  1217. clamp to [0,1]. Otherwise perform min-max normalization as before.
  1218. """
  1219. unnormalized = unnormalize_tensor(tensor)
  1220. # Small tolerance to treat values extremely close to 0/1 as inside (0,1)
  1221. min_val = float(unnormalized.min())
  1222. max_val = float(unnormalized.max())
  1223. normalized = (unnormalized - min_val) / (max_val - min_val)
  1224. print(f"Unnormalized range: [{min_val:.6f}, {max_val:.6f}] -> Min-max normalized to [0,1].")
  1225. return normalized
  1226. def normalize_clamp(tensor):
  1227. unnormalized = unnormalize_tensor(tensor)
  1228. return unnormalized.clamp(0.0,1.0)
  1229. if normalize_way=='rescale':
  1230. original_display = normalize_rescale(normalized)
  1231. transformed_display = normalize_rescale(transformed)
  1232. elif normalize_way=='clamp':
  1233. original_display = normalize_clamp(normalized)
  1234. transformed_display = normalize_clamp(transformed)
  1235. elif normalize_way=='rescale_all':
  1236. original_display = normalize_rescale_all(normalized)
  1237. transformed_display = normalize_rescale_all(transformed)
  1238. else:
  1239. raise ValueError('Invalid normalize way!')
  1240. # Convert to grayscale for luminance analysis (luma approximation)
  1241. def rgb_to_luminance(rgb_tensor):
  1242. """Convert RGB tensor to grayscale (luma) using standard weights (BT.601-like)."""
  1243. return rgb_to_grayscale(rgb_tensor, num_output_channels=1).squeeze(0)
  1244. original_gray = rgb_to_luminance(original_display)
  1245. transformed_gray = rgb_to_luminance(transformed_display)
  1246. # Calculate global contrast (RMS contrast) on grayscale
  1247. def calculate_rms_contrast_gray(gray_tensor):
  1248. """Calculate RMS contrast (std dev of grayscale/luma) for a single image."""
  1249. gray_np = gray_tensor.cpu().numpy()
  1250. mean_intensity = np.mean(gray_np)
  1251. rms_contrast = np.sqrt(np.mean((gray_np - mean_intensity) ** 2))
  1252. return rms_contrast
  1253. original_contrast = calculate_rms_contrast_gray(original_gray)
  1254. transformed_contrast = calculate_rms_contrast_gray(transformed_gray)
  1255. # Create figure with subplots
  1256. # fig, axes = plt.subplots(1, 4, figsize=figsize)
  1257. fig, axes = plt.subplots(1, 3, figsize=figsize)
  1258. # 1. Original image (RGB)
  1259. if show_gray:
  1260. display_img = original_gray.squeeze().cpu().numpy()
  1261. axes[0].imshow(display_img,cmap='gray', vmin=0.0, vmax=1.0)
  1262. else:
  1263. display_img = original_display.cpu().permute(1, 2, 0).numpy()
  1264. axes[0].imshow(display_img)
  1265. axes[0].set_title(f'Original', fontsize=22)
  1266. axes[0].axis('off')
  1267. # 2. Transformed image (RGB)
  1268. if show_gray:
  1269. display_img = transformed_gray.squeeze().cpu().numpy()
  1270. axes[1].imshow(display_img,cmap='gray', vmin=0.0, vmax=1.0)
  1271. else:
  1272. display_img = transformed_display.cpu().permute(1, 2, 0).numpy()
  1273. axes[1].imshow(display_img)
  1274. axes[1].set_title('CMG Modulated', fontsize=22)
  1275. axes[1].axis('off')
  1276. # 3. Global contrast comparison (bar plot) - grayscale luminance
  1277. contrast_values = [original_contrast, transformed_contrast]
  1278. labels = ['Original', 'CMG']
  1279. colors = ["#060606", "#2bc62b"]
  1280. bars = axes[2].bar(labels, contrast_values, color=colors, alpha=0.7)
  1281. axes[2].set_ylabel('RMS Contrast',fontsize=16)
  1282. axes[2].set_title('RMS Contrast (Luma)',fontsize=22)
  1283. axes[2].grid(True, alpha=0.3)
  1284. axes[2].set_ylim(bottom=0.1)
  1285. axes[2].tick_params(axis='x', labelsize=16)
  1286. # Add value labels on bars
  1287. for bar, value in zip(bars, contrast_values):
  1288. height = bar.get_height()
  1289. axes[2].annotate(f'{value:.4f}',
  1290. xy=(bar.get_x() + bar.get_width() / 2, height),
  1291. xytext=(0, 3), textcoords="offset points",
  1292. ha='center', va='bottom', fontsize=16)
  1293. # 4. Pixel intensity distribution comparison - grayscale (luma)
  1294. orig_gray_flat = original_gray.cpu().numpy().flatten()
  1295. trans_gray_flat = transformed_gray.cpu().numpy().flatten()
  1296. # Create smoothed density curves using KDE
  1297. from scipy.stats import gaussian_kde
  1298. # Create x-axis for smooth curves
  1299. x_min = min(orig_gray_flat.min(), trans_gray_flat.min())
  1300. x_max = max(orig_gray_flat.max(), trans_gray_flat.max())
  1301. x = np.linspace(x_min, x_max, 200)
  1302. # Compute KDE for both distributions
  1303. kde_orig = gaussian_kde(orig_gray_flat)
  1304. kde_trans = gaussian_kde(trans_gray_flat)
  1305. # Plot smooth curves
  1306. # axes[3].plot(x, kde_orig(x), label='Original', color="#060606", linewidth=2)
  1307. # axes[3].plot(x, kde_trans(x), label='CMG', color="#2bc62b", linewidth=2)
  1308. # axes[3].set_xlabel('Luma',fontsize=16)
  1309. # axes[3].set_ylabel('Density',fontsize=16)
  1310. # axes[3].set_title('Luma Distribution',fontsize=22)
  1311. # axes[3].legend(fontsize=16)
  1312. # axes[3].grid(True, alpha=0.3)
  1313. plt.tight_layout()
  1314. plt.show()
  1315. # Print summary statistics
  1316. print(f"\n=== Analysis Summary for {class_name.capitalize()} ===")
  1317. print("Global contrast (RMS on grayscale/luma):")
  1318. print(f" Original = {original_contrast:.6f}")
  1319. print(f" CM-GLLF = {transformed_contrast:.6f}")
  1320. print(f" Change = {((transformed_contrast - original_contrast) / original_contrast * 100):+.2f}%")
  1321. print(f"\nGrayscale (luma) intensity statistics:")
  1322. print(f" Original: Mean = {orig_gray_flat.mean():.6f}, Std = {orig_gray_flat.std():.6f}")
  1323. print(f" CM-GLLF: Mean = {trans_gray_flat.mean():.6f}, Std = {trans_gray_flat.std():.6f}")
  1324. print(f" Mean Change: {((trans_gray_flat.mean() - orig_gray_flat.mean()) / orig_gray_flat.mean() * 100):+.2f}%")
  1325. print(f" Std Change: {((trans_gray_flat.std() - orig_gray_flat.std()) / orig_gray_flat.std() * 100):+.2f}%")
  1326. # %%
  1327. gllf_best_params = get_best_model(dataset='CIFAR10',
  1328. network='MLP_CIFAR10_A',
  1329. mode='CM-GLLF-all-valuewise-minmax-mapping',
  1330. scheduler='power_to_linear',
  1331. epochs=150,
  1332. seeds=[0,1,2],
  1333. init_method='uniform')
  1334. # for class_name in ['dog','frog','horse','airplane','truck','bird','cat','deer','ship']:
  1335. for class_name in ['frog','bird',]:
  1336. create_detailed_comparison_figure(gllf_best_params['tuner'], class_name=class_name,normalize_way='clamp',show_gray=False)
  1337. # summary = create_multi_class_comparison_figure(
  1338. # tuner=gllf_best_params['tuner'],
  1339. # class_names=['frog','bird'],
  1340. # normalize_way='clamp',
  1341. # title='CMG vs. Original across classes'
  1342. # )
  1343. # %%
  1344. def create_multi_class_comparison_figure(tuner, class_names, show_gray=False, normalize_way='clamp', dataset_root='./data',
  1345. title=None, row_height=5.0, figsize_width=20.0, kde_points=200, verbose=False):
  1346. """Create a stacked comparison figure for multiple classes using CM-GLLF outputs.
  1347. Each row reproduces the four-panel layout of ``create_detailed_comparison_figure``:
  1348. original image, CM-GLLF image, RMS contrast bar, and luma distribution. Column titles
  1349. are shared across rows. A summary table with per-class statistics is returned.
  1350. Parameters
  1351. ----------
  1352. tuner : torch.nn.Module
  1353. Trained CM-GLLF tuner.
  1354. class_names : Sequence[str] or str
  1355. CIFAR-10 class names to visualize. A single string is promoted to a list.
  1356. show_gray : bool, optional
  1357. Display grayscale versions of the images instead of RGB.
  1358. normalize_way : {'clamp', 'rescale', 'rescale_all'}, optional
  1359. Method used to map normalized tensors back into display space (matches
  1360. ``create_detailed_comparison_figure``).
  1361. dataset_root : str, optional
  1362. Location where CIFAR-10 is stored or downloaded.
  1363. title : str or None, optional
  1364. Optional suptitle for the entire figure.
  1365. row_height : float, optional
  1366. Height in inches allocated per class row.
  1367. figsize_width : float, optional
  1368. Width in inches of the full figure.
  1369. kde_points : int, optional
  1370. Number of evaluation points for the KDE curves.
  1371. verbose : bool, optional
  1372. When True, prints min/max ranges encountered during normalization.
  1373. Returns
  1374. -------
  1375. pandas.DataFrame
  1376. Summary statistics per class (RMS contrast and luma means).
  1377. """
  1378. import torch
  1379. import numpy as np
  1380. import pandas as pd
  1381. import matplotlib.pyplot as plt
  1382. from torchvision import datasets, transforms
  1383. from torchvision.transforms.functional import rgb_to_grayscale
  1384. from scipy.stats import gaussian_kde
  1385. if isinstance(class_names, str):
  1386. class_names = [class_names]
  1387. if not class_names:
  1388. raise ValueError("class_names must contain at least one entry.")
  1389. device = torch.device('cpu')
  1390. tuner = tuner.to(device).eval()
  1391. mean = torch.tensor([0.4914, 0.4822, 0.4465], dtype=torch.float32, device=device).view(3, 1, 1)
  1392. std = torch.tensor([0.2470, 0.2435, 0.2616], dtype=torch.float32, device=device).view(3, 1, 1)
  1393. dataset = datasets.CIFAR10(root=dataset_root, train=False, download=True, transform=transforms.ToTensor())
  1394. available_classes = set(dataset.classes)
  1395. missing = [cls for cls in class_names if cls not in available_classes]
  1396. if missing:
  1397. raise ValueError(f"Unknown classes {missing}. Available classes: {dataset.classes}")
  1398. needed = set(class_names)
  1399. class_samples = {}
  1400. for image, label in dataset:
  1401. cls_name = dataset.classes[label]
  1402. if cls_name in needed and cls_name not in class_samples:
  1403. class_samples[cls_name] = image.to(device)
  1404. if len(class_samples) == len(needed):
  1405. break
  1406. if len(class_samples) != len(needed):
  1407. unresolved = sorted(needed.difference(class_samples.keys()))
  1408. raise RuntimeError(f"Failed to collect samples for classes: {unresolved}")
  1409. def _apply_tuner(module, normalized_tensor):
  1410. flat = normalized_tensor.view(1, -1)
  1411. with torch.no_grad():
  1412. transformed_flat = module(flat)
  1413. return transformed_flat.view_as(normalized_tensor)
  1414. def _unnormalize(tensor):
  1415. return tensor * std + mean
  1416. def _normalize_display(tensor):
  1417. unnormalized = _unnormalize(tensor).detach()
  1418. min_val = float(unnormalized.min())
  1419. max_val = float(unnormalized.max())
  1420. if verbose:
  1421. print(f"[{normalize_way}] range: [{min_val:.6f}, {max_val:.6f}]")
  1422. if normalize_way == 'clamp':
  1423. return unnormalized.clamp(0.0, 1.0)
  1424. elif normalize_way == 'rescale':
  1425. if (min_val > 0.0) and (max_val < 1.0):
  1426. if verbose:
  1427. print("Values already in (0,1); skipping min-max rescale.")
  1428. return unnormalized.clamp(0.0, 1.0)
  1429. if max_val > min_val:
  1430. if verbose:
  1431. print("Applying per-image min-max rescale.")
  1432. return (unnormalized - min_val) / (max_val - min_val)
  1433. if verbose:
  1434. print("Nearly constant tensor; returning zeros.")
  1435. return torch.zeros_like(unnormalized)
  1436. elif normalize_way == 'rescale_all':
  1437. if max_val > min_val:
  1438. if verbose:
  1439. print("Applying per-image min-max rescale (rescale_all).")
  1440. return (unnormalized - min_val) / (max_val - min_val)
  1441. if verbose:
  1442. print("Nearly constant tensor; returning zeros.")
  1443. return torch.zeros_like(unnormalized)
  1444. else:
  1445. raise ValueError("Invalid normalize_way. Choose from {'clamp', 'rescale', 'rescale_all'}.")
  1446. def _rms_contrast(gray_tensor):
  1447. flat = gray_tensor.view(-1)
  1448. mean_intensity = flat.mean()
  1449. return torch.sqrt(torch.mean((flat - mean_intensity) ** 2)).item()
  1450. num_rows = len(class_names)
  1451. fig_height = max(row_height * num_rows, row_height)
  1452. fig, axes = plt.subplots(num_rows, 4, figsize=(figsize_width, fig_height))
  1453. if num_rows == 1:
  1454. axes = axes[np.newaxis, :]
  1455. summary_rows = []
  1456. for row_idx, class_name in enumerate(class_names):
  1457. image = class_samples[class_name]
  1458. normalized = (image - mean) / std
  1459. transformed = _apply_tuner(tuner, normalized)
  1460. original_display = _normalize_display(normalized)
  1461. transformed_display = _normalize_display(transformed)
  1462. original_gray = rgb_to_grayscale(original_display, num_output_channels=1).squeeze(0)
  1463. transformed_gray = rgb_to_grayscale(transformed_display, num_output_channels=1).squeeze(0)
  1464. orig_mean = float(original_gray.mean().item())
  1465. trans_mean = float(transformed_gray.mean().item())
  1466. if show_gray:
  1467. original_img = original_gray.detach().cpu().numpy()
  1468. transformed_img = transformed_gray.detach().cpu().numpy()
  1469. cmap = 'gray'
  1470. else:
  1471. original_img = original_display.detach().cpu().permute(1, 2, 0).numpy()
  1472. transformed_img = transformed_display.detach().cpu().permute(1, 2, 0).numpy()
  1473. cmap = None
  1474. axes[row_idx, 0].imshow(original_img if cmap else np.clip(original_img, 0, 1), cmap=cmap, vmin=0.0, vmax=1.0)
  1475. # axes[row_idx, 0].set_title(f'Original', fontsize=22)
  1476. axes[row_idx, 0].axis('off')
  1477. axes[row_idx, 1].imshow(transformed_img if cmap else np.clip(transformed_img, 0, 1), cmap=cmap, vmin=0.0, vmax=1.0)
  1478. # axes[row_idx, 1].set_title(f'CMG Transformed', fontsize=22)
  1479. axes[row_idx, 1].axis('off')
  1480. original_rms = _rms_contrast(original_gray)
  1481. transformed_rms = _rms_contrast(transformed_gray)
  1482. bars = axes[row_idx, 2].bar(['Original', 'CMG'], [original_rms, transformed_rms],
  1483. color=['#060606', "#2bc62b"], alpha=0.7)
  1484. axes[row_idx, 2].set_ylabel('RMS contrast (luma)', fontsize=16)
  1485. axes[row_idx, 2].grid(True, alpha=0.3)
  1486. axes[row_idx, 2].tick_params(axis='x', labelsize=16)
  1487. for bar, value in zip(bars, [original_rms, transformed_rms]):
  1488. height = bar.get_height()
  1489. axes[row_idx, 2].annotate(f'{value:.4f}',
  1490. (bar.get_x() + bar.get_width() / 2, height),
  1491. xytext=(0, 3), textcoords='offset points',
  1492. ha='center', va='bottom', fontsize=16)
  1493. orig_gray_flat = original_gray.detach().cpu().numpy().flatten()
  1494. trans_gray_flat = transformed_gray.detach().cpu().numpy().flatten()
  1495. x = np.linspace(0.0, 1.0, kde_points)
  1496. kde_orig = gaussian_kde(orig_gray_flat)
  1497. kde_trans = gaussian_kde(trans_gray_flat)
  1498. axes[row_idx, 3].plot(x, kde_orig(x), label='Original', color='#060606', linewidth=2)
  1499. axes[row_idx, 3].plot(x, kde_trans(x), label='CMG', color='green', linewidth=2)
  1500. axes[row_idx, 3].set_xlim(0.0, 1.0)
  1501. axes[row_idx, 3].set_xlabel("Luma", fontsize=16)
  1502. axes[row_idx, 3].set_ylabel('Density', fontsize=16)
  1503. axes[row_idx, 3].grid(True, alpha=0.3)
  1504. if row_idx == 0:
  1505. axes[row_idx, 3].legend(fontsize=16)
  1506. axes[row_idx, 0].text(-0.12, 0.5, class_name.capitalize(), transform=axes[row_idx, 0].transAxes,
  1507. fontsize=14, rotation=90, ha='center', va='center')
  1508. summary_rows.append({
  1509. 'class': class_name,
  1510. 'rms_original': original_rms,
  1511. 'rms_cmg': transformed_rms,
  1512. 'rms_difference': transformed_rms - original_rms,
  1513. 'mean_luma_original': orig_mean,
  1514. 'mean_luma_cmg': trans_mean,
  1515. 'mean_difference': trans_mean - orig_mean
  1516. })
  1517. axes[0, 0].set_title('Original', fontsize=22)
  1518. axes[0, 1].set_title('CMG transformed', fontsize=22)
  1519. axes[0, 2].set_title('RMS contrast (Luma)', fontsize=22)
  1520. axes[0, 3].set_title('Luma distribution', fontsize=22)
  1521. if title:
  1522. fig.suptitle(title, fontsize=18)
  1523. plt.tight_layout(rect=(0, 0, 1, 0.97) if title else None)
  1524. summary_df = pd.DataFrame(summary_rows).set_index('class')
  1525. return summary_df
  1526. # %%
  1527. import torch
  1528. import numpy as np
  1529. import pandas as pd
  1530. from torchvision import datasets, transforms
  1531. from torchvision.transforms.functional import rgb_to_grayscale
  1532. from torch.utils.data import DataLoader, Subset
  1533. from scipy import stats
  1534. from collections import defaultdict
  1535. def analyze_global_contrast_statistics(tuner, dataset_root='./data', normalize_way='clamp', max_images=None, batch_size=512, num_workers=2):
  1536. """
  1537. Compute dataset-wide RMS contrast statistics before and after CM-GLLF.
  1538. Parameters
  1539. ----------
  1540. tuner : torch.nn.Module
  1541. Trained CM-GLLF tuner to evaluate.
  1542. dataset_root : str, optional
  1543. Location to download/load CIFAR-10.
  1544. normalize_way : str, optional
  1545. Strategy used to map normalized tensors back into display space. Options:
  1546. 'clamp', 'rescale', or 'rescale_all'.
  1547. max_images : int or None, optional
  1548. If provided, limits the number of images analyzed (for quick tests).
  1549. batch_size : int, optional
  1550. Mini-batch size used for processing images through the tuner.
  1551. num_workers : int, optional
  1552. Number of subprocesses for data loading.
  1553. """
  1554. device = torch.device('cuda:0') if torch.cuda.is_available() else torch.device('cpu')
  1555. tuner = tuner.to(device).eval()
  1556. mean = torch.tensor([0.4914, 0.4822, 0.4465], dtype=torch.float32, device=device).view(3, 1, 1)
  1557. std = torch.tensor([0.2470, 0.2435, 0.2616], dtype=torch.float32, device=device).view(3, 1, 1)
  1558. dataset = datasets.CIFAR10(
  1559. root=dataset_root,
  1560. train=False,
  1561. download=True,
  1562. transform=transforms.ToTensor()
  1563. )
  1564. total_available = len(dataset)
  1565. total_images = total_available if max_images is None else min(max_images, total_available)
  1566. indices = list(range(total_images))
  1567. data_subset = Subset(dataset, indices)
  1568. loader = DataLoader(data_subset, batch_size=batch_size, shuffle=False, num_workers=num_workers)
  1569. def _unnormalize(tensor):
  1570. return tensor * std + mean
  1571. def _to_display_batch(tensor_batch):
  1572. unnormalized = _unnormalize(tensor_batch)
  1573. if normalize_way == 'clamp':
  1574. return unnormalized.clamp(0.0, 1.0)
  1575. elif normalize_way in {'rescale', 'rescale_all'}:
  1576. flat = unnormalized.view(unnormalized.size(0), -1)
  1577. mins = flat.min(dim=1)[0].view(-1, 1, 1, 1)
  1578. maxs = flat.max(dim=1)[0].view(-1, 1, 1, 1)
  1579. ranges = maxs - mins
  1580. safe_ranges = torch.where(ranges < 1e-8, torch.ones_like(ranges), ranges)
  1581. normalized = (unnormalized - mins) / safe_ranges
  1582. zero_tensor = torch.zeros_like(normalized)
  1583. return torch.where(ranges < 1e-8, zero_tensor, normalized)
  1584. else:
  1585. raise ValueError(f"Unsupported normalize_way '{normalize_way}'.")
  1586. def _rms_contrast_batch(gray_tensor):
  1587. # RMS contrast = std dev of grayscale (luma) intensities per image
  1588. flat = gray_tensor.view(gray_tensor.size(0), -1)
  1589. means = flat.mean(dim=1, keepdim=True)
  1590. return torch.sqrt(((flat - means) ** 2).mean(dim=1))
  1591. original_contrasts_list = []
  1592. transformed_contrasts_list = []
  1593. label_batches = []
  1594. class_stats = defaultdict(lambda: {'total': 0, 'improved': 0, 'decreased': 0, 'tied': 0})
  1595. for images, labels in loader:
  1596. images = images.to(device)
  1597. labels = labels.to(device)
  1598. normalized = (images - mean) / std
  1599. with torch.no_grad():
  1600. transformed = tuner(normalized.view(images.size(0), -1)).view_as(normalized)
  1601. original_display = _to_display_batch(normalized)
  1602. transformed_display = _to_display_batch(transformed)
  1603. original_gray = rgb_to_grayscale(original_display, num_output_channels=1).squeeze(1)
  1604. transformed_gray = rgb_to_grayscale(transformed_display, num_output_channels=1).squeeze(1)
  1605. orig_contrast_batch = _rms_contrast_batch(original_gray)
  1606. trans_contrast_batch = _rms_contrast_batch(transformed_gray)
  1607. orig_contrast_np = orig_contrast_batch.cpu().numpy()
  1608. trans_contrast_np = trans_contrast_batch.cpu().numpy()
  1609. labels_np = labels.cpu().numpy()
  1610. original_contrasts_list.append(orig_contrast_np)
  1611. transformed_contrasts_list.append(trans_contrast_np)
  1612. label_batches.append(labels_np)
  1613. for orig_c, trans_c, label_idx in zip(orig_contrast_np, trans_contrast_np, labels_np):
  1614. class_name = dataset.classes[label_idx]
  1615. class_stats[class_name]['total'] += 1
  1616. if trans_c > orig_c:
  1617. class_stats[class_name]['improved'] += 1
  1618. elif trans_c < orig_c:
  1619. class_stats[class_name]['decreased'] += 1
  1620. else:
  1621. class_stats[class_name]['tied'] += 1
  1622. original_contrasts = np.concatenate(original_contrasts_list) if original_contrasts_list else np.array([])
  1623. transformed_contrasts = np.concatenate(transformed_contrasts_list) if transformed_contrasts_list else np.array([])
  1624. all_labels = np.concatenate(label_batches) if label_batches else np.array([])
  1625. contrast_diff = transformed_contrasts - original_contrasts
  1626. improved_mask = contrast_diff > 0
  1627. decreased_mask = contrast_diff < 0
  1628. tied_mask = contrast_diff == 0
  1629. improved_count = int(improved_mask.sum())
  1630. decreased_count = int(decreased_mask.sum())
  1631. tied_count = int(tied_mask.sum())
  1632. mean_original = float(np.mean(original_contrasts)) if original_contrasts.size else float('nan')
  1633. mean_transformed = float(np.mean(transformed_contrasts)) if transformed_contrasts.size else float('nan')
  1634. mean_diff = float(np.mean(contrast_diff)) if contrast_diff.size else float('nan')
  1635. if contrast_diff.size > 1:
  1636. std_diff = float(np.std(contrast_diff, ddof=1))
  1637. t_stat, p_two_sided = stats.ttest_rel(transformed_contrasts, original_contrasts)
  1638. df = contrast_diff.size - 1
  1639. # One-sided p-value for alternative: CM-GLLF > Original
  1640. p_value_one_sided = float(stats.t.sf(t_stat, df))
  1641. # Cohen's dz (paired) = mean(diff) / std(diff)
  1642. effect_size = mean_diff / std_diff if std_diff > 0 else float('nan')
  1643. ci_half_width = stats.t.ppf(0.975, df) * (std_diff / np.sqrt(contrast_diff.size)) if std_diff > 0 else float('nan')
  1644. ci_low = mean_diff - ci_half_width if std_diff > 0 else float('nan')
  1645. ci_high = mean_diff + ci_half_width if std_diff > 0 else float('nan')
  1646. else:
  1647. std_diff = float('nan')
  1648. t_stat = float('nan')
  1649. p_two_sided = float('nan')
  1650. p_value_one_sided = float('nan')
  1651. effect_size = float('nan')
  1652. ci_low = float('nan')
  1653. ci_high = float('nan')
  1654. per_class_rows = []
  1655. for class_name in dataset.classes:
  1656. stats_row = class_stats[class_name]
  1657. total = stats_row['total']
  1658. per_class_rows.append({
  1659. 'class': class_name,
  1660. 'total': stats_row['total'],
  1661. 'improved': stats_row['improved'],
  1662. 'decreased': stats_row['decreased'],
  1663. 'tied': stats_row['tied'],
  1664. 'improved_pct': (stats_row['improved'] / total) if total > 0 else np.nan
  1665. })
  1666. per_class_summary = pd.DataFrame(per_class_rows).set_index('class')
  1667. summary = {
  1668. 'total_images': int(total_images),
  1669. 'improved_count': improved_count,
  1670. 'decreased_count': decreased_count,
  1671. 'tied_count': tied_count,
  1672. 'improved_rate': improved_count / total_images if total_images else float('nan'),
  1673. 'mean_original_contrast': mean_original,
  1674. 'mean_transformed_contrast': mean_transformed,
  1675. 'mean_difference': mean_diff,
  1676. 'effect_size_dz': effect_size,
  1677. 't_statistic': t_stat,
  1678. 'p_value_two_sided': p_two_sided,
  1679. 'p_value_one_sided': p_value_one_sided,
  1680. 'confidence_interval_95': (ci_low, ci_high),
  1681. 'per_class_summary': per_class_summary
  1682. }
  1683. return summary
  1684. contrast_stats = analyze_global_contrast_statistics(
  1685. tuner=gllf_best_params['tuner'],
  1686. dataset_root='./data',
  1687. normalize_way='clamp',
  1688. max_images=None,
  1689. batch_size=512,
  1690. num_workers=4
  1691. )
  1692. print(f"Images analyzed: {contrast_stats['total_images']}")
  1693. print(f"Improved contrast: {contrast_stats['improved_count']} ({contrast_stats['improved_rate']*100:.2f}%)")
  1694. print(f"Decreased contrast: {contrast_stats['decreased_count']}")
  1695. print(f"No change: {contrast_stats['tied_count']}")
  1696. print(f"Mean RMS contrast (original): {contrast_stats['mean_original_contrast']:.6f}")
  1697. print(f"Mean RMS contrast (CM-GLLF): {contrast_stats['mean_transformed_contrast']:.6f}")
  1698. print(f"Mean difference: {contrast_stats['mean_difference']:.6f}")
  1699. print(f"Paired t-test (one-sided, CM-GLLF > original): t = {contrast_stats['t_statistic']:.4f}, p = {contrast_stats['p_value_one_sided']:.3e}")
  1700. print(f"Paired t-test (two-sided): p = {contrast_stats['p_value_two_sided']:.3e}")
  1701. print(f"Effect size (Cohen's dz, paired): {contrast_stats['effect_size_dz']:.4f}")
  1702. ci_low, ci_high = contrast_stats['confidence_interval_95']
  1703. print(f"95% CI for mean difference: [{ci_low:.6f}, {ci_high:.6f}]")
  1704. display(contrast_stats['per_class_summary'].style.format({
  1705. 'improved_pct': '{:.2%}'
  1706. }))
  1707. # %% [markdown]
  1708. # # Generate Excel Files with Hyperparameter Results
  1709. # %% [markdown]
  1710. # ## Notes on qualitative metrics and tests
  1711. # - RMS contrast: computed as the standard deviation of grayscale intensities (luma). This is appropriate for image-wide contrast when mean luminance varies.
  1712. # - Luma vs. luminance: we use grayscale luma as a perceptual proxy for luminance via rgb_to_grayscale (BT.601-like weights). If you prefer, call the plot "grayscale luma intensity distribution" or simply "grayscale intensity distribution".
  1713. # - "Luminance distribution" wording is acceptable in vision contexts, but "grayscale intensity (luma) distribution" is more precise for sRGB images in code.
  1714. # - Paired t-test: compares the mean of within-image differences (CM-GLLF minus Original RMS contrast). The t-statistic is mean(diff)/(sd(diff)/sqrt(n)); the p-value quantifies evidence against no mean change. One-sided p tests CM-GLLF > Original; two-sided tests any change. Assumptions: pairs are matched per image, differences are approximately normal or n is large (CLT).
  1715. # %%
  1716. def create_hyperparameter_table(dataset, network, mode, scheduler, epochs, seeds, init_method=None,
  1717. lrs=[0.025, 0.01, 0.001], batch_sizes=[32, 64, 128],
  1718. weight_decay=5e-4, is_tuned=True):
  1719. """
  1720. Create a table of average accuracies for different hyperparameter combinations.
  1721. Returns a pandas DataFrame with batch_sizes as rows and learning rates as columns.
  1722. """
  1723. import pandas as pd
  1724. # Initialize results dictionary
  1725. results = {}
  1726. for bs in batch_sizes:
  1727. results[bs] = {}
  1728. for lr in lrs:
  1729. combo_accs = []
  1730. for seed in seeds:
  1731. try:
  1732. if is_tuned:
  1733. # Get file path for tuned model
  1734. result_path, _ = get_fname_tuner(
  1735. dataset, network, bs, lr, lr, epochs, seed,
  1736. True, mode, True, init_method, weight_decay,
  1737. False, False, scheduler_w=scheduler
  1738. )
  1739. else:
  1740. # Get file path for untuned model
  1741. result_path, _ = get_fname_tuner(
  1742. dataset, network, bs, lr, None, epochs, seed,
  1743. False, None, False, None, weight_decay,
  1744. False, False, scheduler_w=scheduler
  1745. )
  1746. # Load test results
  1747. test_file = os.path.join(result_path, 'res_test.mat')
  1748. if os.path.exists(test_file):
  1749. test_data = recover_dict(loadmat(test_file))
  1750. max_test_acc = np.max(test_data["top1"])
  1751. combo_accs.append(max_test_acc)
  1752. except Exception as e:
  1753. print(f"Error processing {dataset}-{network}-{mode}-bs{bs}-lr{lr}-seed{seed}: {e}")
  1754. continue
  1755. # Calculate average accuracy for this combination
  1756. if combo_accs:
  1757. avg_acc = np.mean(combo_accs)
  1758. std_acc = np.std(combo_accs)
  1759. results[bs][lr] = f"{avg_acc:.2f}"
  1760. else:
  1761. results[bs][lr] = "N/A"
  1762. # Convert to DataFrame
  1763. df = pd.DataFrame(results).T # Transpose so batch_sizes are rows
  1764. df.index.name = 'Batch Size\\LR'
  1765. df.columns = [f'{lr}' for lr in lrs]
  1766. return df
  1767. # %%
  1768. def create_excel_for_dataset(dataset, network, scheduler, epochs, seeds,
  1769. modes_and_init_methods, output_dir='excel_results'):
  1770. """
  1771. Create an Excel file for a dataset with separate sheets for each mode and the original MLP.
  1772. """
  1773. import pandas as pd
  1774. from openpyxl import Workbook
  1775. from openpyxl.styles import Font, PatternFill, Border, Side, Alignment
  1776. from openpyxl.utils.dataframe import dataframe_to_rows
  1777. # Create output directory if it doesn't exist
  1778. os.makedirs(output_dir, exist_ok=True)
  1779. # Create Excel file
  1780. excel_filename = os.path.join(output_dir, f'{dataset}_hyperparameter_results.xlsx')
  1781. with pd.ExcelWriter(excel_filename, engine='openpyxl') as writer:
  1782. # Add sheet for original MLP (no tuner)
  1783. print(f"Processing Original MLP for {dataset}...")
  1784. original_df = create_hyperparameter_table(
  1785. dataset, network, None, scheduler, epochs, seeds,
  1786. init_method=None, is_tuned=False
  1787. )
  1788. original_df.to_excel(writer, sheet_name='Original_MLP')
  1789. # Add sheets for each tuned mode
  1790. for mode, init_method in modes_and_init_methods:
  1791. print(f"Processing {mode} for {dataset}...")
  1792. try:
  1793. tuned_df = create_hyperparameter_table(
  1794. dataset, network, mode, scheduler, epochs, seeds,
  1795. init_method=init_method, is_tuned=True
  1796. )
  1797. # Clean sheet name (Excel has character limits and restrictions)
  1798. sheet_name = mode.replace('-', '_')[:30] # Limit to 30 chars
  1799. tuned_df.to_excel(writer, sheet_name=sheet_name)
  1800. except Exception as e:
  1801. print(f"Error processing {mode}: {e}")
  1802. continue
  1803. # Apply formatting to the Excel file
  1804. wb = writer.book if hasattr(writer, 'book') else None
  1805. if wb is None:
  1806. # Re-open the file for formatting
  1807. from openpyxl import load_workbook
  1808. wb = load_workbook(excel_filename)
  1809. # Format each sheet
  1810. for sheet_name in wb.sheetnames:
  1811. ws = wb[sheet_name]
  1812. # Header formatting
  1813. header_font = Font(bold=True, color="FFFFFF")
  1814. header_fill = PatternFill(start_color="366092", end_color="366092", fill_type="solid")
  1815. # Apply formatting to header row and column
  1816. for cell in ws[1]: # First row
  1817. if cell.value:
  1818. cell.font = header_font
  1819. cell.fill = header_fill
  1820. cell.alignment = Alignment(horizontal="center")
  1821. for cell in ws['A']: # First column
  1822. if cell.value and cell.row > 1:
  1823. cell.font = Font(bold=True)
  1824. cell.alignment = Alignment(horizontal="center")
  1825. # Auto-adjust column widths
  1826. for column in ws.columns:
  1827. max_length = 0
  1828. column_letter = column[0].column_letter
  1829. for cell in column:
  1830. try:
  1831. if len(str(cell.value)) > max_length:
  1832. max_length = len(str(cell.value))
  1833. except:
  1834. pass
  1835. adjusted_width = min(max_length + 2, 20)
  1836. ws.column_dimensions[column_letter].width = adjusted_width
  1837. wb.save(excel_filename)
  1838. print(f"Excel file saved: {excel_filename}")
  1839. return excel_filename
  1840. # %%
  1841. # Generate Excel files for all datasets
  1842. import pandas as pd
  1843. # Define the modes and their corresponding initialization methods
  1844. modes_and_init_methods = [
  1845. ('CM-GLLF-all-valuewise-minmax-mapping', 'uniform'),
  1846. ('linear-valuewise-positive', 'uniform'),
  1847. ('cubic-valuewise-positive', 'uniform'),
  1848. ('adaptive-sigmoid-valuewise', 'uniform'),
  1849. ('adaptive-tanh-valuewise', 'uniform'),
  1850. ('srelu_valuewise_positive', 'uniform'),
  1851. ]
  1852. # Dataset configurations
  1853. datasets_config = [
  1854. ('CIFAR10', 'MLP_CIFAR10_A'),
  1855. ('CIFAR100', 'MLP_CIFAR100_A')
  1856. ]
  1857. # Parameters
  1858. scheduler = 'power_to_linear'
  1859. epochs = 150
  1860. seeds = [0, 1, 2]
  1861. # Generate Excel files for each dataset
  1862. for dataset, network in datasets_config:
  1863. print(f"\n{'='*50}")
  1864. print(f"Generating Excel file for {dataset}")
  1865. print(f"{'='*50}")
  1866. excel_file = create_excel_for_dataset(
  1867. dataset=dataset,
  1868. network=network,
  1869. scheduler=scheduler,
  1870. epochs=epochs,
  1871. seeds=seeds,
  1872. modes_and_init_methods=modes_and_init_methods
  1873. )
  1874. print(f"Completed: {excel_file}")
  1875. print(f"\n{'='*50}")
  1876. print("All Excel files generated successfully!")
  1877. print(f"{'='*50}")
  1878. # %%

plot_results.ipynb at commit b060764, no license · at the source

Overview

  1. Center for Complex Network Intelligence (CCNI), Tsinghua Laboratory of Brain and Intelligence (THBI), Department of Psychological and Cognitive Sciences, Tsinghua University, Beijing, China
  2. School of Biomedical Engineering, Tsinghua University, Beijing, China
  3. Department of Computer Science and Technology, Tsinghua University, Beijing, China
Institutions: Tsinghua University (China)
Journal: Frontiers in artificial intelligence, volume 9, article 1785867
Dates: received 12 January 2026; accepted 21 May 2026; published online 9 June 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.3389/frai.2026.1785867 · PMID 42344008 · PMCID PMC13287133 · OpenAlex W4415783342
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: methods / tools (subfield)
Methods: Statistics, Machine learning
Keywords: generalized logistic-logit function, input feature modulator, multi-layer perceptron, neural network, neuron segmentation
Topic: Advanced Neural Network Applications (Computer Vision and Pattern Recognition, Computer Science), according to OpenAlex
Citations: not cited yet (Europe PMC); 30 references in the paper

Abstract

Logistic and logit functions play important roles in modern science, serving as foundational tools in various applications, such as artificial neural networks (ANN). While there are functions that could produce distinct logistic and logit curves, no single, unified framework has been developed to generate both logistic and logit curves. We introduce a Cannistraci–Muscoloni–Gu generalized logistic–logit function (CMG-GLLF) to fill this gap. CMG-GLLF provides four interpretable and trainable parameters that allow explicit control over: curve type and steepness, asymmetry, and upper and lower limits of x- and y-axes. CMG-GLLF’s potential is explored in basic machine intelligence tasks. As a proof-of-concept on how this function can improve the performance of deep learning, we propose a trainable input feature modulator (IFM) that consists of learning the parameters of the CMG-GLLF for each input layer node during backpropagation for a multi-layer perceptron (MLP), which is a fundamental building block of many complex network architectures. Compared to various other learnable functions, across three different optimizers, CMG-GLLF allows superior MLP’s accuracy and stable training behavior on CIFAR-10 and CIFAR-100 image classification, but at the cost of increased computational time. Hence, we identified limitations to address in future studies, notably the need to derive an explicit mathematical expression for the logit phase, which could: (i) mitigate numerical instability in more complex architectures (e.g., CNNs) while reducing computational overhead and (ii) enable a systematic evaluation of CMG as an activation function across all layers. Furthermore, CMG-GLLF, adopted as a data transformation function, enhances the accuracy of affinity-graph-based neuron segmentation. CMG-GLLF combines in a unique framework the ability of logistic and logit functions to modulate signals or variables, covering a full spectrum of attenuation or amplification transformations. CMG-GLLF is flexible and trainable, has the potential to advance machine learning models, and can inspire further applications in other data analysis challenges in different domains of science.

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

biomedical-cybernetics/CMG-GLLF

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: b060764541ff16aad6da0883d07d8943d6b417bd, 27 April 2026
Languages: Python (226), Shell (7), MATLAB (4), Jupyter (2)
Size: 294 files, 239 scripts
Software Heritage: not archived
Found in: the text, “Datasets”
Holds: README, environment (optimizers/environment_train_ifm.yml), 2 notebooks
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: PyTorch (113 files), NumPy (67 files), Matplotlib (30 files), SciPy (12 files), Pillow (5 files), pandas (4 files), scikit-learn (2 files), h5py (1 file), Parallel Computing Toolbox (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
240 files

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;
  • 239 scripts, each with its path and the digest of its content;
  • 17 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

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

Data

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

Data availability statement

Publicly available datasets were analyzed in this study. This data can be found at: CIFAR-10 & CIFAR-100: https://www.cs.toronto.edu/~kriz/cifar.html; CREMI: https://cremi.org/data/.

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 2, 28 September 2026

  • Authors: added Wenqi Gu (0009-0007-2283-2261); Yingtao Zhang (0009-0008-5803-3020); Alessandro Muscoloni (0000-0002-9238-3357); Carlo Vittorio Cannistraci (0000-0003-0100-8410); removed Wenqi Gu; Yingtao Zhang; Alessandro Muscoloni; Carlo Vittorio Cannistraci
  • Funding: added Tsinghua University

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, pages, dates, 4 authors, 5 keywords, 17 references.

Cite

This paper

Gu, W., Zhang, Y., Muscoloni, A., & Cannistraci, C. V. (2026). A generalized logistic-logit function and its application to multi-layer perceptron and neuron segmentation. Frontiers in artificial intelligence, 9, 1785867. https://doi.org/10.3389/frai.2026.1785867

BibTeX

@article{gu2026generalized,
author = {Gu, Wenqi and Zhang, Yingtao and Muscoloni, Alessandro and Cannistraci, Carlo Vittorio},
title = {{A generalized logistic-logit function and its application to multi-layer perceptron and neuron segmentation}},
journal = {Frontiers in artificial intelligence},
year = {2026},
month = jun,
volume = {9},
pages = {1785867},
publisher = {Frontiers Media SA},
issn = {2624-8212},
doi = {10.3389/frai.2026.1785867},
url = {https://doi.org/10.3389/frai.2026.1785867},
pmid = {42344008},
pmcid = {PMC13287133}
}

RIS

TY - JOUR
AU - Gu, Wenqi
AU - Zhang, Yingtao
AU - Muscoloni, Alessandro
AU - Cannistraci, Carlo Vittorio
TI - A generalized logistic-logit function and its application to multi-layer perceptron and neuron segmentation
T2 - Frontiers in artificial intelligence
J2 - Front Artif Intell
PY - 2026
DA - 2026/06/09
VL - 9
SP - 1785867
SN - 2624-8212
PB - Frontiers Media SA
DO - 10.3389/frai.2026.1785867
UR - https://doi.org/10.3389/frai.2026.1785867
LA - en
ER -

CSL-JSON

{
"id": "10.3389/frai.2026.1785867",
"type": "article-journal",
"title": "A generalized logistic-logit function and its application to multi-layer perceptron and neuron segmentation",
"container-title": "Frontiers in artificial intelligence",
"author": [
{
"family": "Gu",
"given": "Wenqi"
},
{
"family": "Zhang",
"given": "Yingtao"
},
{
"family": "Muscoloni",
"given": "Alessandro"
},
{
"family": "Cannistraci",
"given": "Carlo Vittorio"
}
],
"container-title-short": "Front Artif Intell",
"volume": "9",
"page": "1785867",
"DOI": "10.3389/frai.2026.1785867",
"PMID": "42344008",
"PMCID": "PMC13287133",
"ISSN": "2624-8212",
"publisher": "Frontiers Media SA",
"URL": "https://doi.org/10.3389/frai.2026.1785867",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
9
]
]
}
}

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.3389/frai.2026.1771088 [code]
Few-shot deployment of pretrained MRI transformers in brain imaging tasks.
Journal: Frontiers in artificial intelligence
In common: h5py, Pillow, PyTorch, 5 other tools, methods / tools, 1 reference
[2] doi:10.1016/j.patter.2026.101590 [code]
Density-based longitudinal neuron tracking in high-density electrophysiological recordings.
Journal: Patterns (New York, N.Y.)
In common: Parallel Computing Toolbox, h5py, Pillow, 6 other tools
[3] doi:10.1371/journal.pone.0344600 [code]
Robust disease prognosis via diagnostic knowledge preservation: A sequential learning approach.
Journal: PloS one
In common: h5py, Pillow, PyTorch, 5 other tools, 1 reference
[4] doi:10.1038/s41467-026-71270-w [code]
Spatiotemporal dynamics of the human cortical functional hierarchy across the lifespan.
Journal: Nature communications
In common: Parallel Computing Toolbox, h5py, Pillow, 5 other tools
[5] doi:10.7554/elife.110588 [code]
Opening the black box toward a modular approach to spike sorting.
Journal: eLife
In common: h5py, Pillow, PyTorch, 5 other tools, methods / tools
[6] 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: h5py, Pillow, PyTorch, 5 other tools, methods / tools
[7] doi:10.1093/bioinformatics/btag540 [code]
Deciphering spatial heterogeneity by multimodal spatial transcriptomics modelling with SpatialModal.
Journal: Bioinformatics (Oxford, England)
In common: h5py, Pillow, PyTorch, 5 other tools, methods / tools
[8] doi:10.1007/s12021-026-09803-3 [code]
NeuroFusion: A Unified Framework for Generalized Visual Stimulus Decoding from fMRI Across Datasets and Subjects.
Journal: Neuroinformatics
In common: h5py, Pillow, PyTorch, 5 other tools, methods / tools
[9] doi:10.1162/imag.a.1299 [code]
A modular semantic-structural pipeline for visual decoding from primate spiking data via selective temporal integration.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: h5py, Pillow, PyTorch, 5 other tools, methods / tools
[10] doi:10.1038/s41598-026-57519-w [code]
Automated segmentation of neurons and spinal cord structures in immunofluorescence images using SpineDL.
Journal: Scientific reports
In common: h5py, Pillow, PyTorch, 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.