Lecture 04 · LAB 09 · 学得去噪器与完整采样

LAB 09 · 学得去噪器与完整采样

先阅读保存结果和解释,再按本册步骤选择是否运行。在 Deepnote 阅读与运行 · 下载 Notebook

LAB09|把解析去噪换成学得的网络

先看完整链条改变了什么

本册把 LAB05 的“已知解析预测器”换成真正训练的小型时间条件网络。网络从许多带噪练习题中学习;生成时从一个新噪声状态出发,重复调用它。训练题的正确噪声、生成时的当前状态、评价时使用的解析参照,在代码里是三个分开的角色。

同一批2048个起点下,初始化、2000次与8000次更新后的采样,与同条件解析参照比较。上排20步,下排100步;灰点是真实分布,橙点是网络采样,青点是解析预测采样。

先固定下排100步横向比较:初始化只能形成一团散点,训练2000步后出现环状结构,8000步后八个模式更清楚。再固定8000步纵向比较:20步仍把较多点留在簇间,100步更集中。最后看最右列:即使预测器是解析的,有限步采样仍会改变终点分布。因此“预测够不够准”和“怎样把预测接成生成”必须分别检查。

图中初始化约1.46%(20步)/3.71%(100步)的点在坐标范围之外,结果数组与指标均保留这些点;其余训练态没有出框点。绘图没有删除离群样本后重算分数。

同一个八模数据,为什么能自己出题

数据与LAB08相同:均匀选八个半径2的中心之一,再加标准差0.12的二维高斯噪声。每次训练重新抽一个x₀、一个时间t和一个标准正态噪声ε,然后构造:

\[ x_t=\alpha_t x_0+\sigma_t\epsilon,\qquad \alpha_t=e^{-5t},\quad\sigma_t=\sqrt{1-e^{-10t}}. \]

因为ε是程序刚抽出来的,训练标签无需人工标注。网络只接收xₜ与t,输出对ε的预测;x₀与真实ε用于制造练习题和计算损失,不作为网络额外输入。时间告诉网络这次题大致损坏到什么程度。

时间按 t=0.001+0.999u²、u∼Uniform(0,1) 抽样,让固定预算多见一些低噪声题;损失是在这个时间分布下对样本与两个坐标一起求平均的噪声MSE。这里没有声称采用均匀时间的原始DDPM损失权重。

x0 = draw_data(rng, batch)
t = .001 + .999 * rng.uniform(size=batch)**2
epsilon = rng.normal(size=x0.shape)
a, s = schedule(t)
xt = a*x0 + s*epsilon
prediction, _ = predict(model, xt, t)
loss = np.mean((prediction - epsilon)**2)

网络到底学了什么,初始化已经知道什么

输入由两个坐标、时间和四组时间正弦/余弦特征组成,共11维;后接64→64→2的tanh MLP。为避免把大量预算用在高噪声下近似恒等映射,预测器含一个固定线性支路:

\[ \widehat\epsilon_\theta(x,t) =\underbrace{\frac{\sigma_t}{\alpha_t^2v+\sigma_t^2}x}_{\text{同均值、同方差的高斯基线}} +\alpha_t f_\theta(x,t),\qquad v=2+0.12^2=2.0144. \]

初始化不是完全无知识的零网络。 线性支路知道数据均值为0、每轴方差为2.0144,这些量取自本实验数据定义;它不知道八个模式的方向与位置分配。MLP的末层小随机初始化;真正被训练的是MLP所有共享权重与偏置。学习要补出高斯基线不能表达的多峰结构。代码和结果把这个初始化完整保留,不把预置部分当作网络学到的能力。

网络从抽样ε的MSE反传。没有使用解析函数来给学习器蒸馏标签;解析预测只出现在评价和单独的参照采样分支。

为什么“正确噪声”也不是唯一可推回的答案

同一个附近的xₜ可能来自不同x₀和不同ε。只给xₜ、t,网络无法知道本题历史上究竟抽中了哪一对。平方误差下最佳预测是条件均值 E[ε|xₜ,t];即使精确给出这个条件均值,针对逐题真实ε的MSE也一般不为0。

因此本册同时记录两种误差:对抽样ε的误差,回答它在真实练习题上预测得怎样;对解析条件均值的误差,回答网络离这个任务的最佳均方预测还有多远。4096个固定未见测试题从独立随机流生成,训练从未读取它们。

参数状态 对抽样ε的MSE 对解析条件均值的MSE
初始化 0.34665 0.11890
更新2000次 0.27845 0.05058
更新8000次 0.24462 0.01875
解析条件均值 0.22569 0

低噪声t<0.05处,最终网络对解析均值的误差为0.06720;0.5≤t≤1处为0.000649。整体平均下降掩盖了不同时间段难度,也提醒我们:采样末段恰恰可能遇到网络仍不精确的位置。

每100次更新记录的训练批次噪声MSE。虚线是固定未见样本上解析条件均值对逐题噪声的MSE;它们的样本集不同,不能把瞬时曲线穿过虚线理解成打败最优预测器。

从当前预测走向下一状态

