WARNING
🧪 Beta公测版本提示:教程主体已完成,正在优化细节,欢迎大家提Issue反馈问题或建议。
LeWM — demo.py 代码详解
运行方式
cd docs/world-models/abstract/lewm/code
python demo.pyCPU 一两分钟。先二维质点(线性编码器 + 预测器,MSE + 高斯代理正则,潜空间 CEM),再倒立摆(CompactLeWM 恒等编码,火柴杆图只做可视化)。不是论文级 SIGReg(Epps–Pulley);玩具用均值/方差或随机投影逼近
和 PETS 对比:PETS 在观测/状态上拟合
代码逐段详解
第1步:观测为什么要手工特征
def observe(pos):
feat = np.array([
pos[0], pos[1], pos[0] ** 2, pos[1] ** 2,
np.sin(pos[0]), np.cos(pos[1]), pos[0] * pos[1], 1.0,
])
return feat + 0.02 * np.random.randn(8)线性编码器 1.0 是常数特征,给偏置一条通路(be 已经有偏置,多一维无妨)。
第2步:SIGReg 代理 — 为什么需要第二项损失
只有 MSE 时,编码器可以把所有
def sigreg_proxy(z):
mu = z.mean(axis=0)
std = z.std(axis=0) + 1e-6
return float(np.mean(mu ** 2) + np.mean((std - 1.0) ** 2))逼各维零均值、单位方差。+1e-6 防 std=0。这是对角高斯正则,不是完整特征函数 SIGReg。
倒立摆版:
dirs = np.random.randn(d, n_proj)
dirs /= np.linalg.norm(dirs, axis=0, keepdims=True) + 1e-8
h = z @ dirsCramér–Wold:多维分布由一维投影决定。随机方向上逼近 keepdims=True 才能按列广播除范数。
第3步:LinearLeWM.train_step — 停梯度与手写反传
z = self.encode(o)
nz_tgt = self.encode(no)
hat = self.predict(z, a)
err = hat - nz_tgt预测器拟合「下一观测的嵌入」,不是原始 Wp 输入是 concat(z, a),a 二维加速度,所以 Wp 形状 (Z_DIM+2, Z_DIM)。
self.Wp -= lr * (x.T @ err) / len(o)
self.bp -= lr * err.mean(axis=0)MSE 对线性层的梯度:x.T @ err:(feat, B) × (B, Z)。除以 len(o) 当 batch 平均。
编码器:
gz = (err @ Wp_z.T) / len(o)
gz = gz + LAMBDA_REG * (2.0 * mu) / len(o)
self.We -= lr * (o.T @ gz)链式法则:预测误差先流过 Wp 里属于 Wp[:Z_DIM]。再加上 no 同样推一把,避免「当前
注释里的「停梯度」:本实现 没有 nz_tgt.detach()(NumPy 没有计算图),意思是预测器损失不通过 nz_tgt 再改编码器去追 hat——目标嵌入只当固定靶(梯度只从 hat 一侧和正则来)。若让 err 同时改两端,encoder 可以把 hat 好猜的地方。
第4步:潜空间 CEM
for a in s:
z = model.predict(z, a)
scores.append(-np.sum((z - zg) ** 2))在嵌入里滚 HORIZON 步,终点对齐 zg = encode(observe(goal))。分数是负距离,CEM 仍取 argsort 最大。规划不用真实位置,只信模型。
闭环仍 true_step 执行 mu[0],下一步重新 encode(observe(pos))——MPC。
第5步:倒立摆 CompactLeWM — 恒等编码的教学妥协
def encode(self, o):
return np.asarray(o, dtype=np.float64)像素 JEPA/LeWM 的编码器在 CPU 上难训稳,这里 [1,0,0]。火柴杆 render 不进入训练,只给终局图。
残差预测:
nxt = z + x @ self.Wp + self.bp学
if np.ndim(a) == 0:
a = np.full((z.shape[0], 1), float(a))语法:规划时 a 是 Python float,训练时是 (B,)。统一成 (B,1) 才能 concatenate。np.ndim(a)==0 判断标量。
第6步:cem_plan_latent 的代价
cost += np.sum((z - zg) ** 2)
cost -= z[0] * np.exp(-0.05 * z[2] ** 2)既追直立嵌入,又加一项「像 PETS 的 cos 奖励」(z[0] 是 cos)。纯距离有时会绕远;奖励项把直立方向写进每一步。scores = -cost 再取精英。
collect 里 |θ|>0.8 则 reset:和 PETS 一样,只在近似线性区采样。
关键概念速查表
| 概念 | 本章落地 |
|---|---|
| 两项损失 | MSE( |
| 潜空间规划 | CEM 滚 predict,对齐 |
| 停梯度目标 | 不让目标嵌入追预测 |
| 恒等编码 | 摆的 CPU 妥协 |
x.T @ err | 线性层 MSE 梯度 |
np.ndim | 标量动作对齐 batch |
源码位置
clone 后打开(相对仓库根目录):
docs/world-models/abstract/lewm/code/demo.py