Skip to content

Can Circuit Alignment Predict OOD Generalization?

Conference: NeurIPS2026
arXiv: 2609.31996
Code: https://github.com/ayanban011/ACE
Area: Interpretability
Keywords: circuit alignment, out-of-distribution generalization, graph kernels, class-conditional structure, model ranking

TL;DR

The paper defines Circuit Alignment Score (CAS) as same-class cross-domain circuit similarity minus cross-class circuit similarity, achieving a mean Spearman correlation of 0.88 on PACS for source-domain model selection without target data, but its OOD ranking consistency requires an additional monotonicity assumption and is not a fully data-free weight diagnostic.

Background & Motivation

Selecting a classifier that will survive a new domain before deployment is often harder than reporting source validation accuracy: models that recognize objects in photographs may fail on sketches or cartoons. ATC, ProjNorm, and ALine-D can estimate target accuracy using unlabeled target samples, but they are not directly available when the target distribution has not yet been observed. CKA, SVCCA, and RSA offer another seemingly natural clue: if representation geometry is similar across source domains, might the computation also remain stable across domains?

The authors argue that representation similarity is insufficient. Two models can produce similar final-layer features while relying on different neurons and connection pathways; aggregate feature similarity may miss the changes when domain shift affects those pathways. Checking only whether each class preserves its pathway is also insufficient: if dog and horse circuits become increasingly similar, same-class similarity can remain high while class discrimination deteriorates. Same-class preservation and cross-class entanglement therefore need separate measurements rather than a single average over every circuit pair.

The paper uses class-specific circuits from mechanistic interpretability for model selection rather than only explaining individual predictions. Its actual inputs include trained models and source samples with class information: circuits are extracted for each source domain, then assessed for same-class preservation and resistance to cross-class confusion. Core Idea: compare class-specific circuits as graphs, use the gap between same-class cross-domain preservation and cross-class entanglement as an OOD ranking proxy, and separately analyze estimation error and the validity of the proxy itself.

Method

Overall Architecture

CAS is neither a new classification network nor a training objective that directly improves accuracy by optimizing CAS. It takes candidate classifiers and multiple source domains through source-domain circuit extraction, graph-kernel comparison, class-decomposed scoring, and cross-domain ranking aggregation, producing a model ranking and class-level vulnerability information. Target accuracy is used only to evaluate the ranking experimentally, not to compute source-domain CAS.

ACE extracts the actual circuits: a pretrained backbone is frozen, adapters are inserted at selected layers, and the adapters and classification head are trained; class-conditional activation interventions are then performed in each source domain. Graph kernels compare the resulting sparse directed graphs rather than scoring final predictions or activation matrices alone. The diagram below represents the diagnostic pipeline, not the network's forward architecture.

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    A["Candidate models<br/>and source domains"] --> B["Source-domain<br/>circuit extraction"]
    Y["Source labels"] -.->|Adapter training and class interventions| B
    B --> C["Graph-kernel comparison"]
    C --> D["Class-decomposed scoring"]
    D --> E["Cross-domain<br/>ranking aggregation"]
    E --> F["Model ranking<br/>and class diagnostics"]

Key Designs

1. Source-domain circuit extraction: identify class-relevant nodes and connections through interventions

ACE places adapters at at least three sites near the end of the backbone, allowing circuits to contain cross-layer structure rather than merely a list of important neurons in one layer. Each adapter is a two-layer bottleneck MLP with a residual connection. CNN and MLP circuits primarily contain adapter neurons, while ViT circuits additionally contain attention heads from frozen attention sublayers. These are local circuits over selected computational units, not exhaustive reconstructions of every pathway in the full backbone.

For a class, ACE first records the original output, then zeros each candidate unit's output and measures the mean change in that class's output. A positive change indicates support for the class and a negative change indicates suppression; selection uses the largest absolute changes, so suppressive units can also enter TopK. Node importance is defined as:

\[ \Delta_{\ell}^{(u)}(k)=\mathbb{E}_{\mathbf{x}\sim\mathcal{D}_{k}}\!\left[\hat{y}_{k}(\mathbf{x})-\hat{y}_{k}\!\left(\mathbf{x}\mid\mathrm{do}(u=\mathbf{0})\right)\right]. \]

Here \(\mathcal{D}_{k}\) is the set of source samples labeled \(k\). This data dependence matters: the current implementation of the paper's weight-derived prediction does not mean reading a parameter file without any samples or labels.