本册用η=0的确定性DDIM形式:在时间t预测ε,先得到当前对干净样本的估计,再合成较小时间s的状态:

\[ \widehat x_0=\frac{x_t-\sigma_t\widehat\epsilon(x_t,t)}{\alpha_t},\qquad x_s=\alpha_s\widehat x_0+\sigma_s\widehat\epsilon(x_t,t). \]

这里α表示信号幅度,等于DDIM论文中累计噪声系数的平方根;不能把两个记号直接替换而漏掉平方根。时间网格从1均匀降到0.001,随后用最后一次干净估计到0。20步或100步指预测器调用次数,分别用20或100个正时间点。

for t, next_t in pairs(times):
    eps_hat = epsilon_fn(current_x, t)
    x0_hat = (current_x - sigma(t)*eps_hat) / alpha(t)
    current_x = alpha(next_t)*x0_hat + sigma(next_t)*eps_hat

采样器只取得当前坐标、时间、预测器和预定日程。 它没有目标x₀,没有逐题真实ε,也没有“离哪个模式最近”的标签。在学得分支中,epsilon_fn只调用已训练网络;解析分支明确使用另一函数,不能把它偷偷混入学得轨迹。

所有图使用同一组从N(0,I)抽出的2048个xT。t=1时真实前向边缘与标准正态已经很近,但并非严格相等;这一近似对两个分支相同。有限网格、最后从0.001到0的处理和起点近似都保留,不能宣称解析参照就等于无误差的数据生成器。

固定起点,只改变采样步数

固定前12个随机起点的逐步轨迹。空心圆为相同xT,黑叉为终点;上排20步,下排100步,左列学得预测,右列解析预测。每条线的所有中间状态均来自采样器实际更新。

选一条轨迹读:从空心圆出发,每次只知道此刻状态与时间;不是沿着早已指定的目标点移动。20步的折线跨度较大,100步有更多中间修正。学得与解析分支有时走向不同终点,这可以来自预测近似误差,也可能在模式分界附近被步进放大。

沿用LAB08的覆盖规则:半径0.36邻域、至少总样本1%的点才算覆盖一个模式。覆盖与近模式比例同时保留;距离更近也不能代替模式内部方差与占用均衡的检查。

预测器/训练状态 采样步数 覆盖 近模式比例 最近中心距离中位数
初始化 20 4/8 8.69% 0.8500
初始化 100 6/8 9.38% 0.8466
2000次更新 20 8/8 61.23% 0.2839
2000次更新 100 8/8 62.45% 0.2877
8000次更新 20 8/8 58.89% 0.3006
8000次更新 100 8/8 80.03% 0.1654
同条件解析参照 20 8/8 70.95% 0.2604
同条件解析参照 100 8/8 94.68% 0.1027
独立真实样本 — 8/8 98.93% 0.1450

在8000次更新后,100步比20步有更多点落在模式附近。但2000→8000次训练,20步的近模式比例反而由61.23%降到58.89%,尽管测试噪声MSE下降。这直接回答“训练损失更低,生成就一定更好吗”:单个预测器的改善和某个离散采样器的终点分布不是同一指标。

解析100步的最近中心距离甚至比真实样本更小,说明“更集中”并非完整质量目标。若只用这个距离排名,会奖励把真实模式压得过窄。我们用这些指标观察误差来源,不用教学小模型给GAN和扩散家族排名。

解析参照与 LAB05 怎样接上

参照与学习器严格使用本册同一个八模数据、α/σ参数化、时间网格、初始xT和DDIM更新。解析式先计算给定xₜ、t时的八个分量后验权重,再求 E[x₀|xₜ,t],最后用 (xₜ−αE[x₀|xₜ,t])/σ 得到条件平均噪声。

它不是把采样终点吸附到最近中心,也不是知道某个生成起点对应的真答案。它只因本课数据分布被明确给定,才可以精确计算条件平均预测。LAB05帮助理解已知密度下如何核对预测与积分;本册重新提供匹配八模与本日程的解析函数。LAB05旧图若使用其他分布、速度参数化或更新规则,不能直接拿数值与本册横比。

学生常问:解析已经能算,为何还要学? 在这个小世界里,它就是一个有价值的参照;真实图像的完整数据密度并没有这样的已知公式。网络从样本获得可反复调用的局部预测,正是为了在无法手写精确密度时构造生成过程。

运行、改变量与留下证据

默认训练8000步、批量256、Adam学习率0.0008,单次固定训练约6.69秒;包括两种步数和三个检查点的采样与绘图。环境为Python 3.12.14、NumPy 2.3.5、单线程BLAS;设备不同会改变时间。训练前、训练中、训练后都用同测试题与同xT,不依赖其他Notebook状态。完整脚本只需NumPy和Matplotlib。

第一次运行保持全部默认值,先定位保存的test_xt、test_t、test_epsilon及各检查点预测,再对应上表。默认 run() 会重新训练并生成20/100步对照;如果只想改采样步数,不要再次调用 run()。先运行定义类与函数的单元,跳过重新训练单元,再使用文末独立的“只重采样”单元,从已保存的权重与xT读取。文件可以来自随包保存结果,或本册第一次运行生成的 outputs/;若缺文件,先完成一次默认运行。

