Source-linked AI summary
Canalization Before Generalization: Grokking as a Dynamical Probe
Yiming Lin
TL;DR
Overparameterized networks can fit training data while selecting different functions, and grokking exposes this selection before visible generalization. The paper scans temporary WD pulses across that plateau and finds stable dose-ordered timing effects alongside collapsing test-loss barriers, calling this canalization of function selection.
Problem
Grokking provides a window for studying how optimization selects among multiple training-fitting solutions with different unseen-sample behavior.
Method
The paper scans temporary fixed-duration WD pulses across the pre-generalization plateau and measures their later effects on generalization time and function selection.
Results
Across three algorithmic grokking tasks, responses shift from unordered early in the plateau to stable dose ordering before visible generalization, while test-loss barriers collapse toward zero.
Takeaways & Limitations
The paper terms this combination of increasingly constrained solution selection and persistent, adjustable timing sensitivity the canalization of function selection.
Takeaways & Limitations
The experiments use three small algorithmic tasks with sharp, single generalization transitions, so applicability to broader optimization settings remains unclear.
Abstract
from arXiv · showhide
For overparameterized neural networks, many solutions can fit the training data equally well while behaving very differently on unseen samples. Grokking separates training fit from visible generalization, providing a window for studying how this selection develops during training. We sweep short, fixed-duration weight-decay (WD) pulses across this plateau and measure how they shift later generalization time. Across three grokking tasks, these shifts are unordered early in the plateau but later form a stable dose ordering, with stronger WD increases leading to earlier generalization and stronger WD decreases leading to later generalization. This ordering emerges before visible generalization in all three tasks. Meanwhile, test-loss barriers between perturbed and baseline generalization checkpoints collapse toward zero while the ordered timing effects persist. We call this combination of increasingly constrained solution selection and persistent dose-ordered timing sensitivity the canalization of function selection.
1 Introduction
The paper uses grokking’s separation between training fit and visible generalization to study how optimization selects among equally training-fitting solutions. Temporary WD pulses reveal a pre-generalization transition from unordered responses to stable, dose-dependent timing effects.
- Motivation: Overparameterized networks can fit training data equally well while selecting functions with different unseen-sample behavior.This makes optimization’s implicit bias central to understanding generalization.
- Motivation: Grokking creates a long low-test-performance plateau after training fit and before rapid generalization.This temporal separation enables perturbing optimization before visible generalization.
- Motivation: Earlier perturbation studies suggest training becomes less sensitive to local changes as it enters a more stable and constrained regime.The paper uses downstream effects of local perturbations to probe this dynamical change.
- Motivation: Weight decay is a natural intervention because its effects depend on application timing and influence whether and when grokking generalizes.The intervention therefore targets a training variable with established temporal relevance.
- Approach: The framework scans temporary WD pulses across the pre-generalization plateau and measures their later effects on generalization timing and selected functions.Each pulse changes WD briefly, restores baseline training, and continues optimization.
- Main result: Across three algorithmic tasks, responses are initially unordered but later become dose-ordered before visible generalization; stronger positive and negative pulses advance and delay generalization, respectively.Test-loss barriers collapse toward zero while ordered timing shifts persist, and the pattern reproduces across initializations and tasks.
2 Grokking as a Dynamical Probe
The experiments construct WD-response maps by branching from saved training states, applying short changes in weight decay, and measuring shifted threshold-crossing times. The maps compare pulse timing and magnitude across three algorithmic grokking tasks.
- Tasks: The study tests parity matching, sparse parity, and modular addition as algorithmic grokking tasks with distinct training and test index pairs.These settings require learning a rule that transfers to unseen combinations.
- Intervention: AdamW permits temporarily changing weight decay without changing the loss function or gradient computation.Weight decay is decoupled from AdamW’s loss-gradient update.
- Sliding-Window WD Interventions: For each baseline trajectory, a fixed-duration intervention window slides along the pre-generalization plateau from the first time full training accuracy is reached.The start-time grid is defined by t0 = tmin + kδt, with window length τ and stride δt.
- Sliding-Window WD Interventions: Each pulse branches from a saved full training state containing model, optimizer, random-number-generator, and data-sampling states.This preserves the branch’s complete training state at the intervention time.
- Intervention: Across all tasks, the analysis includes 70 baseline runs, 23,926 pulse starts, and 239,260 intervention branches.The same short intervention design is evaluated across many baseline trajectories and pulse settings.
- Generalization-Time Response: The response ΔTα measures how a pulse changes the first time the perturbed branch reaches the baseline’s test-accuracy threshold.Negative values indicate earlier generalization, while positive values indicate later generalization.
- Generalization-Time Response: The WD-response map indexes pulse start time and perturbation magnitude, with each location reporting the resulting advance or delay in generalization.This map exposes how intervention effects vary over training.
- Generalization-Time Response: Figure 2 summarizes dose–response curves across Task 1 runs and shows corresponding response maps for Tasks 2 and 3.The accompanying analysis compares response structure across runs, tasks, and training phases.
3 Results
Across tasks, WD-pulse responses reorganize from unordered early effects into stable dose-ordered timing shifts before visible generalization. Meanwhile, perturbed and baseline checkpoints become increasingly connected by low-barrier paths, while timing sensitivity persists.
- WD responses become ordered before visible generalization: Early plateau responses show no stable relationship between WD-pulse direction or magnitude and generalization-time shifts.By the middle of the plateau, responses begin forming a persistent dose-dependent structure.
- WD responses become ordered before visible generalization: Dose–response linearity increases across 30 Task 1 runs, approaching 1 near the end of the interval.This indicates that WD responses gradually develop an approximately linear relationship with perturbation dose across initializations.
- WD responses become ordered before visible generalization: Across tasks and initializations, WD responses transition from unordered early patterns to stable dose ordering before baseline test accuracy visibly rises.Task 2 and Task 3 reproduce the same qualitative behavior despite differing time scales and response details.
- Loss barriers collapse while ordered WD responses persist: Test-loss barriers between perturbed and baseline generalization checkpoints decrease toward zero as interventions move later.The decline is clearest in Task 1 and appears as a weaker late-stage trend in Tasks 2 and 3.
- Loss barriers collapse while ordered WD responses persist: Ordered timing effects persist after test-loss barriers become close to zero: increasing WD advances generalization, while decreasing WD delays it.The timing shift remains ordered by pulse magnitude even when checkpoints are connected by nearly low-loss paths.
- Loss barriers collapse while ordered WD responses persist: After ordering emerges, perturbed and baseline update-direction trajectories approximately overlap, differing mainly in timing along the trajectories.This pattern is observed in PCA projections of parameter-update directions.
4 Discussion
The paper treats grokking as a temporal window for probing function selection with transient WD interventions. Across small algorithmic tasks, later perturbations become less able to separate selected solutions while still changing generalization timing, motivating the term canalization of function selection.
- Discussion: Overparameterized networks can fit training data with multiple functions that differ on unseen samples, making optimization’s implicit bias central to generalization.Grokking separates training fit from visible generalization, allowing this selection process to be observed before generalization.
- Discussion: The WD-response map reveals a dynamical transition before visible generalization across several small algorithmic grokking tasks.Responses move from unstable early effects to stable dose ordering later in the plateau.
- Discussion: Later perturbations become less able to steer training toward linearly separated generalization solutions, while perturbation strength still adjusts generalization timing.The paper presents this combination as an analogy to Waddington’s increasingly narrow and deep developmental valleys.
- Discussion: The canalization analogy describes perturbations losing influence over lateral direction while retaining the ability to speed or slow progress along a trajectory.This interpretation connects collapsed loss barriers with persistent ordered timing responses.
- Discussion: Perturbation responses can reveal changes in training dynamics well before standard progress measures show them.The discussion relates stable dose ordering qualitatively to broader ideas about trajectory selection under irreversible optimization dynamics.
- Discussion: The evidence is limited to three small algorithmic tasks with single, relatively sharp generalization transitions and long pre-generalization plateaus.Broader tasks may require a more general definition of the generalization transition.
- Discussion: Using identical pulse parameters across tasks produced broadly similar response structures without task-specific intervention tuning.This supports the reported robustness within the studied settings, not beyond them.
- Discussion: The paper names this pre-generalization reorganization the canalization of function selection.The term refers to increasingly constrained solution selection alongside persistent, ordered timing sensitivity.
LLM Usage Claim
The authors used LLMs for language polishing and a reproducibility sanity check, while retaining responsibility for verification, ideas, experiments, and conclusions.
- LLM Usage Claim: LLMs assisted with language polishing and a reproducibility sanity check, but the authors checked the output and retained ownership of the research.The AI agent attempted replication from the paper description to assess clarity.
A.1 Three Tasks
The paper evaluates WD-pulse responses across three delayed-generalization tasks and selects baseline runs using sustained accuracy and final-performance criteria. The main analysis includes task-specific architectures, datasets, and screened trajectory sets.
- Experimental setup: Three algorithmic tasks exhibit delayed generalization and are evaluated with full-batch AdamW training without a learning-rate scheduler.Metrics are evaluated every 25 epochs until test accuracy reaches 0.60, then every epoch.
- Task 1: Parity match: Task 1 is parity matching over two indices from 0 to 47, with labels indicating whether the indices share parity.Inputs are concatenated 48-dimensional one-hot encodings.
- Task 2: Sparse parity: Task 2 is sparse parity on 50 binary features, with labels determined by three fixed relevant coordinates and 47 distractors.The relevant coordinates are selected using a fixed seed, and the model is a four-hidden-layer tanh MLP.
- Task 3: Factored modular addition: Task 3 is factored modular addition modulo 67, using unordered input pairs and a shared embedding-based architecture with tied unembedding.The dataset uses 570 training examples and 1708 test examples.
- Baseline run selection: 105 of 1,000 Task 1 runs satisfy all screening criteria, and the main analysis uses the first 30 qualifying runs; Task 3 uses seeds 1–16 directly.The criteria require a sufficiently long plateau, a concentrated generalization transition, and final test accuracy of at least 0.95.
B.1 Dose–Response Linearity Analysis
The analysis aligns pulse starts by training phase and measures whether five WD doses produce ordered generalization-time shifts. It aggregates chance-corrected linearity across runs while tracking weight-norm changes and reporting a late-phase coverage limitation.
- Phase alignment: Pulse phase ϕ runs from training accuracy reaching 1.0 to baseline test accuracy reaching 0.60, enabling comparison across runs with different time scales.The phase interval is divided into bins of width 0.01.
- Dose–response linearity: Five perturbation magnitudes are analyzed separately for positive and negative WD changes at each pulse start.The squared Pearson correlation R2 measures dose–response linearity, with chance correction based on response permutations.
- Aggregation: Phase-bin summaries use within-run medians, averages across runs, and 4,000 bootstrap resamples to form 95% confidence intervals.Random initializations are treated as independent samples.
- Coverage limitation: The final four phase bins, ϕ ≥ 0.965, are omitted because each contains fewer than 15 valid runs; plotted curves end at ϕ = 0.955.Pulse starts are treated as missing when branches fail to reach the required threshold within the continuation limit.
- Weight-norm comparison: Relative weight norm is measured against the first checkpoint reaching 1.0 training accuracy and aggregated using the same phase-binned bootstrap procedure.Synchrony is assessed by correlating changes in mean chance-corrected linearity with changes in log relative weight norm.
C.1 Implementation of WD Pulse Interventions
The WD-pulse protocol resumes exact saved training states, applies fixed-duration decay perturbations across the pre-generalization plateau, and then restores baseline conditions. Generalization checkpoints are compared using straight-line test-loss barriers.
- State-controlled intervention: Each pulse branch resumes from the complete saved state, including parameters, AdamW moments, random-number state, and data-sampling state.Branches at the same pulse start differ only in the weight-decay coefficient.
- Sliding-window protocol: Pulse starts are scanned every 25 epochs from the first 1.0 training-accuracy checkpoint to 1,000 epochs before baseline reaches 0.95 test accuracy.The pulse duration is fixed at τ = 500 epochs.
- Post-pulse continuation: After each pulse, the branch returns to baseline training conditions and continues until epoch 40,000 while test accuracy is monitored at increasing temporal resolution.Branches that fail to reach the threshold within the limit are marked as missing.
- Test-loss barrier: Barrier comparisons connect the baseline and perturbed checkpoints where each first reaches 95% test accuracy, using 51 equally spaced parameter interpolates.Test loss is evaluated on the full test set along the interpolation path.
- Test-loss barrier: The test-loss barrier is the maximum excess test loss along the path relative to linear interpolation between endpoint losses.This excludes the endpoint loss difference itself from the barrier.
D Update-Direction Analysis
The update-direction analysis compares baseline and WD-pulse branches before, at, and after ordering onset using common per-run PCA coordinates. After onset, trajectories overlap approximately, with timing rather than direction distinguishing branches.
- Ordering onset: Ordering onset is the first checkpoint where all eight dose–response sequences are complete, directionally consistent, and have absolute Spearman correlation at least 0.8.The sequences cover four test-accuracy thresholds for both positive and negative WD perturbations.
- Trajectory construction: Each run contributes the baseline and 30 perturbed trajectories from three pulse starts, with 10 WD doses per start.The starts represent pre-ordering, ordering-onset, and post-ordering phases.
- Update directions: Update directions are centered differences over t − 25 to t + 25 and are normalized, so they encode movement direction rather than magnitude or parameter position.Directions are sampled every 50 epochs for joint PCA visualization.
- Results: Before ordering, pulse branches deviate clearly from baseline; from ordering onset onward, their projected trajectories approximately overlap.The same trend appears across all three analyzed Task 1 runs.
- Results: After ordering emerges, branches mainly differ in when they traverse corresponding parts of a shared projected trajectory rather than following clearly distinct directional paths.This timing distinction is consistent across the three runs.
E Comparison with Existing Progress Measures
In Task 1, standard checkpoint-based progress measures provide little advance warning of generalization, whereas the WD response becomes dose-ordered substantially earlier.
- Comparison: The analysis tests whether restricted/excluded losses, parameter movement, and hidden-layer Fourier energy reveal progress before visible generalization.All measures are computed from checkpoints saved during training.
- Measurement interpretation: Restricted BCE isolates the rule-aligned output component, whereas excluded BCE tracks the remaining training-fit component after removing that component.Restricted logits retain the global bias and joint-parity rule component; excluded logits remove the joint-parity component.
- Nanda-style decomposition: Restricted test BCE stays near chance through almost the entire plateau and declines only about 50 epochs before held-out test BCE.Excluded train BCE continues decreasing and rises again near the generalization transition.
- Barak-style amplification: Parameter movement saturates early, while fourth-layer joint-parity Fourier energy remains near zero until around visible generalization.The joint-parity fraction is computed from the (24, 24) Fourier mode relative to total nonconstant Fourier energy.
- Comparison: The WD response’s dose-dependent structure is established roughly 4,000 epochs before baseline test accuracy crosses α = 0.90.This provides substantially earlier dynamical information than the static progress measures examined in Task 1.
F Additional Runs Across Random Initializations
Across additional random initializations, later pulse timing reduces test-loss barriers while preserving directional and dose-dependent generalization-time responses. The pattern is broadly reproduced, with especially consistent responses across unscreened Task 3 seeds.
- Additional runs: The reported relationship is broadly reproduced across 30 Task 1, 24 Task 2, and 16 Task 3 runs.Endpoint selection, interpolation points, and test-loss barrier definitions match the main analysis.
- Additional runs: Later pulse starts reduce the test-loss barrier between perturbed and baseline generalization checkpoints while preserving directional and dose-dependent ΔT0.95 structure.Some Task 1 perturbed branches do not reach 0.95 test accuracy within the observation limit.
- Task 1: Task 1 exhibits substantial initialization-to-initialization variation in perturbation sensitivity.This heterogeneity coexists with the broader persistence of directional and dose-dependent response structure.
- Task 3: Task 3 shows particularly consistent WD timing responses and test-loss barriers across seeds 1–16 despite using no seed screening.Task 1 and Task 2 instead use predefined baseline-screening criteria.