ACE next intervenes on selected units in an earlier layer and observes the mean activation changes of selected units in the next circuit layer. A directed edge is included only when the absolute change exceeds threshold \(\epsilon\); edges preserve the sign of the influence, and nodes preserve their importance. For ViT, attention-to-attention and adapter-neuron-to-adapter-neuron cross-layer edges are estimated separately; this step does not reconstruct every mixed attention–MLP connection.

This graph construction is closer to asking what changes when a unit is removed than activation correlation alone, but it remains sensitive to sampling, zero interventions, and discrete TopK selection. A single source-node intervention exposes changes in all selected downstream nodes, avoiding a separate forward pass for every candidate edge; random candidate subsets can also reduce node-scoring cost in high-dimensional layers.

2. Graph-kernel comparison: preserve pathway structure rather than only representation geometry

Class-specific circuits of the same model on two source domains are treated as graphs with node and edge attributes and compared using a normalized graph kernel \(\kappa\) with values in \([0,1]\). Node attributes are intervention importance scores, and edge attributes are causal influence weights. Normalization uses the geometric mean of graph self-similarities to reduce the effect of each graph's own scale.

The paper compares the Treelet Kernel (TK), Random Walk Kernel (RWK), and Optimal Transport kernel (OT), using TK by default. TK emphasizes shared local hierarchical substructures, RWK emphasizes common paths, and OT emphasizes global matching of node attributes and structure. In the reported experiments, TK separates same-class and cross-class circuit similarities more clearly instead of giving high scores to most circuit pairs.

The theoretical motivation is that metrics computed only through activation embeddings can conflate circuits with identical source behavior but different internal pathways. Graph comparison provides information for detecting such differences, but adopting a graph kernel does not automatically identify every functional difference: kernel discrimination, node annotation, and circuit extraction quality remain prerequisites.

3. Class-decomposed scoring: reward same-class preservation and penalize cross-class entanglement

Comparing class circuits across two domains produces a \(c\times c\) matrix, where \(c\) is the number of classes. Its entries are \(S_{ij}=\kappa(C_{1}^{(i)},C_{2}^{(j)})\): diagonal entries compare the same class across domains, while off-diagonal entries compare different classes across domains. Low diagonal similarity indicates circuit drift; high off-diagonal similarity indicates circuit entanglement. An average over the full matrix should not conceal these distinct failures.

\[ \mathrm{CAS}(\mathcal{C}_{1},\mathcal{C}_{2})=\frac{1}{c}\sum_{i=1}^{c}S_{ii}-\frac{1}{c(c-1)}\sum_{i\neq j}S_{ij}. \]

Both terms are means rather than unnormalized sums, so the larger number of off-diagonal entries does not overwhelm the same-class term merely because there are more classes. Uniform class weights make the score invariant to jointly renaming classes; this definition requires \(c\geq2\). The score lies in \([-1,1]\), and self-comparison of a circuit family need not equal 1 because different classes can share structure.

The formula determines the directions directly: increasing same-class similarity increases CAS, while increasing cross-class similarity decreases CAS. Section 4.3 of the main text reverses these coordinate-wise directions; this note follows Eq. (4) and the derivatives in Appendix B.2 while retaining the inconsistency warning. The perturbation-order guarantee covers comparable configurations with decreasing same-class and increasing cross-class similarity, not a correct total ordering of arbitrary graph changes.

Retaining the matrix also supports diagnosis: a low diagonal entry identifies a class whose pathway is unstable across domains, while a high off-diagonal entry identifies a particular pair of potentially entangled classes. This is more specific than one accuracy value, but a high average CAS does not imply that every class satisfies the authors' class-wise circuit-robustness thresholds.

4. Cross-domain ranking aggregation: estimation stability and OOD validity require separate arguments

The practical experiment averages CAS over all source-domain pairs, ranks 48 learners by that score, and compares the ranking with leave-one-domain-out accuracy. It needs neither unlabeled target samples nor target labels, but it does require multiple available source domains. Whether invariance among those sources covers a future target still depends on the task's distributional relationships.

The theory separately defines population CAS under a domain distribution and approximates it by a Monte Carlo average over \(M\) independently and identically sampled domains. For two learners with different population CAS values, let the mean score difference be \(m_{ij}>0\) and its cross-domain variance be \(\sigma_{ij}^{2}\). Chebyshev's inequality gives:

\[ p_{\mathrm{inv}}^{(M)}(\ell_i,\ell_j)\leq\frac{\sigma_{ij}^{2}}{M\,m_{ij}^{2}}=O\!\left(\frac{1}{M}\right). \]