下面是该单元的核心步骤,完整单元还保存新结果并检查权重没有改变:

saved = np.load('outputs/lab09-data.npz')
weights = np.load('outputs/lab09-weights.npz')
loaded_model = MLP(CONFIG['architecture'], CONFIG['seed_model'])
for index, parameter in enumerate(loaded_model.p):
    parameter[...] = weights[f'p{index}']
fixed_xT = saved['fixed_xT'].copy()
samples50, trajectory50, times50 = ddim(
    lambda x, t: predict(loaded_model, x, t)[0], fixed_xT, 50
)
print(coverage(samples50))

这一只重采样路径已经实际运行:训练更新0次,50次预测调用,约0.106秒,参数逐位未改变,起点与保存的xT完全一致。50步结果覆盖8/8、近模式比例76.07%、最近中心距离中位数0.1656。它位于已保存20步58.89%和100步80.03%之间,但单个距离指标并不需要严格随步数单调变化。

这个干预只改变生成阶段的采样步数。它不同于LAB07改变训练目标中的β,也不同于LAB08改变训练过程的D:G更新次数;后二者必须重新训练,不能只加载同一份最终权重就声称完成对照。先提出一个关于轨迹或模式占用的预测,再读取结果;不要同时改模型、训练时间分布和步数。

自检。 为什么不把ε的MSE降到0作为验收?因为给定xₜ、t仍有条件不确定性。为什么解析参照也要改变步数?这样才能观察同一预测器下离散更新的影响。为什么不从学习器里调用解析式修正末态?那会改变被检验的生成方法,不能再把结果全部归给网络学习。

来源:Ho等,Denoising Diffusion Probabilistic Models,关注前向造题与噪声预测;Song、Meng、Ermon,Denoising Diffusion Implicit Models,重点看式12与η=0的更新。本册的连续指数日程、固定高斯线性支路、二维数据和有限训练预算是明确的教学实现选择。

重算时,run() 用 threadpoolctl 将这一次实验的 BLAS 计算限为单线程,并在结果 JSON 保存实际线程信息。Notebook 内核可能已加载 NumPy,只在代码中设置环境变量不能可靠改变已启动的线程池。这个执行控制不改变数据、随机种子、训练预算或采样规则;运行时间仍按当前机器实测,不能将本地秒数当作云端承诺。

云端复核(2026-09-30)。 本册在 Deepnote 的独立内核按自身步骤完整运行。完成8000次训练更新;固定最终权重与2048份起点,20/100步的近模式比例分别为58.89%/80.03%,两者均覆盖8/8。另用相同权重与起点重算50步,近模式比例76.07%,该格没有训练更新。上方原始保存图继续标为本地基线,两次测量分别记录;大图在插件的运行快照预览中可能因尺寸被省略,这不代替学生在同册查看自己的输出。

完整可运行实现

下面代码与下载版 .py 同源。先读上面的已保存实测结果;从新内核按顺序运行以下代码,可重新训练并保存自己的图、权重、数组与指标。这里显示的既有图来自独立 Python 进程,不冒称已在当前 Deepnote 计算实例中运行。

查看可执行代码
# Dependencies: Python 3.10+, numpy, matplotlib, threadpoolctl. No downloads or hidden notebook state.
import os
os.environ.setdefault('OPENBLAS_NUM_THREADS','1')
os.environ.setdefault('OMP_NUM_THREADS','1')
import json, time, math, hashlib, platform
from pathlib import Path
import numpy as np
from threadpoolctl import threadpool_limits, threadpool_info
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt

CENTERS = 2.0*np.stack([np.cos(np.arange(8)*2*np.pi/8),np.sin(np.arange(8)*2*np.pi/8)],axis=1)
DATA_STD = .12

def draw_data(rng,n):
    k=rng.integers(0,8,n)
    return CENTERS[k]+DATA_STD*rng.normal(size=(n,2))

def coverage(x):
    dist=np.linalg.norm(x[:,None,:]-CENTERS[None,:,:],axis=2)
    nearest=dist.argmin(1); inside=dist.min(1)<=3*DATA_STD
    counts=np.bincount(nearest[inside],minlength=8)
    return {'n':len(x),'covered_modes':int(np.sum(counts>=.01*len(x))),
            'mode_counts_within_3sigma':counts.tolist(),'on_mode_fraction':float(inside.mean()),
            'median_distance_to_center':float(np.median(dist.min(1))),
            'off_plot_fraction':float(np.mean(np.any(abs(x)>3.1,axis=1)))}

