We study the problem of estimating the guidance that steers the distribution learned by a diffusion generative model toward a tilted target q0∝wp0 at inference time. Relying on the stochastic optimal control approach, we observe that the exact drift correction is the gradient of the logarithm of Doob's h-function, and we study the problem of estimating it from a sample. In the present paper, we assume that the score of the pretrained model is available, that the tilting weight is bounded and positive, and that the reference distribution has a bounded support, no smoothness of the weight is required. Introducing a penalized least-squares risk in which the penalty is the residual of the space-time harmonicity equation satisfied by the h-function, measured in a dual Sobolev norm, we derive high-probability bounds on the squared error of the resulting guidance estimate. Since the penalty vanishes at the target, the estimator is free of regularization bias, and in favourable scenarios its rate of convergence is faster than the minimax rate of estimating first-order derivatives of a smooth regression function. Assuming that w is bounded and positive with Ep0[w−s]<∞ for some s∈(0,∞], and that the reference data are compactly supported, we prove that the guidance is estimable in squared L2 at rate εns/(s+4), where εn=n−2(β−1)/(2(β−1)+d). We also transfer the obtained bounds to the total variation distance between the marginals of the estimated and the exactly guided samplers, and illustrate the performance of the suggested approach with numerical experiments.
Figures & tables
Penalty P(h)
order in h
value at h∗
approximation floor
Sobolev, ∫I∥∇h∥ρt2ν(dt)
1
∥∇h∗∥2=0
— (nonzero gap for λ>0 )
strong residual, ∥R[h]∥L2(ν⊗ρ)2
2
0
N−2(β−2) , non-informative at β=2
weak residual, ∥R[h]∥V∗2
1
0
N−2(β−1)
Table 1: Why the residual is measured weakly. The floor is infh∈HNP(h) at resolution N, only the last is both gap-free and informative at β=2 .
Figure 1 : (a) Guidance error at d=1 , ±1 s.d. over 8 seeds, against slopes 0.857 and 0.667 . (b) Held-out reward shift, as a share of the exact tilt’s, against the FID it costs over that dataset’s base ( 24.5 , 7.4 ), so grey is both base models and lower is cheaper. Colour is the method, line style the dataset, c∈[0.25,20] , bars ±1 s.d. over 3 seeds, DOIT is off the panel at 4.8 and 21 times its base, absolute values in Figure 3 . (c) The five largest rises in held-out probability among 64 samples, base above and ours at c=20 below on the same initial noise, all 64 in Figure 4 .
Appendix figures & tables12 assets
Supplementary material from the paper’s appendix.
Appendix
checkpoint
parameters
share
SIR-32
FID
CIFAR-10
google/ddpm-cifar10-32
35.7 M
0.653
0.854
24.5
MNIST
dvgodoy/ddpm-cifar10-32-mnist
35.7 M
0.523
0.630
7.4
Appendix
Table 2 : Base models. “Share” is the fraction of base samples the held-out head assigns to the target superclass, “SIR-32” the same after 32 -fold importance resampling with weights w , and FID is to the model’s own reference sample (Appendix J.3 ).
penalty
error at n=103
error at n=106
predicted exponent
value LS ( μ=0 )
6.46×10−2
9.71×10−5
0.667
Sobolev
4.41×10−2
1.40×10−4
—
strong residual
4.46×10−3
2.34×10−5
—
weak residual (ours)
6.51×10−3
1.64×10−5
0.857
Appendix
Table 3 : Penalty ablation at d=1 in the cubic class ( β=4 ), 8 seeds, each penalty at its own oracle (N,μ) : relative guidance error at the smallest and the largest sample size, against the exponent the theory predicts for it. A rate is predicted only for the weak residual and, from the classical derivative-recovery bound, for least squares, the Sobolev penalty carries a regularization bias and the strong residual an approximation floor (Table 1 ), and neither is analysed here. “Value LS” is the μ=0 baseline of Figure 2 .
Figure 2 : Guidance error against n at d=1 in the tensor cubic B-spline class ( β=4 ), for the weak-residual estimator and for least squares on the values followed by differentiation, against the predicted slopes 0.857 and 0.667 , intercepts fitted only. Bars are ±1 standard deviation over 8 seeds. The weak residual’s fitted exponent, 0.842±0.026 , matches its prediction, least squares runs above its own, 0.909±0.038 against 0.667 , so over this range it is measured far from its asymptote and its steeper slope is a climb out of a worse level rather than a faster rate.
weak residual (ours)
differentiate
ratio
d
n=103
largest n
n=103
largest n
smallest
largest
1
6.5×10−3
1.6×10−5
6.5×10−2
9.7×10−5
5.9×
13.5×
2
2.3×10−2
6.9×10−5
3.5×10−1
4.2×10−4
5.4×
15.3×
3
8.3×10−2
3.7×10−4
2.5
3.6×10−3
9.7×
30.2×
Appendix
Table 4 : The two estimators at equal budget, 8 seeds, each at its own oracle (N,μ) on the same grids. “Differentiate” is least squares on the values, then differentiated. Entries are relative squared guidance error, the largest n is 106 , and 3.2⋅105 at d=3 .
ours
DEFT
DOIT
c
%
FID
%
FID
%
FID
0.01†
—
—
—
—
−43.8±5.1
43.1±0.7
0.05†
—
—
—
—
38.7±2.0
87.0±1.2
0.1†
—
—
—
—
116.0±2.5
221.6±1.0
0.25
3.2±0.3
24.5±0.2
—
—
—
—
0.5
6.1±0.4
24.6±0.2
—
—
—
—
Appendix
Table 5 : CIFAR-10, animals, along the guidance scale ( γ for DOIT): held-out reward shift as % of the exact tilt’s, and FID to the model’s own reference sample. Base FID 24.5±0.2 , share 0.653 . Bold marks, at each scale, the largest shift and the smallest FID, a row holds c fixed, which leaves the three methods at different FIDs, so the operative reading is the frontier of Figure 3 and not the row. DOIT is entered at the three γ of the shared sweep, where its cap binds on every guided step so that the three agree to 0.2 points, and at the small γ its authors’ own setting calls for. Daggered rows are that setting — η=0.4 and no cap, read against the base of the same sampler — and are not comparable with the rest of the row they sit in, the wider grid over γ , τ and the sampler is in Appendix J.8 . Here and in Table 6 , ± is one standard deviation over the three seeds, which dominates the standard error over the 6000 paired samples (at most 2.3 for the shift).
Figure 3 : What the alignment costs. Horizontal is the held-out reward shift as a percentage of the exact tilt’s, vertical the FID to the model’s own reference sample, so the grey point is the base model and a lower curve is cheaper, bars are ±1 s.d. over three seeds on both axes. Left: CIFAR-10, c∈[0.25,20] . Right: MNIST, c∈{1,5,20} . DOIT is off both panels, at FID 117 and 153 against bases of 24.5 and 7.4, drawing it would compress the region the figure exists to show, and its numbers are in Tables 5 and 6 .
Figure 4 : CIFAR-10, animals, the first 64 samples of seed 0 , uncurated. Base and ours at c=20 above, DEFT and SIR-32 below, each panel labelled with the share of animals it carries and with its FID. DEFT is drawn at c=2 , the scale of its sweep whose FID is closest to the one our own panel is drawn at, because a matched departure from the base distribution is the reading under which the two frontiers are compared, at a matched c it would instead sit at FID 98 . The guided panels share the initial noise with the base panel, so a sample can be compared with its counterpart in the same position, SIR-32 does not, by construction.
ours
DEFT
DOIT
c
%
FID
%
FID
%
FID
0.01†
—
—
—
—
88.5±3.7
14.1±0.8
0.05†
—
—
—
—
280.0±7.9
99.4±2.2
0.1†
—
—
—
—
283.4±7.9
108.0±1.4
1
13.9±1.0
7.4±0.1
59.8±3.5
9.3±0.3
296.4±4.1
153.3±3.6
5
62.8±4.7
7.7±0.1
145.6±4.6
25.2±0.2
296.4±4.0
153.1±3.5
Appendix
Table 6 : MNIST, odd digits, tilted under the recipe of Appendix J.1 : held-out reward shift as % of the exact tilt’s and FID to the model’s own reference sample. Base share 0.523 , FID 7.4±0.1, SIR-32 share 0.630 . Bold marks, at each scale, the largest shift and the smallest FID, and as in Table 5 the row, which holds c fixed, is not the operative comparison. Daggered rows are DOIT at the small γ and in the setting its authors use, η=0.4 and no cap, read against the base of that sampler, γ=0.01 is the one point at which DOIT reaches a usable frontier at all. ± is one standard deviation over the three seeds.
Figure 5 : MNIST, odd digits, the first 64 samples of seed 0 , uncurated. Base and ours at c=20 above, DEFT at c=1 and SIR-32 below, each panel labelled with the share of odd digits and with its FID, as in Figure 4 , DEFT is drawn at the scale of its sweep whose FID is closest to ours. The guided panels share the initial noise with the base panel. The base model draws some digits mirrored.
sampler
method
∣δ∣
binds
capped → uncapped
η=0
ours
0.031
5.3%
0.831/32.9→0.739/43.4
DEFT
0.158
77.9%
0.941/97.7→ —
DOIT
2.4
100%
0.689/116.7→0.995/463.2
η=0.4
ours
0.031
5.0%
0.854/34.6→0.805/46.2
DOIT, γ=1
2.3
100%
0.464/70.2→0.996/437.1
DOIT, γ=0.1
0.23
95.2%
0.462/70.0→0.889/221.6
Appendix
Table 7 : The cap at c=20 on CIFAR-10 ( γ=1 for DOIT, and γ=0.1 , the small- γ setting of its authors), where it binds at all: the mean uncapped correction against the cap of 0.1 , the share of guided steps it binds on, and what removing it does. At c=5 and below it binds on at most 1% of the steps of either trained method. The uncapped point was not sampled for the reported DEFT head, nor was that head run under the stochastic sampler, so those cells are empty.
β (order)
weak floor
strong floor
strong / weak at N=24
4 ( 3 )
N−1.44
N−1.89
7.4
3 ( 2 )
N−1.81
N−1.00
311
2 ( 1 )
N−2.46
N+0.29
8500
Appendix
Table 8 : Approximation floors against resolution, d=1 , 4 replicates, slopes over N∈[2,24] , predictions −2(β−1) and −2(β−2) .
minw
n=103
n=106
measured
envelope
clip
0.5 (base)
3.0×10−3
1.0×10−5
0.787
0.779
0.00
1
1.2×10−2
2.2×10−5
0.919
0.779
0.00
0.3
2.6×10−2
2.0×10−5
1.014
0.779
0.00
0.1
4.3×10−2
2.2×10−5
1.076
0.337
0.00
0.03
1.5×10−1
2.9×10−3
0.547
0.004
0.40
3×10−3
2.6×10−1
2.1×10−1
0.028
0.002
0.44
Appendix
Table 9 : The finite- s ladder, d=1 , 4 replicates, each cell at its own oracle (N,μ), slopes over the seven sample sizes, “clip” the mass on which h∗<b at n=106 .