📄 Paper: arXiv
- End-to-end training, adversarial data generation, error prioritization, unified evaluation, and Neural Collapse measurement/validation across image and text tasks.
- Supported datasets: MNIST, Fashion-MNIST, CIFAR-10/100, Tiny-ImageNet-200, IMDb, AGNews.
- Prioritization toolbox: Uncertainty (Entropy, DeepGini, Softmax margin/PCS), Surprise sufficiency (LSA/DSA), and NCIP (TVP+Margin).
- Metrics: RAUC + error-type diversity; optional small-scale retraining to gauge practical gains.
Fun fact: NCIP ranks “which test samples to look at first” by tracking how predictions wiggle across checkpoints and how close decisions sit to the boundary. 🧭
- Training with periodic checkpoints: main.py
- Neural Collapse metrics and visualization:
- Error prioritization:
- Uncertainty: uncertainty.py, get_rank_idx.py
- LSA/DSA: sa.py
- NCIP (TVP+Margin): ncip_priority.py
- Unified evaluation/export:
- Adversarial data:
- Generation: adv_dataset.py
- Merge: combined_adv.py
Workflow peek 🔎
train ➜ checkpoints ➜ evaluate logits ➜ rank errors (uncertainty | LSA/DSA | NCIP)
↳ CSV reports (RAUC/diversity/retraining gains)
↳ NC metrics & plots over epochs
- model/: LeNet, ResNet, VGG, DenseNet, text classifiers
- e.g., ResNet.py, VGG.py, TextClassifier.py
- src/: training, evaluation, utilities
- datasets: get_dataset.py
- models: get_model.py
- evaluation: evaluation_utils.py
- utilities: utils.py
- uncertainty ranking: uncertainty.py
- LSA/DSA: sa.py
- training: main.py
- neuronal_collapse/: metrics and validation
- core library: neuronal_collapse_lib.py
- visualization: validate_nc.py
- NCIP prioritization: ncip_priority.py
- adv_example/: adversarial dataset pipeline
- generation: adv_dataset.py
- merge: combined_adv.py
- results/: CSV reports, execution times, NCIP outputs
- data/: dataset directories (see below)
- Python ≥ 3.9
- PyTorch ≥ 1.12, torchvision
- Packages: tqdm, numpy, matplotlib, scipy, pillow, pickle5, torchattacks, dnn_tip
Installation 🧰
pip install -r requirements.txt- Auto-downloaded into
../data:- MNIST/Fashion-MNIST/CIFAR-10/100
- Tiny-ImageNet-200:
../data/tiny-imagenet-200/{train,val} - IMDb:
../data/aclImdb/{train,test}/(pos|neg)/*.txt - AGNews:
../data/ag_news_csv/{train.csv,test.csv}
See loaders: get_dataset.py
Periodic checkpoints and cosine LR scheduling ⚙️
python -m src.main \
--model_save_path ./model/CIFAR10/ResNet18 \
--model_name ResNet18 \
--dataset cifar10 \
--epochs 30 \
--batch_size 128 \
--learning_rate 0.1Code: main.py
- Supported
--datasetvalues in training:mnist,fmnist,cifar10,cifar100,tiny-imagenet,imdb,agnews. - Checkpoints are saved every epoch as
epoch_{epoch:03d}.pthunder--model_save_path/{timestamp}/. - Reproducibility note: training script fixes seed to
42internally (no--seedCLI flag).
- Unified interface 📊
- Accuracy, error indices, logits: evaluate_model
- RAUC: compute_rauc_metrics, utils.calculate_rauc
- Diversity (error pair coverage): compute_diverse_scores
- Shared CLI for
uncertainty.py/sa.py/ncip_priority.py:--model_path,--model_name,--dataset,--attack {NONE,ADV},--batch_size,--num_ratio
- CSV output per method:
- save_results_to_csv
- Path:
results/csv/{dataset}/{model}/{attack}/
- When
--attack ADV, test data is loaded from./adv_example/adv_dataset/{dataset}-{model_name}/adv_samples.pt.
- Uncertainty metrics ☁️: uncertainty.py
- Random / DeepGini / Entropy / PCS / Vanilla Softmax: get_rank_idx.py
python -m src.uncertainty \
--model_path ./model/CIFAR10/ResNet18 \
--model_name ResNet18 \
--dataset CIFAR10 \
--attack NONE \
--batch_size 128- Surprise sufficiency (LSA/DSA) 🎯: sa.py
python -m src.sa \
--model_path ./model/CIFAR10/ResNet18 \
--model_name ResNet18 \
--dataset CIFAR10 \
--attack NONE \
--batch_size 128- NCIP (TVP+Margin) 🧭: ncip_priority.py
- Build an ensemble from representative checkpoints selected by classifier-weight NC geometry
- Score each sample by
zscore(TVP across ensemble heads) + zscore(1 - final-head margin) - Ranking key:
nc_tvp - Output path:
results/nc_tvp/{dataset}/{model_name}/{attack}/
python -m neuronal_collapse.ncip_priority \
--model_path ./model/CIFAR10/ResNet18 \
--model_name ResNet18 \
--dataset CIFAR10 \
--attack NONE \
--batch_size 128 \
--num_ratio 0.7Note: --num_ratio is currently reserved in CLI; checkpoint sampling in ncip_priority.py uses internal selection logic.
References:
- Weights/features and extraction: neuronal_collapse_lib.py, penultimate features
- Equiangularity/equinormity: neuronal_collapse_lib.py, neuronal_collapse_lib.py:L207-L213
- TV distance + margin: ncip_priority.py
- Compute curves over checkpoints (Top-1, equiangularity/equinormity, classifier–mean convergence, within-class variation, NN mismatch), save plots 🧪:
validate_nc.pyscans--checkpoints_folderand evaluates every 5th checkpoint (sorted(checkpoints)[::5]).
python -m neuronal_collapse.validate_nc \
--dataset CIFAR10 \
--model_name ResNet18 \
--checkpoints_folder ./model/CIFAR10/ResNet18/{timestamp}Code: validate_nc.py
- Generate adversarial samples for the test set ⚔️:
python -m adv_example.adv_dataset \
--model_path ./model/CIFAR10/ResNet18 \
--model_name ResNet18 \
--dataset cifar10 \
--attack PGD \
--batch_size 128-
adv_dataset.pysupports datasets:mnist,fmnist,cifar10,cifar100. -
Saved files:
./adv_example/{attack}/{dataset}-{model_name}/adv_samples.pt(and an accuracy-stamped copy). -
Merge multiple attacks 🧩:
python -m adv_example.combined_advcombined_adv.pyhas no CLI args; it merges fixed attacks (FGSM,PGD,BIM,CW) into:./adv_example/adv_dataset/{dataset}-{model_name}/adv_samples.pt
- This merged file is what evaluation scripts consume when
--attack ADV.
Attacks: get_attack.py
- Add prioritized Top-k test samples to training set; small retraining; compare accuracy on random subset 🛠️:
- API: retrain_with_priority_samples
- Batch runner: evaluate_retrain_methods
- Scripts save retraining gains alongside RAUC/diversity in CSV reports.
- Fixed seed (42) via set_random_seed
- Report PyTorch/torchvision versions, GPU model/driver
- Provide dataset sources/paths as above
- Share the exact command lines used to produce tables/plots