class MLP:
    def __init__(self,sizes,seed,final_scale=1.):
        rng=np.random.default_rng(seed); self.p=[]
        for k,(a,b) in enumerate(zip(sizes[:-1],sizes[1:])):
            scale=final_scale if k==len(sizes)-2 else 1.
            self.p.extend([rng.normal(size=(a,b))*np.sqrt(2/(a+b))*scale,np.zeros(b)])
        self.m=[np.zeros_like(p) for p in self.p]; self.v=[np.zeros_like(p) for p in self.p]; self.step=0
    def forward(self,x):
        acts=[x]
        for k in range(0,len(self.p),2):
            x=x@self.p[k]+self.p[k+1]
            if k<len(self.p)-2: x=np.tanh(x)
            acts.append(x)
        return x,acts
    def backward(self,acts,grad):
        grads=[None]*len(self.p)
        for layer in range(len(acts)-2,-1,-1):
            if layer<len(acts)-2: grad=grad*(1-acts[layer+1]**2)
            grads[2*layer]=acts[layer].T@grad; grads[2*layer+1]=grad.sum(0)
            grad=grad@self.p[2*layer].T
        return grad,grads
    def update(self,grads,lr,b1=.9,b2=.999):
        self.step+=1
        for k,(p,g) in enumerate(zip(self.p,grads)):
            self.m[k]=b1*self.m[k]+(1-b1)*g; self.v[k]=b2*self.v[k]+(1-b2)*g*g
            p-=lr*(self.m[k]/(1-b1**self.step))/(np.sqrt(self.v[k]/(1-b2**self.step))+1e-8)
    def flat(self): return np.concatenate([p.ravel() for p in self.p])
    def state(self): return {f'p{i}':p.copy() for i,p in enumerate(self.p)}

def sigmoid(x): return 1/(1+np.exp(-np.clip(x,-50,50)))
def bce_logits(logit,label): return np.mean(np.logaddexp(0,logit)-label*logit)
def dump_json(path,data): path.write_text(json.dumps(data,indent=2,ensure_ascii=False),encoding='utf8')
def check_gradient():
    m=MLP([2,5,2],980); x=np.random.default_rng(981).normal(size=(3,2)); target=np.ones((3,2))
    y,a=m.forward(x); _,g=m.backward(a,(y-target)/y.size)
    errors=[]
    for k,ij in [(0,(0,1)),(2,(1,0)),(3,(0,))]:
        old=m.p[k][ij]; h=1e-5
        m.p[k][ij]=old+h; plus=.5*np.mean((m.forward(x)[0]-target)**2)
        m.p[k][ij]=old-h; minus=.5*np.mean((m.forward(x)[0]-target)**2)
        m.p[k][ij]=old; errors.append(abs((plus-minus)/(2*h)-g[k][ij]))
    assert max(errors)<1e-7,errors
    return max(errors)

def scatter_base(ax,real):
    ax.scatter(real[:,0],real[:,1],s=4,c='#bcc8cf',alpha=.35,rasterized=True)
    ax.scatter(CENTERS[:,0],CENTERS[:,1],s=35,c='#e5ac33',marker='x',zorder=10)
    ax.set(xlim=(-3.1,3.1),ylim=(-3.1,3.1),aspect='equal',xlabel='x₁',ylabel='x₂')

def setup_style():
    plt.rcParams.update({'font.family':'DejaVu Sans','font.size':10,'axes.spines.top':False,'axes.spines.right':False,'figure.facecolor':'white','savefig.facecolor':'white'})
查看可执行代码
# Fixed, finite training budget. No analytic target is used by the learner.
CONFIG={'seed_model':800,'seed_training':801,'seed_evaluation':802,'steps':8000,'batch':256,
        'learning_rate':.0008,'architecture':[11,64,64,2],'checkpoints':[0,2000,8000],
        'sampling_steps':[20,100],'beta':10.,'t_min':.001,
        'time_training_distribution':'t=.001+.999*Uniform(0,1)^2; deliberately spends more training examples at low noise',
        'noise_parameterization':'x_t=alpha*x0+sigma*epsilon; alpha=exp(-5t); sigma=sqrt(1-exp(-10t))',
        'predictor':'Gaussian-covariance linear skip plus alpha(t)*MLP(x_t,time_features)',
        'data':'8 equal Gaussian modes; radius=2; sigma=.12',
        'coverage':'At least 1% of all 2048 generated points within radius .36 of a center',
        'sampler':'deterministic DDIM eta=0; uniform t grid from 1 to .001, then final x0 estimate',
        'initialization_note':'Tiny random neural residual plus a fixed Gaussian second-moment baseline; not a zero predictor',
        'reference_note':'Exact conditional expected epsilon for the SAME mixture and SAME schedule; used only in reference branch/evaluation'}
VAR0=2.+DATA_STD**2

def schedule(t):
    t=np.asarray(t).reshape(-1,1); a=np.exp(-.5*CONFIG['beta']*t);s=np.sqrt(-np.expm1(-CONFIG['beta']*t));return a,s

def features(x,t):
    t=np.asarray(t).reshape(-1,1)
    return np.concatenate([x/np.sqrt(2),t,np.sin(2*np.pi*t*np.array([1,2,4,8])),np.cos(2*np.pi*t*np.array([1,2,4,8]))],axis=1)

def predict(model,x,t):
    a,s=schedule(t); correction,cache=model.forward(features(x,t))
    linear=s*x/(a*a*VAR0+s*s)
    return linear+a*correction,(cache,a)

