Sharpness-Aware Geometric Defense (SaGD)
- The paper introduces SaGD, a defense framework that smooths the adversarial loss landscape using Riemannian Sharpness-aware Minimization to enhance OOD detection.
- SaGD employs dual geometric projections—a hypersphere head with vMF-based classification and a hyperbolic head—for learning separable latent embeddings under attack.
- SaGD integrates Jitter-based adversarial training to keep perturbed ID samples near the manifold, yielding significant improvements in FPR95 and AUC on CIFAR benchmarks.
Sharpness-aware Geometric Defense (SaGD) is a training-time defense and geometry-aware scoring framework for robust out-of-distribution (OOD) detection in the setting where both in-distribution (ID) and OOD samples are subjected to adversarial perturbations. It was introduced to address a failure mode of contemporary geometry-projection OOD detectors: adversarially perturbed ID samples can be pushed far from their class prototype or neighborhood in latent space and are therefore incorrectly flagged as OOD. SaGD smooths the rugged loss landscape induced by adversarial training and learns multi-geometry latent embeddings whose projected structure remains separable for ID, including adversarial ID, versus OOD. Its central components are multi-geometry projection, Riemannian Sharpness-aware Minimization (RSAM), Jitter-based adversarial training, and a nearest-neighbor geometric OOD score in the learned embedding space (Li et al., 24 Aug 2025).
1. Problem setting and motivation
OOD detection is intended to ensure safe and reliable model deployment by distinguishing samples drawn from the training distribution from inputs lying outside it. In geometry-projection OOD detection, a sample is mapped into a latent geometry and scored by distance or density relative to the ID manifold. Under adversarial perturbations, however, an ID image can be pushed far from its class prototype or neighborhood in latent space, making its OOD score large. As a result, many geometry-based detectors misclassify adversarial ID as OOD, conflating “maliciously perturbed ID” with “unknown-class OOD” and thereby increasing false positives on adversarial ID (Li et al., 24 Aug 2025).
SaGD is designed for this underexplored but practically critical regime. The stated goal is to distinguish adversarial ID samples from OOD ones rather than treating both as a single rejected category. The motivation rests on the observation that adversarial training, while a strong defense for classification, often sharpens the loss landscape: the inner maximization introduces large gradient norms and increases curvature around local minima, which degrades convergence and stability. In a geometry-projection setting, this appears as unstable latent embeddings, poor prototype compactness, and reduced separability between ID and OOD embeddings, especially under attack (Li et al., 24 Aug 2025).
The framework therefore targets two coupled objectives. First, it seeks to smooth the adversarial loss landscape directly in geometry-projection learning. Second, it seeks to preserve ID compactness and inter-class disparity while maintaining large geometric margins to OOD so that adversarial ID remains near the ID manifold and OOD remains far from it. This suggests that SaGD should be understood not merely as an adversarially trained classifier with an auxiliary detector, but as a joint geometric representation-learning and robust OOD-detection procedure.
2. Multi-geometry latent representation
SaGD uses a backbone encoder and projection heads to learn two latent geometries simultaneously. Let denote the backbone encoder and heads collectively. Given an input , the method extracts a penultimate embedding , which is -normalized for scoring, and simultaneously learns two geometry projections (Li et al., 24 Aug 2025).
The first projection is a hypersphere head that produces on via a vMF-based classifier with class prototypes . The class-conditional vMF score for unit is
with classification probabilities
0
Its compactness loss encourages 1 to align with the correct class prototype,
2
and its disparity loss pushes prototypes apart,
3
The hypersphere loss is
4
The second projection is a hyperbolic head that produces 5 on a Poincaré ball 6 with curvature 7, corresponding to constant negative curvature 8. With Möbius operations 9 and geodesic distance
0
the supervised hyperbolic contrastive loss on an augmented set 1 is
2
The two geometry heads are combined with cross-entropy in a multi-geometry projection (MGP) objective,
3
The stated purpose of this MGP design is to learn richer latent structure than a single geometry. In the formulation of SaGD, that richer structure is intended to improve ID characterization and preserve separability under attack (Li et al., 24 Aug 2025).
3. Sharpness-aware optimization and adversarial example generation
A defining feature of SaGD is the use of Riemannian Sharpness-aware Minimization to regularize adversarial training on the manifold defined by the geometry heads. The framework defines manifold sharpness as
4
where 5 lies on a manifold 6 with retraction 7 and metric tensor 8. Using a first-order approximation,
9
the inner maximizer is the normalized Riemannian gradient,
0
The outer update performs Riemannian gradient descent on the perturbed loss,
1
In the method’s interpretation, RSAM flattens sharp minima across the multi-geometry heads, stabilizes latent embeddings, and improves ID/OOD separation under attack. The paper further frames this in gradient-Hessian terms: with adversarial samples 2, the gradients 3 and 4 differ by 5, and the curvature term 6 appears in the change of feature gradients. Large 7 and large 8 increase sharpness and harm convergence during adversarial training (Li et al., 24 Aug 2025).
SaGD instantiates adversarial training with Jitter attack rather than standard PGD or FGSM. The standard min-max objective is written as
9
Adversarial examples are generated as 0 with 1, 2, step size 3, and 4 steps. Let 5 be logits, and define
6
Jitter aims to maximize the distance to the one-hot ground-truth 7 with Gaussian jitter noise 8 injected into the softmax output:
9
Its adaptive perturbation rule is
0
where 1 reduces perturbation magnitude once the attack is effective. The stated reason is to keep adversarial examples close to the ID manifold and prevent distribution shift that could harm embedding learning. Empirically, Jitter-based adversarial training is reported to generalize better to unseen attacks than training based on standard PGD or FGSM (Li et al., 24 Aug 2025).
4. OOD scoring, training algorithm, and implementation choices
After training, SaGD constructs an ID embedding bank from penultimate-layer features and scores a test sample by nearest-neighbor geometric distance. With normalized embedding 2 from the penultimate layer, the OOD score is
3
where 4 is the 5-th nearest neighbor among ID training embeddings. A threshold 6 determines OOD versus ID. In the main experiments, 7 is taken as the nearest neighbor. The method reports that this score is empirically robust, computationally simple, and more stable than alternatives adopted in some baselines, such as Mahalanobis variants (Li et al., 24 Aug 2025).
The training procedure comprises adversarial example generation, joint geometry learning, and RSAM optimization. For each minibatch, adversarial examples are first generated via Jitter with Gaussian softmax jitter and adaptive factor 8. Clean and adversarial inputs are then forwarded through the backbone and geometry heads to produce 9, 0, and penultimate 1. The multi-geometry loss is computed, followed by an RSAM inner step that computes the Riemannian gradient, perturbs parameters within a tangent-ball of radius 2, retracts to 3, and evaluates the worst-case loss. An RSAM outer step then updates the parameters via retraction-based gradient descent on that worst-case loss. At inference time, the ID embedding bank is built from L2-normalized penultimate features, and threshold calibration is performed on a held-out ID set, optionally with a small validation OOD set, to report FPR@95%TPR, AUROC, and AUPR (Li et al., 24 Aug 2025).
The implementation details reported for the main experiments are specific. CIFAR-10 uses ResNet-18, and CIFAR-100 uses ResNet-34. Optimization uses SGD with momentum 4, weight decay 5, initial learning rate 6, and RSAM regularization. Training lasts 7 epochs with batch size 8, penultimate dimension 9, and hyperbolic curvature 0. Feature clipping of the form 1 is applied to stabilize training near the ball boundary. Attacks considered in training and evaluation include PGD, FGSM, FAB, Jitter, and CW via TorchAttacks, all with 2, step size 3, and 4 steps; additional evaluation includes APGD-100, APGD-1000, and AutoAttack (Li et al., 24 Aug 2025).
The paper also provides practical tuning guidance. It recommends Jitter-based adversarial training for robustness to unseen attacks; small RSAM radii, exemplified by 5 relative to gradient norms; curvature 6 for stability; temperatures 7 for vMF and contrastive losses; and small 8 in k-NN, with 9 in the main scoring. RSAM is reported to roughly double the per-iteration cost compared to vanilla SGD, since it requires one extra forward/backward pass at 0 plus retraction, while Jitter adversarial generation adds 1 inner steps per batch (Li et al., 24 Aug 2025).
5. Empirical evaluation
The experimental setup uses CIFAR-10 and CIFAR-100 as ID datasets and six OOD datasets: Tiny-ImageNet, Places365, LSUN, LSUN-Resize, iSUN, and Textures. For ATD and ATOM, which are exposure-based baselines, Food-101 is used for auxiliary outlier training and SVHN for validation, following their protocols. The primary metrics are FPR@95%TPR (FPR95) and AUROC, with additional reporting of Inlier AUC and Outlier AUC (Li et al., 24 Aug 2025).
The main reported quantitative results are averages over the six OOD datasets. For CIFAR-10 ID, without adversarial training, MGP-Mahalanobis achieves FPR95 2 and AUC 3, and KNN+, ASH, ODIN, and GODIN are reported to degrade significantly under attacks. With adversarial training, ATOM obtains FPR95 4 and AUC 5, ATD obtains FPR95 6 and AUC 7, and SaGD obtains FPR95 8 and AUC 9. The gains over ATD are reported as 0 FPR95 and 1 AUC. Per-condition results for SaGD on CIFAR-10 are: Clean 2, PGD 3, Jitter 4, FAB 5, FGSM 6, and CW 7 (Li et al., 24 Aug 2025).
For CIFAR-100 ID, without adversarial training, CIDER-Maha reports FPR95 8 and AUC 9, while MGP-KNN reports FPR95 00 and AUC 01. With adversarial training, ATOM reports FPR95 02 and AUC 03, ATD reports FPR95 04 and AUC 05, and SaGD reports FPR95 06 and AUC 07. The gains over ATD are 08 FPR95 and 09 AUC, and the gains over ATOM are 10 FPR95 and 11 AUC (Li et al., 24 Aug 2025).
On CIFAR-10 under adaptive attacks, the average adversarial results over five attacks are ATD: FPR95 12, AUC 13, AUCIn 14, AUCOut 15; and SaGD: FPR95 16, AUC 17, AUCIn 18, AUCOut 19. Under APGD-100, ATD reports 20 versus SaGD 21; under APGD-1000, ATD reports 22 versus SaGD 23; and under AutoAttack, ATD reports 24 versus SaGD 25 (Li et al., 24 Aug 2025).
| Setting | Comparator | SaGD |
|---|---|---|
| CIFAR-10 average over six OOD datasets | ATD: 42.59 / 87.36 | 27.68 / 94.83 |
| CIFAR-100 average over six OOD datasets | ATD: 67.58 / 77.41 | 49.87 / 87.59 |
| CIFAR-10 AutoAttack | ATD: 47.95 / 83.86 | 32.18 / 93.01 |
The table entries are FPR95 / AUC. The reported pattern is that SaGD improves both false-positive behavior and ranking quality across clean, standard adversarial, and adaptive attack settings. A plausible implication is that the method improves both separability and score calibration in robust OOD detection, although the paper itself states this more directly in terms of reduced FPR and increased AUC (Li et al., 24 Aug 2025).
6. Ablations, interpretation, and limitations
The ablation study on CIFAR-10 isolates the contributions of RSAM and Jitter. MGP + Jitter without RSAM achieves average FPR95 26 and AUC 27. MGP + RSAM without Jitter achieves average FPR95 28 and AUC 29, which the paper interprets as evidence that sharpness reduction alone without adversarial training is insufficient. CIDER + RSAM + Jitter achieves average FPR95 30 and AUC 31, indicating that both RSAM and Jitter help even with single-geometry heads. SaGD, combining MGP + RSAM + Jitter, achieves the best result: average FPR95 32 and AUC 33 (Li et al., 24 Aug 2025).
A separate ablation shows that the choice of adversarial training attack matters. RSAM+PGD training yields average FPR95 34 and AUC 35 but is very strong on FGSM, with FPR95 36 and AUC 37, indicating overfitting to attack type. RSAM+FAB yields average FPR95 38 and AUC 39; RSAM+FGSM yields average FPR95 40 and AUC 41; and RSAM+CW yields average FPR95 42 and AUC 43. The reported conclusion is that Jitter-based adversarial training generalizes better across attack families than these alternatives (Li et al., 24 Aug 2025).
The paper also studies attack intensity sensitivity under PGD on CIFAR-10. Increasing 44 from 45 to 46 is reported to double FPR95 and reduce AUC by approximately 47. SaGD is described as remaining relatively stable in AUROC while FPR95 is more sensitive, which the authors use to reinforce the importance of reporting FPR@95%TPR in robust OOD detection (Li et al., 24 Aug 2025).
Visualization results further support the method’s interpretation. Histograms of 48 for ID versus OOD under clean and adversarial settings, specifically FGSM and FAB, show that SaGD maintains clear separation, while MGP without defense and CIDER-RSAM-Jitter collapse under FGSM and FAB. RSAM+PGD yields distinct peaks only against PGD but fails on FAB, again illustrating overfitting of standard adversarial training to the seen attack (Li et al., 24 Aug 2025).
The limitations and deployment guidance are stated with similar specificity. FGSM remains a relatively challenging case because of its single-step, large-direction perturbations, although SaGD still improves AUROC and maintains separability. Stronger attacks with larger 49 increase FPR, so both FPR95 and AUROC should be reported. The framework is recommended for systems in which adversarial ID must be kept as ID while OOD rejection is also required, such as safety-critical gating or triage. It is also presented as advantageous when outlier exposure is impractical, because it avoids dependence on large auxiliary OOD datasets while outperforming exposure-based baselines ATD and ATOM in the reported experiments (Li et al., 24 Aug 2025).