LAB 09 · 学得去噪器与完整采样
先阅读保存结果和解释,再按本册步骤选择是否运行。在 Deepnote 阅读与运行 · 下载 Notebook
LAB09|把解析去噪换成学得的网络
先看完整链条改变了什么
本册把 LAB05 的“已知解析预测器”换成真正训练的小型时间条件网络。网络从许多带噪练习题中学习;生成时从一个新噪声状态出发,重复调用它。训练题的正确噪声、生成时的当前状态、评价时使用的解析参照,在代码里是三个分开的角色。

先固定下排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。整体平均下降掩盖了不同时间段难度,也提醒我们:采样末段恰恰可能遇到网络仍不精确的位置。

从当前预测走向下一状态
本册用η=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的处理和起点近似都保留,不能宣称解析参照就等于无误差的数据生成器。
固定起点,只改变采样步数

选一条轨迹读:从空心圆出发,每次只知道此刻状态与时间;不是沿着早已指定的目标点移动。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"
}