def analytic_epsilon(x,t):
    # An independent oracle/reference, never used to form learned-network training labels.
    a,s=schedule(t); v=a*a*DATA_STD**2+s*s
    diff=x[:,None,:]-a[:,None,:]*CENTERS[None,:,:]
    logp=-np.sum(diff*diff,axis=2)/(2*v)
    weights=np.exp(logp-logp.max(1,keepdims=True)); weights/=weights.sum(1,keepdims=True)
    mean_center=weights@CENTERS
    # E[x0 | xt, component] = mu + alpha*sigma_data²/v*(xt-alpha*mu)
    posterior=mean_center+a*DATA_STD**2/v*(x-a*mean_center)
    return (x-a*posterior)/s

def ddim(epsilon_fn,initial,steps):
    # Legal inputs only: current x, scalar t, predictor, schedule. No target sample.
    x=initial.copy(); trajectory=[x[:24].copy()]
    times=np.r_[np.linspace(1.,CONFIG['t_min'],steps),0.]
    for t,snext in zip(times[:-1],times[1:]):
        ts=np.full(len(x),t); a,s=schedule(ts);eps=epsilon_fn(x,ts)
        x0hat=(x-s*eps)/a
        an,sn=schedule(np.full(len(x),snext)); x=an*x0hat+sn*eps
        trajectory.append(x[:24].copy())
    return x,np.stack(trajectory),times
查看可执行代码
# Training labels are sampled epsilon; all evaluation examples are held out.
@threadpool_limits.wrap(limits=1, user_api='blas')
def run(output_dir='outputs'):
    begin=time.time();out=Path(output_dir);out.mkdir(parents=True,exist_ok=True);setup_style();check_gradient()
    rng=np.random.default_rng(CONFIG['seed_training']); erng=np.random.default_rng(CONFIG['seed_evaluation'])
    real=draw_data(erng,2048);initial=erng.normal(size=(2048,2))
    test_x0=draw_data(erng,4096);test_t=CONFIG['t_min']+(1-CONFIG['t_min'])*erng.uniform(size=4096)**2
    test_eps=erng.normal(size=(4096,2));ta,ts=schedule(test_t);test_xt=ta*test_x0+ts*test_eps
    oracle_test=analytic_epsilon(test_xt,test_t)
    model=MLP(CONFIG['architecture'],CONFIG['seed_model'],final_scale=.01)
    saved={'real_evaluation':real,'centers':CENTERS,'fixed_xT':initial,'test_x0':test_x0,'test_xt':test_xt,'test_t':test_t,'test_epsilon':test_eps,'oracle_test_prediction':oracle_test}
    records={};curve=[];fig,axes=plt.subplots(len(CONFIG['sampling_steps']),4,figsize=(15,3.75*len(CONFIG['sampling_steps'])),constrained_layout=True,squeeze=False)
    cols={0:0,2000:1,8000:2}
    for step in range(CONFIG['steps']+1):
        if step in CONFIG['checkpoints']:
            prediction=predict(model,test_xt,test_t)[0]
            bytime=[]
            for lo,hi in [(0,.05),(.05,.2),(.2,.5),(.5,1.01)]:
                mask=(test_t>=lo)&(test_t<hi)
                bytime.append({'t_range':[lo,hi],'n':int(mask.sum()),'noise_target_MSE':float(np.mean((prediction[mask]-test_eps[mask])**2)),
                               'conditional_mean_error_MSE':float(np.mean((prediction[mask]-oracle_test[mask])**2))})
            record={'noise_target_MSE':float(np.mean((prediction-test_eps)**2)),
                    'conditional_mean_error_MSE':float(np.mean((prediction-oracle_test)**2)),
                    'by_time':bytime,'samplings':{}}
            saved[f'step{step}_heldout_prediction']=prediction
            for row,n in enumerate(CONFIG['sampling_steps']):
                result,traj,times=ddim(lambda x,t:predict(model,x,t)[0],initial,n)
                metrics=coverage(result);record['samplings'][str(n)]=metrics
                saved[f'learned_step{step}_n{n}_samples']=result;saved[f'learned_step{step}_n{n}_trajectory']=traj;saved[f'times_n{n}']=times
                ax=axes[row,cols[step]];scatter_base(ax,real);ax.scatter(result[:,0],result[:,1],c='#df713e',s=4,alpha=.4,rasterized=True)
                ax.set_title(f'Network updates {step} · {n} steps\ncoverage {metrics["covered_modes"]}/8 · near modes {metrics["on_mode_fraction"]:.1%}')
            records[str(step)]=record
        if step==CONFIG['steps']:break
        x0=draw_data(rng,CONFIG['batch']);t=CONFIG['t_min']+(1-CONFIG['t_min'])*rng.uniform(size=CONFIG['batch'])**2
        epsilon=rng.normal(size=x0.shape);a,s=schedule(t);xt=a*x0+s*epsilon
        pred,(cache,alpha)=predict(model,xt,t)
        # mean((pred-epsilon)^2) over batch AND coordinate; sampled epsilon is known.
        grad_eps=2*(pred-epsilon)/pred.size
        _,grads=model.backward(cache,grad_eps*alpha);model.update(grads,CONFIG['learning_rate'])
        if step%100==0:curve.append([step+1,float(np.mean((pred-epsilon)**2))])
    oracle_records={}
    for row,n in enumerate(CONFIG['sampling_steps']):
        result,traj,times=ddim(analytic_epsilon,initial,n);metrics=coverage(result);oracle_records[str(n)]=metrics
        saved[f'oracle_n{n}_samples']=result;saved[f'oracle_n{n}_trajectory']=traj
        ax=axes[row,3];scatter_base(ax,real);ax.scatter(result[:,0],result[:,1],c='#267e83',s=4,alpha=.5,rasterized=True)
        ax.set_title(f'Matched analytic reference · {n} steps\ncoverage {metrics["covered_modes"]}/8 · near modes {metrics["on_mode_fraction"]:.1%}')
    fig.suptitle('Same x_T in every panel · gray: real · orange: learned · teal: analytic reference',fontsize=13)
    fig.savefig(out/'lab09-training-and-steps.png',dpi=170);plt.close(fig)
    fig,axes=plt.subplots(len(CONFIG['sampling_steps']),2,figsize=(10,4*len(CONFIG['sampling_steps'])),constrained_layout=True,squeeze=False)
    for row,n in enumerate(CONFIG['sampling_steps']):
        for col,kind in enumerate(['learned_step8000','oracle']):
            traj=saved[f'{kind}_n{n}_trajectory'];ax=axes[row,col];scatter_base(ax,real)
            for j in range(12):ax.plot(traj[:,j,0],traj[:,j,1],lw=1.2,alpha=.85)
            ax.scatter(initial[:12,0],initial[:12,1],marker='o',facecolors='none',edgecolors='black',s=30,label='same xT')
            ax.scatter(traj[-1,:12,0],traj[-1,:12,1],c='black',marker='x',s=20,label='final')
            ax.set_title(f'{"Learned" if col==0 else "Analytic"} · {n} DDIM steps');ax.legend(fontsize=8)
    fig.savefig(out/'lab09-trajectories.png',dpi=170);plt.close(fig)
    curve=np.array(curve);saved['training_losses']=curve
    fig,ax=plt.subplots(figsize=(8,3.4),constrained_layout=True);ax.plot(curve[:,0],curve[:,1],lw=1,label='sampled minibatch noise MSE')
    ax.axhline(float(np.mean((oracle_test-test_eps)**2)),c='#267e83',ls='--',label='held-out analytic expected-noise MSE')
    ax.set(xlabel='Network updates',ylabel='Mean squared error');ax.legend();fig.savefig(out/'lab09-loss.png',dpi=170);plt.close(fig)
    report={'config':CONFIG,'checkpoints':records,'analytic_reference':oracle_records,
            'analytic_noise_target_MSE':float(np.mean((oracle_test-test_eps)**2)),
            'runtime_seconds':time.time()-begin,'python':platform.python_version(),'numpy':np.__version__,
            'blas_threadpools':threadpool_info(),
            'run_utc':time.strftime('%Y-%m-%dT%H:%M:%SZ',time.gmtime()),'real_reference_coverage':coverage(real)}
    np.savez_compressed(out/'lab09-weights.npz',**model.state());np.savez_compressed(out/'lab09-data.npz',**saved)
    dump_json(out/'lab09-metrics.json',report);print(json.dumps(report,ensure_ascii=False,indent=2));return report