This first guarantees convergence of empirical CAS rankings to population CAS rankings, not directly correct OOD rankings. Smaller gaps and greater domain-wise variation require more domains; learner pairs tied in population CAS are outside this nonzero-gap guarantee. Practical source-domain pairs share domains and cannot be treated without qualification as the independent domain samples in the theorem.

Recovering true OOD rankings additionally requires Assumption 1: \(g(\ell_i)>g(\ell_j)\Rightarrow\overline{\mathrm{CAS}}(\ell_i)>\overline{\mathrm{CAS}}(\ell_j)\), where \(g\) is expected accuracy under the domain distribution. The authors explicitly do not prove this finite-accuracy assumption from first principles. High observed rank correlation offers approximate support, not verification of strict monotonicity for every learner pair. The perfect-accuracy endpoint argument also depends on overlapping class supports and graph-kernel properties and cannot replace a full finite-error theorem.

Loss & Training

Candidate learners use ERM, IRM, CORAL, and DANN, respectively employing cross-entropy, an IRM gradient penalty, cross-domain covariance alignment, and a domain-adversarial objective; CAS itself is a post-training diagnostic score. The main experiment combines four architectures, four objectives, and three regularization settings: none, Dropout 0.3, and weight decay \(10^{-4}\), giving 48 configurations.

The appendix freezes the backbone and updates only adapters and the classification head; the default bottleneck ratio is \(r=4\), with zero-initialized up-projections. Training lasts 30 epochs with Adam at learning rate \(5\times10^{-5}\) and batch size 32 per domain; the DANN discriminator uses learning rate \(10^{-4}\). PACS defaults are \(K=30\) and edge threshold \(\epsilon=10^{-4}\); preferred \(K\) values for Office-Home and DomainNet are 60 and 164. Experiments use one NVIDIA A100 80GB.

Two budgets must be distinguished: CAS aggregation does not retrain a model, but the complete ACE pipeline is neither training-free nor cost-free. The appendix also lists VGG19 and Mixup, while some subsequent analyses mention Mixer/ViT-S, which do not fully match the main text's four-architecture, four-objective list. These additions should not silently expand the definition of the 48 main-experiment configurations.

Key Experimental Results

Main Results

The table follows Table 2 in the main text. Every entry is Spearman \(\rho_S\) between model rankings and target-accuracy rankings, not classification accuracy. PACS contains four domains and seven classes; each evaluation uses three source domains and holds out the fourth.

Method Photo Art Cartoon Sketch Mean
ATC 0.65 0.54 0.62 0.51 0.58
ProjNorm 0.72 0.61 0.68 0.55 0.64
ALine-D 0.79 0.70 0.78 0.61 0.72
CKA 0.81 0.48 0.58 0.45 0.58
SVCCA 0.28 0.21 0.21 0.22 0.23
RSA 0.12 0.14 0.19 0.11 0.14
CAS (TK) 0.91 0.86 0.93 0.83 0.88

CAS's 0.88 is the four-target mean; Figure 2's 0.93 is the Cartoon result and is not interchangeable with that mean. On Sketch, the gap over the next-best method, ALine-D, is 0.22. ATC, ProjNorm, and ALine-D use unlabeled target data; ALine-D uses the entire 48-model pool, while ProjNorm requires two extra training runs per fold, so information and compute budgets differ.

Table 3 additionally reports a CAS MAE of 2.14 percentage points. However, CAS is defined as a ranking proxy, and that comparison does not explain equally clearly how scores are calibrated into accuracy; MAE is therefore not treated as an intrinsic, calibration-free property. Appendix Figure 10/J gives Photo, Art, and Sketch correlations as 0.93, 0.91, and 0.88, conflicting with Table 2. This note retains Table 2's values rather than merging the inconsistent reports.

Ablation Study

The table selects PACS single-parameter ablations from Appendix I, with other parameters at their defaults. OOD Acc is the percentage reported in the original table. Accuracy changes accompanying extraction parameters should not be interpreted directly as CAS scoring improving the capability of an unchanged classifier.

Config Mean CAS OOD Acc (%) \(\rho_S\)
Default: \(K=30,r=4,\epsilon=10^{-4}\) 0.38 77.83 0.88
\(K=5\) 0.12 68.41 0.71
\(K=164\) 0.44 78.19 0.83
\(r=1\) 0.31 78.52 0.80
\(r=16\) 0.48 75.30 0.82
\(\epsilon=10^{-2}\) 0.22 74.61 0.79
\(\epsilon=10^{-6}\) 0.41 78.02 0.82

