Large language models are trained with backpropagation, whose global gradient coordinates all layers but forces each to hold its activations and wait for the gradient to pass back through every deeper layer. Conventional local learning removes this update locking by training each module to predict the target through its own readout, but has not scaled to billion-parameter pretraining. We identify these private readouts as a key weakness, since they leave each module without information from deeper modules. We propose Shared-Output LOcal learning (SOLO), which replaces them with a shared, read-only copy of the final module's readout, the only one trained on the output of the whole network. Taken from the previous step, the copy transmits information from the final module without passing gradients between modules or reintroducing update locking. SOLO approaches backpropagation on Transformers of 340M to 2B parameters pretrained on 15B tokens, staying within one point in average zero-shot accuracy with a perplexity gap that narrows with scale. Readout ablations attribute SOLO's improvement over private readouts to sharing. Without update locking, each of p pipeline stages holds activations for O(1) micro-batches instead of O(p). The freed memory permits larger micro-batches, which reach up to 1.44x the best measured throughput of pipeline backpropagation on the same partition. To our knowledge, SOLO is the first local learning method to show such memory and throughput gains in billion-parameter language-model pretraining. Local learning thus becomes a practical alternative to backpropagation for large-scale pretraining.
Figures & tables
Figure 1: Three ways to train a network of K modules. (a) BP trains every module with the final loss L through the terminal readout W . (b) Local learning stops the gradient between modules and gives each auxiliary head a trainable private readout Wk with its own loss ℓk . (c) SOLO keeps the stop-gradients, but every head predicts through sg(Wˉ) , a read-only copy of W from the previous step. Gradients flow through the copy into its head but do not update it; only L updates W , so no head waits for its current update.
Variant
Readout of head k
Shared
Source
Trainable
rand
Wkrand , frozen
no
random
0
priv
Wkpriv , learned
no
own head
Vd
SOLO
Wˉ , read-only copy
yes
live terminal readout
0
pretrained (SOLO)
WSOLO⋆ , frozen
yes
fully trained SOLO
0
pretrained (BP)
WBP⋆ , frozen
yes
fully trained BP
0
Table 1: Readout sources. Each variant changes only the readout of the auxiliary heads; ϕk and τk are common to all variants. The last column counts trainable readout parameters per head.
Perplexity ↓
Zero-shot accuracy (%) ↑
Cost
Scale
Variant
Wiki.
LMB-p
LMB-a
PIQA
Hella.
Wino.
ARC-e
ARC-c
Avg.
tok/s ↑
Mem. GB ↓
340M
BP (DP)
28.04
39.5
31.8
64.2
34.6
49.9
43.4
25.0
41.5
247k
29.6
SOLO, K=2
29.05
44.6
29.8
64.0
34.0
51.9
45.6
23.6
41.5
208k
15.8
priv , K=2
29.47
46.3
29.2
65.1
33.8
51.0
44.4
22.5
41.0
209k
15.9
SOLO, K=4
31.19
47.9
28.6
63.3
33.1
51.3
46.0
23.7
41.0
156k
11.2
priv , K=4
31.79
55.6
26.6
62.5
32.9
53.4
44.1
23.4
40.5
157k
11.5
Table 2: Pretraining on 15B SlimPajama tokens. Quality values are three-seed means; bold marks the best local-learning variant per column. Throughput and peak memory per GPU are measured on four A100s at micro-batch 8 under data parallelism, with BP sharded at 2B. LMB-p and LMB-a are LAMBADA perplexity and accuracy.
Figure 2: Validation perplexity during pretraining at (a) 340M and (b) 1.3B; insets show the final 9 to 15B tokens on a linear scale.
Figure 3: Readout source across scale and modality. (a) Perplexity gap to BP on the one-epoch WikiText-103 sweep with K=4 , three-seed means. pre-SOLO and pre-BP are the pretrained variants, which freeze the terminal readout of a fully trained SOLO or BP model. (b) Cost of a random readout relative to a private one, as the relative increase in validation perplexity (language) or test error (vision), against the number of output classes per unit width. Numbers are in Appendix D .
K
Variant
Perplexity
Cost
Val ↓
Test ↓
tok/s ↑
Peak mem (GB) ↓
Δ par (M)
1
BP (end-to-end)
23.15 (38.30)
23.59 (41.59)
88k
16.6
—
2
SOLO
23.49 (38.95)
24.02 (42.48)
53k
12.9
+4.8
priv
23.64 (39.25)
24.20 (42.87)
52k
13.0
+16.7
pretrained (SOLO)
23.59 (39.10)
24.06 (42.56)
52k
13.0
+4.8\lx@sectionsign
pretrained (BP)
23.82 (39.59)
24.30 (43.07)
52k
12.9
+4.8\lx@sectionsign
Table 3: Readout sources under a published WikiText-103 recipe, mean over three seeds. Perplexity per subword token, word-level in parentheses; Δ par denotes trainable parameters beyond BP’s 62.3M, and § an additional 11.9M-parameter frozen pretrained readout.
Figure 4: Gradient alignment on the 40M WikiText-103 sweep, three seeds per variant. (a) Alignment during training; (b) alignment at the end, by module; (c) norms of the three terms of Eq. ( 4 ) at the end, divided by the norm of the gap; (d) terminal perplexity against alignment. The terms can partly cancel, so their norms are not shares of the gap. Setting aside τk , the readout term is zero under SOLO by construction.
Figure 5: SOLO against pipeline BP under the same partition. (a, b) Schedules of 1F1B and SOLO with p=4 stages and eight micro-batches, ignoring the auxiliary heads. Each cell is one forward or backward pass of one micro-batch; 1F1B takes 22 slots and SOLO 19, and the right column gives the peak number of micro-batches in flight per stage. Only one step is drawn; with S>1 , SOLO’s idle slots at the start and end overlap with the neighboring steps. (c) Throughput relative to 1F1B at p=8 ; VPP- v is interleaved 1F1B with v virtual stages per GPU, and S is the refresh period of the copy in steps. (d) Peak memory on the most loaded GPU, with model state in gray and activations in color. (e) Communication volume per step. Panels (d) and (e) use M=24 ; data-parallel runs are shown for reference.
Appendix figures & tables26 assets
Supplementary material from the paper’s appendix.
Appendix
Figure 6: Relative perplexity gap during pretraining, SOLO validation perplexity divided by BP’s at the same token count, for (a) 340M and (b) 1.3B, from the single-seed training-log curves of Figure 2 ; ratios at matching evaluation steps, no interpolation or smoothing; the dotted line is parity.
Validation perplexity ↓
Variant
40M
64M
128M
256M
340M
512M
V/d=85
64
43
32
32
26
BP (end-to-end)
67.08
54.91
44.46
40.47
39.51
36.81
rand
98.88
81.63
63.76
55.34
54.12
48.99
priv
75.34
62.89
51.10
44.57
43.82
40.24
SOLO
70.62
59.09
48.01
42.49
41.48
38.01
Appendix
Table 4: Language readout sweep across scale, WikiText-103 validation perplexity (word-level, V=32,768 , one epoch, K=4 ), three-seed means. The sub-row gives V/d , which falls from 85 to 26 from left to right. The shaded BP row is the end-to-end reference, and bold marks the best local-learning variant per column. Rows run from worst to best, an order that holds at every size. The two pretrained variants freeze the terminal readout of a fully trained SOLO or BP model (pre-SOLO and pre-BP in Figure 3 ); they are diagnostic probes, not training methods.
ResNet-32 ( d=64 )
ResNet-50 ( d=2048 ) †
Variant
C-10
C-100
Tiny-IN
C-10
C-100
Tiny-IN
Δ par ‡
C/d=0.16
1.56
3.1
0.005
0.049
0.098
R50,M
rand
92.53
67.28
44.30
95.80
78.56
63.54
+41.1
priv
92.51
67.97
47.35
95.71
79.43
63.60
+47.2
SOLO
92.57
67.22
46.14
95.34
79.13
63.31
+41.1
Pretrained SOLO
92.27
67.63
46.48
95.43
79.55
64.59
+41.1
Appendix
Table 5: Vision readout sweep, terminal test accuracy (%), backbones split into 15 gradient-isolated modules. ResNet-32 columns are three-seed means; ResNet-50 columns are single-seed scouting runs ( † ). Datasets are ordered by classes per unit width, C/d within each backbone. Bold: best per column. ‡ Trainable parameters over BP on ResNet-50 Tiny-ImageNet, millions.
Backbone
Variant
Split
Top-1
Peak mem
ResNet-50
BP
—
76.49
25.3 GB
priv , K=2
[5∣11]
76.27
18.3 GB ( −28% )
SOLO, K=2
[5∣11]
76.22
18.3 GB ( −28% )
priv , K=4
[2,3,6,5]
75.15
14.0 GB ( −45% )
SOLO, K=4
[2,3,6,5]
75.03
13.7 GB ( −45% )
ResNet-101
BP
—
76.87
42.3 GB
Appendix
Table 6: ImageNet-1k ( 2242 , batch 256, 90 epochs; mean over three seeds). Memory change is relative to the end-to-end baseline of the same backbone. Bold: best per column within each backbone.
(a) Width sweep, V=32,768
Perplexity
Size
d
V/d
rand
priv
Penalty (%)
40M
384
85.3
98.88
75.34
31.25
64M
512
64.0
81.63
62.89
29.80
128M
768
42.7
63.76
51.10
24.77
256M
1024
32.0
55.34
44.57
24.16
Appendix
Table 7: Language points of Figure 3 b: the random-readout penalty, (ppl\textscrand/ppl\textscpriv−1)×100% , against the number of output classes per unit width, V/d . (a) Width sweep at fixed vocabulary, the runs of Table 4 . (b) Vocabulary sweep at fixed width; for this sweep we report only the penalty. The two sweeps meet at V/d=64 ( d=512 , V=32,768 ), where they give 29.80 and 29.72%, a difference within run-to-run noise.
Scale
Variant
H=0
H=1
H=2
H=3
H=5
H=8
BP
40M, 8L, d=384
SOLO
76.13
72.06
70.05
69.58
70.10
—
67.08
priv
81.75
76.48
75.36
73.77
73.91
—
128M, 16L, d=768
SOLO
49.93
48.31
47.38
46.75
46.22
—
44.46
priv
52.98
51.87
51.06
49.90
49.44
—
256M, 16L, d=1024
SOLO
43.95
42.90
41.95
41.56
41.18
41.48
40.47
priv
46.94
45.54
44.54
44.10
43.87
44.12
Appendix
Table 8: Auxiliary-head depth on the WikiText-103 sweep, validation perplexity (word-level, V=32,768 , one epoch, K=4 , seed 42) with H blocks per head; H=8 is run at 256M only, as an over-deep control. The last column is the end-to-end BP reference of Table 4 . This grid re-reads the shared copy every step, whereas the sweep of Table 4 refreshed it every 100 steps, so the H=2 cells differ slightly. Bold: best local variant per row pair.
Figure 7: Auxiliary-head depth on the WikiText-103 sweep (Table 8 ). (a) Gap to BP, 100(ppl−BP)/BP , for SOLO (solid) and priv (hatched) at H∈{0,1,2,3,5} , and 8 at 256M. (b) Advantage of the shared readout, 100(\textscpriv−SOLO)/SOLO ; the dotted line is the mean over the sixteen cells, 6.6%. Bars are seed-42 runs; error bars are the seed-to-seed standard deviation estimated from paired second-seed runs, 0.40 ppl at 40M and 0.12 at 256M, and inferred at 128M.
Perplexity ↓
Parameter-space cosine ↑
Arm
Readout of module k<K
H=0
H=2
H=0
H=2
BP
end-to-end, no auxiliary readout
67.08
—
—
SOLO
sg(Wˉ) , shared, live copy
76.16
70.08
.80 / .78 / .84
.68 / .80 / .90
priv
Wk , private, learned
81.71
75.36 †
.75 / .64 / .57
—
rot
sg(Wˉ)Rk , same content, per-module basis
89.41
77.44
.80 / .67 / .61
.63 / .74 / .79
rshare
Wrand , frozen, shared by all heads
104.36
94.62
.50 / .31 / .28
.46 / .48 / .37
Appendix
Table 9: Content and basis of the auxiliary readout on the 40M WikiText-103 sweep (8 layers, d=384 , V=32,768 , one epoch, K=4 , copy refreshed every step; seed 42, paired-seed σ=0.40 ppl). Terminal validation perplexity with H=0 and H=2 head blocks, and the parameter-space cosine between the local and the end-to-end gradient of modules 1 to 3 at the end of training. rot uses the terminal readout in a fixed random orthogonal basis Rk per module, so its range equals SOLO’s; rshare freezes one random matrix shared by all heads, rand one per head. † priv at H=2 is the matched run of Table 8 ; the two scripts reproduce each other to within 0.04 ppl.
Figure 8: The five variants of Table 9 at H=0 and H=2 , terminal validation perplexity on the 40M sweep; BP dashed. A per-module basis of the same readout ( rot ) costs 13 ppl at H=0 and 7 at H=2 ; one shared random matrix ( rshare ) beats independent ones ( rand ) by 12 and 4.
WikiText ppl
Battery ppl
Variants
8L
14L
20L
24L
8L
14L
20L
24L
SOLO, K=4
28.20
24.83
24.17
23.59
29.4
28.4
29.0
27.8
SOLO, K=2
—
24.16
—
22.29
—
22.6
—
21.7
BP logit lens
—
—
—
21.61
1347
290
45.9
19.5
Appendix
Table 10: Exits against the logit lens at 1.3B, ppl by depth, mean over three seeds; 24L is the terminal. On the battery (geometric-mean ppl over 27 prompts) every SOLO exit is within 2 ppl of its terminal and the BP lens is not; SOLO exits match the BP terminal’s top prediction 71 to 75% of the time, the lens 16 to 51%.
Figure 9: One passage decoded word by word by every exit of the 1.3B K=4 model. A row gives one readout’s next-word prediction, the word itself when correct and the prediction in italics when not; shading is the change in logp of the true word from the previous exit (legend in the figure). Rows are labeled by model depth; 6, 12, and 18 layers are the 8L, 14L, and 20L exits of Table 10 .
Figure 10: Terminal-readout probes across depth on the 40M WikiText-103 sweep. The output of each module is decoded with the model’s own terminal readout and compared with BP’s logit-lens distribution at the matched depth. (a) Top-1 prediction agreement. (b) KL(p∥pBP) on a logarithmic scale. Means over three seeds; bands are the minimum and maximum across seeds. SOLO and pretrained (local) agree more and diverge less at intermediate boundaries than priv , rand , and pretrained (BP); the differences narrow at the terminal modules ( k=3 ). The probe bypasses the auxiliary heads, so it measures terminal-readout decoding of intermediate representations, not the predictions of the trained auxiliary exits.
S
refreshes/epoch
val ppl
exits, shallow to deep
ens
test ppl
Δ val
1
7242
65.93 (129.0)
71.14 / 68.12 / 67.43 / 65.93
66.70
67.55
—
10
724
65.69 (128.5)
71.00 / 67.88 / 67.17 / 65.69
66.48
67.33
−0.24 ( −0.4% )
50
145
66.16 (129.5)
71.44 / 68.26 / 67.59 / 66.16
66.92
67.73
+0.23 ( +0.3% )
100
72
67.16 (131.8)
72.19 / 69.21 / 68.59 / 67.16
67.91
68.71
+1.23 ( +1.9% )
200
36
67.93 (133.6)
73.04 / 69.99 / 69.35 / 67.93
68.67
69.48
+2.00 ( +3.0% )
Appendix
Table 11: Refresh period S of the shared readout copy (WikiText-103 recipe of Table 3 , K=4 , one epoch of 7242 steps; mean over three seeds). Subword perplexity with word-level in parentheses; exits listed shallow to deep; ens is the exit ensemble. Δ is the change in validation perplexity relative to S=1 .
Figure 11: Relative change of the terminal readout over one refresh period, rS(t)=∥W(t)−W(t−S)∥F/∥W(t−S)∥F , along the WikiText-103 runs of Table 11 for S∈{1,10,50,100,200} . The shaded band is the learning-rate warmup; the fall at the end follows the cosine decay.
Figure 12: Throughput (a) and per-GPU memory (b) against model size on four A100s at matched micro-batch 8, for replicated (DDP) and sharded (FSDP) backpropagation and the SOLO pipelines. DDP does not fit at 2B; its memory there is extrapolated (dashed).
Interconnect
BP-DDP (tok/s)
SOLO K=4 (tok/s)
ratio
pipeline busy
NVLink (reference)
250k
163k
0.65
0.96
socket transport, unshaped
69.2k
150.1k
2.2
0.92
Appendix
Table 12: Throughput off RDMA-class interconnect (340M, four GPUs per arm), a single-node emulation. Ratio is pipeline over DDP; busy is the fraction of wall-clock the pipeline stages spend computing. The NVLink row is this session’s own reference and differs from the sweep of Table 2 (247k, 156k) by 1 to 4%.
BP
SOLO
model
L
d
K
H
total
layers
total
aux. blocks
readouts
ρ
340M
24
1024
2
2
2.61
2.42
2.94
0.20
0.33
1.127
1.3B
24
2048
2
2
8.85
8.46
9.82
0.70
0.66
1.109
2B
24
2560
4
2
13.33
12.83
17.52
3.21
1.48
1.315
96-block 1.2B
96
1024
8
1
8.51
8.46
9.36
0.62
0.29
1.100
Appendix
Table 13: Training FLOPs per token. One forward pass counts as one unit, so a block costs 3b with b=12d2+2Td multiply-accumulate operations per token, the terminal readout costs 3r with r=Vd , and an auxiliary readout costs 2r , since its weights are a detached copy and no gradient with respect to them is computed. One multiply-accumulate is two FLOPs. Embedding lookups are excluded. The SlimPajama rows use context 2048 and a 32 k vocabulary and assume two-block auxiliary heads; the last row is the configuration of Table 15 . The auxiliary blocks, not the readouts, account for most of the difference.
run
parallelism
GPUs
tok/s
GFLOPs/token
MFU
1.3B, micro-batch 8
BP, replicated DP
4
82,500
8.85
58.5%
1.3B
BP, sharded DP
4
92,100
8.85
65.3%
2B
BP, sharded DP
4
62,000
13.33
66.2%
96-block 1.2B
BP, replicated DP
8
65,471
8.51
22.3%
96-block 1.2B
BP, 1F1B pipeline
8
59,441
8.51
20.3%
96-block 1.2B
SOLO, pipeline
8
59,770
9.36
22.4%
Appendix
Table 14: Model FLOPs utilization of the measured runs, computed from Table 13 and the reported throughput, with 312 TFLOPs as the bf16 peak of an A100 and throughput summed over the devices of a run. The SOLO row uses its own FLOPs per token, so the auxiliary heads count as work rather than as overhead. The 96-block configuration used for the pipeline comparison runs at about a third of the utilization of the pretraining runs, because its blocks are narrow and its module is not compiled; its absolute throughput is therefore not representative and only ratios within a row of Table 15 should be read.
p=2
p=4
p=6
p=8
method
M=24
72
24
72
24
72
24
72
BP, 1F1B
1.000
1.000
1.000
1.000
1.000
1.000
1.000
1.000
BP, 1F1B (PyTorch)
–
–
–
–
–
–
0.980
0.978
BP, VPP-2
–
–
1.017
0.982
–
–
1.067
0.990
BP, VPP-4
–
–
1.040
0.956
–
–
1.062
0.859
BP, 1F1B + recomputation
0.766
0.768
–
–
0.750
0.745
0.748
0.736
Appendix
Table 15: Throughput under the same split, relative to 1F1B at the same p and M . Values above one mean the method is faster than 1F1B. Configuration (A): L=96 , d=1024 , T=1024 , V=8192 , H=1 , micro-batch 4 , K=p . A dash means the configuration was not run at that p . The data-parallel runs use the same p devices as full replicas with the same global batch, and are included as context rather than as a same-split comparison.
method
peak GB, max
peak GB, mean
act. GB, max
act. GB, mean
traffic GiB/step
BP, 1F1B
19.8
12.1
17.4
9.8
2.6
BP, 1F1B (PyTorch)
20.1
13.2
17.3
9.7
5.3
BP, VPP-2
28.8
21.6
24.8
17.0
11.3
BP, VPP-4
27.0
–
20.9
–
23.3
BP, 1F1B + recomputation
4.5
3.6
2.1
1.3
2.6
SOLO, S=50
5.1
5.0
2.5
2.5
1.3
Appendix
Table 16: Memory and communication at p=8 and M=24 , configuration (A). Peak memory is given for the most loaded device and as the mean over the eight devices. Activations are the peak minus resident state minus gradients. Traffic is computed from tensor shapes; the pipeline runs scale linearly with M and the data-parallel runs do not. The VPP and PyTorch 1F1B runs send activations in fp32, which doubles their traffic relative to an implementation that sends bf16.
Figure 13: Same-split throughput and the per-module cost. (a) SOLO relative to 1F1B against micro-batches per step at p=8 on the 96-block model of configuration (A); curves are the cost model with ρ fitted at M=72 , dots are measured, the dashed curve adds the synchronization cost of a broadcast at every step, the gray line is 1/ρ . (b) The same ratio against the number of stages at M=24 (solid) and 72 (dashed), for the 96-block model with H=1 and for the 24-block model of the pretraining runs ( d=2048 , T=2048 , V=32 k, micro-batch 1) with H=2 ; the open marker is the 24-block model with H=1 at p=4 . (c) Cost per module (ρ−1)/(K−1) against auxiliary-head depth under data parallelism on the 24-block model (micro-batch 8, M=24 ), for K=2 , 4, 8; solid line the fit, dotted line the FLOP model of ( 6 ), open squares the values recovered from pipeline throughput at p=4 .
BP, 1F1B
SOLO
B
M
tok/s
peak GB
tok/s, S=50
tok/s, S=1
peak GB
1
96
21,980
8.4
20,272
20,994
3.5
2
48
43,125
12.3
40,784
41,218
4.0
4
24
50,576
19.8
59,333
55,537
5.1
8
12
48,042
34.7
68,545
60,361
7.3
16
6
37,943
48.5
72,877
57,387
11.6
Appendix
Table 17: Micro-batch sweep at a fixed global batch of 96 sequences per step, configuration (A), p=8 . Peak memory on the most loaded device; SOLO with the copy refreshed every 50 steps and every step.
Figure 14: Micro-batch sweep of Table 17 . (a) Throughput against peak memory per GPU, points labeled by micro-batch size, circles at each method’s best. (b) Throughput against micro-batch size; the dotted line is SOLO’s throughput times the 1F1B bubble factor M/(M+p−1) .
link
nominal MB/s
1F1B
SOLO
SOLO / 1F1B
NVLink
n/a
57,685
59,606
1.033
socket
n/a
55,679
59,363
1.066
10 Gb/s
1250
52,667
59,318
1.126
5 Gb/s
625
48,401
59,328
1.226
2 Gb/s
250
38,989
59,213
1.519
1 Gb/s
125
28,084
58,901
2.097
Appendix
Table 18: Throughput for different link speeds, configuration (B): L=24 , H=2 , M=8 , p=2 , each method using the fastest of the layer splits 11:13 , 12:12 and 13:11 . This configuration differs from Table 15 and the two should not be compared directly. The rows from 10 Gb/s down shape the loopback link of a single node with tc tbf . At the two slowest settings the measured transfer rate is 88 to 92% of the nominal rate.
Figure 15: Same-split comparison, configuration (A). (a, b) Throughput relative to 1F1B against the number of stages at M=24 and M=72 . (c) The bubble model of ( 7 ) against the number of micro-batches, with ρ fitted at M=72 and the M=24 points held out; the dashed line adds the synchronization cost at S=1 . (d, e) Throughput of each schedule at p=4 and p=8 . (f) Peak memory on the most loaded device, split into resident state and activations. (g) Peak memory at p=8 , most loaded device against the mean over devices. (h) Traffic per step. (i) Throughput lost to readout synchronization against (p−1)/(MS) , with the fitted slope.
parameters
L
d
T
V
M⋆ , H=1
M⋆ , H=2
7B
32
4096
4096
32,000
23.7
13.7
8B
32
4096
8192
128,256
14.7
10.3
70B
80
8192
4096
32,000
66.9
36.5
70B
80
8192
8192
128,256
46.5
29.5
Appendix
Table 19: Crossover M⋆ from ( 7 ) for common model shapes. SOLO is faster than 1F1B when the number of micro-batches per step is below M⋆ . The value does not depend on the number of stages. Computed, not measured.
We explore catastrophic forgetting in the context of large pre-trained models. By considering forgetting as a geometric problem in the input space of each weight matrix, we uncover a natural retention objective under which updates produced by gradient-based optimizers are suboptimal. Following this observation, we propose Local Support Learning (LSL), a general-purpose framework that augments gradient-based training for retention of prior capabilities without access to prior data. During a new learning phase, LSL pairs two components with distinct roles: a standard weight adapter, trained as usual to minimize the loss, and a gating function that enables the adapter only on input activations from its own training distribution, making the update local to that distribution. The key challenge is that this gate must route data from all learning phases while training only on data from the current one. We address this with a gate based on a Gaussian Mixture Model (GMM), whose likelihood decays rapidly away from its training data, giving it a natural tendency to stay closed on data from prior phases. We show that this post-training approach can resolve forgetting in LLMs of up to 7 billion parameters, retaining both pretrained and finetuned capabilities across multiple training phases, while being efficient in memory and compute, robust to hyperparameter choice, and showing scaling potential.
LLM post-training typically propagates task gradients through the full depth of the model. Although this end-to-end structure is simple and general, it couples task adaptation to full-depth activation storage, long-range backward dependencies and direct task-gradient access to pretrained representations. We argue that this full-depth backward coupling can be unnecessarily expensive and intrusive, particularly when post-training supervision is much narrower than pre-training. To this end, we propose \textbf{LoPT}: Local-Learning Post-Training, a simple post-training strategy that makes gradient reach an explicit design choice. LoPT places a single gradient boundary at the transformer midpoint: the second-half block learns from the task objective, while the first-half block is updated by a lightweight feature-reconstruction objective to preserve useful representations and maintain interface compatibility. LoPT shortens the task-induced backward path while limiting direct interference from narrow task gradients on early-layer representations. Extensive experiments demonstrate that LoPT achieves competitive performance with lower memory cost, higher training efficiency and better retention of pretrained capabilities. Our code is available at: https://github.com/HumyuShi/LoPT
Hengyu Shi, Tianyang Han, Peizhe Wang +3
1Independent Researcher · 2D4 Lab · 3Southeast University
Training large language models is generally done on clusters containing thousands of accelerators, communicating over a high-bandwidth interconnect. Scaling up these clusters is expensive and can become impractical, imposing limits on the size of models that can be trained. Several recent studies have proposed training methods that are less communication intensive, avoiding the need for compute clusters with extremely high interconnect speeds. These low communication training methods still employ a global synchronization step for model parameters, which can be too costly with a high number of participants, as the communication cost scales quadratically with group size. In this work, we propose a novel optimization method, NoLoCo, that does not explicitly synchronize all model parameters during training and does not require any collective communication. NoLoCo implicitly synchronizes model weights via a novel variant of the Nesterov momentum optimizer by partially averaging model weights within randomly selected subgroups. We provide both a theoretical convergence analysis of our optimizer and empirical results from language model training. Our method requires significantly less communication than fully sharded data parallel training and DiLoCo, a widely used low-communication baseline. Moreover, our method avoids global blocking communication, thereby reducing accelerator idle time. Our experiments show that NoLoCo is more communication-efficient than DiLoCo, improving final perplexity by up to 4% and converging up to 4× faster in wall-clock time across a range of worker counts, model sizes, and communication bandwidths.