查看可执行代码
report = run('outputs')
查看保存的计算输出
{
  "config": {
    "seed_model": 800,
    "seed_training": 801,
    "seed_evaluation": 802,
    "steps": 8000,
    "batch": 256,
    "learning_rate": 0.0008,
    "architecture": [
      11,
      64,
      64,
      2
    ],
    "checkpoints": [
      0,
      2000,
      8000
    ],
    "sampling_steps": [
      20,
      100
    ],
    "beta": 10.0,
    "t_min": 0.001,
    "time_training_distribution": "t=.001+.999*Uniform(0,1)^2; deliberately spends more training examples at low noise",
    "noise_parameterization": "x_t=alpha*x0+sigma*epsilon; alpha=exp(-5t); sigma=sqrt(1-exp(-10t))",
    "predictor": "Gaussian-covariance linear skip plus alpha(t)*MLP(x_t,time_features)",
    "data": "8 equal Gaussian modes; radius=2; sigma=.12",
    "coverage": "At least 1% of all 2048 generated points within radius .36 of a center",
    "sampler": "deterministic DDIM eta=0; uniform t grid from 1 to .001, then final x0 estimate",
    "initialization_note": "Tiny random neural residual plus a fixed Gaussian second-moment baseline; not a zero predictor",
    "reference_note": "Exact conditional expected epsilon for the SAME mixture and SAME schedule; used only in reference branch/evaluation"
  },
  "checkpoints": {
    "0": {
      "noise_target_MSE": 0.3466462386461925,
      "conditional_mean_error_MSE": 0.11890371735823559,
      "by_time": [
        {
          "t_range": [
            0,
            0.05
          ],
          "n": 908,
          "noise_target_MSE": 0.9195597783847196,
          "conditional_mean_error_MSE": 0.4757193204463834
        },
        {
          "t_range": [
            0.05,
            0.2
          ],
          "n": 943,
          "noise_target_MSE": 0.5202881992324372,
          "conditional_mean_error_MSE": 0.05819971073333699
        },
        {
          "t_range": [
            0.2,
            0.5
          ],
          "n": 1045,
          "noise_target_MSE": 0.08680401267912498,
          "conditional_mean_error_MSE": 0.00018575966494943397
        },
        {
          "t_range": [
            0.5,
            1.01
          ],
          "n": 1200,
          "noise_target_MSE": 0.002967291329670833,
          "conditional_mean_error_MSE": 3.105217327371017e-08
        }
      ],
      "samplings": {
        "20": {
          "n": 2048,
          "covered_modes": 4,
          "mode_counts_within_3sigma": [
            19,
            36,
            20,
            14,
            22,
            13,
            25,
            29
          ],
          "on_mode_fraction": 0.0869140625,
          "median_distance_to_center": 0.8499968196468977,
          "off_plot_fraction": 0.0146484375
        },
        "100": {
          "n": 2048,
          "covered_modes": 6,
          "mode_counts_within_3sigma": [
            25,
            34,
            14,
            25,
            17,
            22,
            30,
            25
          ],
          "on_mode_fraction": 0.09375,
          "median_distance_to_center": 0.8465869181647522,
          "off_plot_fraction": 0.037109375
        }
      }
    },
    "2000": {
      "noise_target_MSE": 0.27844650124564085,
      "conditional_mean_error_MSE": 0.05058195301449061,
      "by_time": [
        {
          "t_range": [
            0,
            0.05
          ],
          "n": 908,
          "noise_target_MSE": 0.6370114333852833,
          "conditional_mean_error_MSE": 0.19032023129840403
        },
        {
          "t_range": [
            0.05,
            0.2
          ],
          "n": 943,
          "noise_target_MSE": 0.4797693768483786,
          "conditional_mean_error_MSE": 0.019883639289979473
        },
        {
          "t_range": [
            0.2,
            0.5
          ],
          "n": 1045,
          "noise_target_MSE": 0.10091981474832874,
          "conditional_mean_error_MSE": 0.014130709723696854
        },
        {
          "t_range": [
            0.5,
            1.01
          ],
          "n": 1200,
          "noise_target_MSE": 0.0035222990069025636,
          "conditional_mean_error_MSE": 0.0007133716805740587
        }
      ],
      "samplings": {
        "20": {
          "n": 2048,
          "covered_modes": 8,
          "mode_counts_within_3sigma": [
            128,
            200,
            96,
            163,
            176,
            185,
            171,
            135
          ],
          "on_mode_fraction": 0.6123046875,
          "median_distance_to_center": 0.283906749946691,
          "off_plot_fraction": 0.0
        },
        "100": {
          "n": 2048,
          "covered_modes": 8,
          "mode_counts_within_3sigma": [
            126,
            192,
            103,
            163,
            172,
            199,
            169,
            155
          ],
          "on_mode_fraction": 0.62451171875,
          "median_distance_to_center": 0.28774819704597443,
          "off_plot_fraction": 0.0
        }
      }
    },
    "8000": {
      "noise_target_MSE": 0.2446225852379858,
      "conditional_mean_error_MSE": 0.01874684663416279,
      "by_time": [
        {
          "t_range": [
            0,
            0.05
          ],
          "n": 908,
          "noise_target_MSE": 0.510103650470085,
          "conditional_mean_error_MSE": 0.06719602036125287
        },
        {
          "t_range": [
            0.05,
            0.2
          ],
          "n": 943,
          "noise_target_MSE": 0.46653668758865086,
          "conditional_mean_error_MSE": 0.010645675059477246
        },
        {
          "t_range": [
            0.2,
            0.5
          ],
          "n": 1045,
          "noise_target_MSE": 0.09053742017363885,
          "conditional_mean_error_MSE": 0.004742436762024722
        },
        {
          "t_range": [
            0.5,
            1.01
          ],
          "n": 1200,
          "noise_target_MSE": 0.003536911692001908,
          "conditional_mean_error_MSE": 0.0006486494400919221
        }
      ],
      "samplings": {
        "20": {
          "n": 2048,
          "covered_modes": 8,
          "mode_counts_within_3sigma": [
            248,
            91,
            122,
            180,
            141,
            130,
            170,
            124
          ],
          "on_mode_fraction": 0.5888671875,
          "median_distance_to_center": 0.3005931038037627,
          "off_plot_fraction": 0.0
        },
        "100": {
          "n": 2048,
          "covered_modes": 8,
          "mode_counts_within_3sigma": [
            292,
            170,
            162,
            218,
            201,
            233,
            208,
            155
          ],
          "on_mode_fraction": 0.80029296875,
          "median_distance_to_center": 0.1654350773147341,
          "off_plot_fraction": 0.0
        }
      }
    }
  },
  "analytic_reference": {
    "20": {
      "n": 2048,
      "covered_modes": 8,
      "mode_counts_within_3sigma": [
        183,
        204,
        159,
        178,
        195,
        169,
        190,
        175
      ],
      "on_mode_fraction": 0.70947265625,
      "median_distance_to_center": 0.26043647358286615,
      "off_plot_fraction": 0.0
    },
    "100": {
      "n": 2048,
      "covered_modes": 8,
      "mode_counts_within_3sigma": [
        228,
        261,
        222,
        248,
        255,
        235,
        250,
        240
      ],
      "on_mode_fraction": 0.94677734375,
      "median_distance_to_center": 0.10271187832943826,
      "off_plot_fraction": 0.0
    }
  },
  "analytic_noise_target_MSE": 0.22568853355937113,
  "runtime_seconds": 20.518462419509888,
  "python": "3.13.12",
  "numpy": "2.4.6",
  "blas_threadpools": [
    {
      "user_api": "blas",
      "internal_api": "openblas",
      "num_threads": 1,
      "prefix": "libscipy_openblas",
      "filepath": "/root/venv/lib/python3.13/site-packages/numpy.libs/libscipy_openblas64_-32a4b2a6.so",
      "version": "0.3.31.188.0",
      "threading_layer": "pthreads",
      "architecture": "Haswell"
    }
  ],
  "run_utc": "2026-09-30T02:23:07Z",
  "real_reference_coverage": {
    "n": 2048,
    "covered_modes": 8,
    "mode_counts_within_3sigma": [
      270,
      251,
      261,
      241,
      235,
      238,
      287,
      243
    ],
    "on_mode_fraction": 0.9892578125,
    "median_distance_to_center": 0.1450050819067562,
    "off_plot_fraction": 0.0
  }
}
查看可执行代码
from IPython.display import display, Image
display(Image(filename='outputs/lab09-training-and-steps.png'))
display(Image(filename='outputs/lab09-trajectories.png'))
display(Image(filename='outputs/lab09-loss.png'))