Key Findings

  • Retaining more nodes does not necessarily improve prediction: from \(K=30\) to \(K=164\), mean CAS rises from 0.38 to 0.44, but \(\rho_S\) falls from 0.88 to 0.83. Absolute score magnitude and ranking quality within a learner pool are different quantities.
  • Narrow adapters can artificially increase similarity: with \(r=16\), CAS is 0.48 rather than the default 0.38, yet accuracy and rank correlation are lower. Insufficient capacity can conceal domain differences, so comparisons across configurations must control extractor capacity.
  • In within-domain re-extraction analysis, PACS has mean noise dissimilarity 0.096, observed cross-domain dissimilarity 0.501, signal dissimilarity 0.405, and mean SNR 4.22; 47/48 learners have SNR above 1. This supports a signal beyond extraction noise, not error-free circuit recovery.
  • Source-domain-count experiments cover only \(M\in\{1,2,3\}\) for PACS/Office-Home and \(M\in\{1,2,3,4,5\}\) for DomainNet, with 10 repetitions per setting; trends over this finite range cannot independently prove an asymptotic law. The continuous LoRA domain experiment tests only SDXL ArtPainting/Photo style weights \(w\in[0,2]\), not every natural domain shift.

Highlights & Insights

  • Class structure is central, not merely a stronger similarity function. Preserving diagonal and off-diagonal identities distinguishes same-class drift from cross-class entanglement, and the matrix itself provides localized diagnostics.
  • Separating estimation consistency from proxy validity clarifies the theory. More source domains can reduce estimation noise but cannot repair a structural proxy that is not monotone with target accuracy to begin with.
  • Extraction-capacity ablations expose a counterintuitive risk: pathways that cannot express domain differences can also appear stable. Stability should be reported jointly with task sufficiency and re-extraction reliability rather than maximizing CAS alone.

Limitations & Future Work

  • Assumption 1 remains unproved; high correlation is not strict monotonicity or a generalization guarantee on an arbitrary unknown target distribution. The endpoint argument additionally needs support overlap and functional discrimination by the kernel; perfect classification does not automatically give a uniform separation gap for every structural kernel.
  • The current implementation requires source samples, labels, and adapter training and mainly concerns classification with a shared class space. It should not be described as fully data-free, applicable to open-set tasks, or directly applicable to language generation.
  • Local ACE node selection, zero interventions, and edge thresholds affect the score. The appendix estimates noise through within-domain re-extraction, but additive noise subtraction for arbitrary graph-kernel dissimilarities needs additional statistical conditions; corrected differences are not direct observations of true structural change.
  • Experimental descriptions conflict: the main text describes graph-embedding inputs for baselines, whereas Appendix D describes final-layer features; appendix architectures do not fully match the main list; PACS figure captions conflict with Table 2. Reproduction should first unify input objects and the learner pool before comparing values.
  • Further work should test fine-grained ranking within the same training objective and provide ranking confidence intervals that jointly account for source re-extraction and domain sampling. Reporting that near-tied models cannot be distinguished is more reliable than forcing a strict ranking.
  • vs CKA / SVCCA / RSA: these primarily measure representation geometry; CAS compares class-specific circuit graphs and explicitly penalizes cross-class similarity. It requires additional interventions and graph construction, and its benefit depends on extracted circuits containing task-relevant structure.
  • vs ATC / ProjNorm / ALine-D: these estimate accuracy using unlabeled target data; CAS predicts model rankings from structural relationships among sources. The former directly observe the target distribution, while the latter can support selection before a target is visible; they are not substitutes under identical information budgets.
  • vs circuit-discovery methods such as ACE: an extractor identifies nodes and edges influencing behavior; CAS asks whether those circuits persist across domains while preserving class separation. A possible extension is to construct control models that preserve accuracy while changing redundant pathways and test whether CAS mistakes harmless rerouting for OOD fragility; this remains an untested research direction.

Rating

  • Novelty: 4/5. Class-conditional circuit decomposition provides a structural perspective on OOD model selection distinct from representation similarity.
  • Experimental Thoroughness: 4/5. Multiple datasets, kernels, extraction settings, and noise analyses are covered, but the finite learner pool and conflicting descriptions limit the strength of conclusions.
  • Writing Quality: 3/5. The metric is clearly defined, but monotonicity directions, input protocols, and some figure/table values require clarification.
  • Value: 4/5. Useful as an auxiliary selection signal before target data is available, not a substitute for target validation or a safety guarantee.