只重采样:加载保存权重,改为50步

已保存结果或已运行过一次实验时,先执行前面的类、配置与函数定义单元,跳过 report = run(...) 重新训练单元,直接执行本单元。它只读取 outputs/lab09-weights.npz 和 outputs/lab09-data.npz,训练更新为0;改变的是生成阶段采样步数。若缺文件,先运行一次默认实验。

查看可执行代码
# Optional sampling-only cell: run the definition cells above, skip run().
# Files come from the saved experiment bundle or one previous run of this notebook.
resample_start = time.time()
saved_arrays = np.load('outputs/lab09-data.npz')
saved_weights = np.load('outputs/lab09-weights.npz')
loaded_model = MLP(CONFIG['architecture'], CONFIG['seed_model'])
for index, parameter in enumerate(loaded_model.p):
    parameter[...] = saved_weights[f'p{index}']
fixed_xT = saved_arrays['fixed_xT'].copy()
parameters_before = loaded_model.flat().copy()

# The only change is the sampler's call count: no backward(), update(), or run().
samples50, trajectory50, times50 = ddim(
    lambda x, t: predict(loaded_model, x, t)[0], fixed_xT, 50
)
assert np.array_equal(loaded_model.flat(), parameters_before)
assert np.array_equal(trajectory50[0], fixed_xT[:24])
metrics50 = coverage(samples50)
np.savez_compressed('outputs/lab09-resample50-data.npz',
                    samples=samples50, trajectory=trajectory50, times=times50,
                    fixed_xT=fixed_xT)
resample_record = {
    'sampling_steps': 50,
    'training_updates_this_run': 0,
    'weights_source': 'lab09-weights.npz; final 8000-update model',
    'initial_states_source': 'lab09-data.npz:fixed_xT',
    'parameters_unchanged': True,
    'trajectory_starts_at_saved_xT': True,
    'metrics': metrics50,
    'runtime_seconds': time.time() - resample_start,
    'run_utc': time.strftime('%Y-%m-%dT%H:%M:%SZ', time.gmtime())
}
dump_json(Path('outputs/lab09-resample50-metrics.json'), resample_record)
print(json.dumps(resample_record, ensure_ascii=False, indent=2))
查看保存的计算输出
{
  "sampling_steps": 50,
  "training_updates_this_run": 0,
  "weights_source": "lab09-weights.npz; final 8000-update model",
  "initial_states_source": "lab09-data.npz:fixed_xT",
  "parameters_unchanged": true,
  "trajectory_starts_at_saved_xT": true,
  "metrics": {
    "n": 2048,
    "covered_modes": 8,
    "mode_counts_within_3sigma": [
      285,
      156,
      154,
      211,
      191,
      216,
      196,
      149
    ],
    "on_mode_fraction": 0.7607421875,
    "median_distance_to_center": 0.16563237005659642,
    "off_plot_fraction": 0.0
  },
  "runtime_seconds": 8.828055620193481,
  "run_utc": "2026-09-30T02:23:30Z"
}