资讯详情

在 Pyro 中实现深度核学习(Deep Kernel Learning):用 CNN 扭曲 RBF 核在 MNIST 上做分类的完整实战

发布时间:2026/9/25 3:48:57

500+
企业客户服务经验
120+
行业领域内容覆盖
3000+
原创页面设计沉淀
98%
客户满意度

在 Pyro 中实现深度核学习(Deep Kernel Learning):用 CNN 扭曲 RBF 核在 MNIST 上做分类的完整实战

人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载本篇技术指南以 Pyro 仓库中 dkl.rst 教程及其嵌入的完整示例脚本 sv-dkl.py 为骨架讲解如何用卷积神经网络CNN对 RBF 核函数做输入扭曲input warping构造深度核并配合变分稀疏高斯过程VariationalSparseGP在 MNIST 数据集上进行 10 分类与二分类。读完本文你将掌握 Pyro GP 模块中Warping核、诱导点inducing points、Binary/MultiClass似然与TraceMeanField_ELBO的完整组合用法并能在本地复现 98.45%10 分类与 99.41%二分类的准确率。一、Deep Kernel Learning 的背景与本文实现思路深度核学习Deep Kernel LearningDKL的核心思想来自 Wilson、Hu、Salakhutdinov 与 Xing 的工作Stochastic Variational Deep Kernel Learning不再使用固定形式的核函数而是让核函数本身由一个深度网络参数化从而在保留高斯过程不确定性建模能力的同时获得深度特征表达。在 Pyro 的 GP 模块中这一思想被具体化为核扭曲kernel warping用一个输入扭曲函数f与一个输出扭曲多项式q构造新核k_new(x, z) q(k(f(x), f(z)))其中f可以是任意可微映射这里就是 CNNq是系数非负的多项式。该公式的实现位于 kernel.py 的Warping类中。需要注意本示例与原论文实现路径的区别这一差异在脚本 docstring 中有明确说明原论文将 CNN 作为特征提取层在其输出之上再加一个高斯过程层因此诱导点位于提取后的特征空间本示例将 CNN 模块与 RBF 核直接拼接为一个深度核因此诱导点位于原始图像空间。从源码结构看本示例用gp.kernels.Warping(rbf, iwarping_fncnn)一次性完成CNN 降维 RBF 计算协方差的整体封装CNN 把高维图像映射为低维向量RBF 核在 CNN 的输出上计算协方差矩阵Warping.forward中self.kern(self.iwarping_fn(X), self.iwarping_fn(Z))正是这一调用链见 kernel.py。二、整体架构CNN 扭曲核 变分稀疏高斯过程示例的核心模型配置如下rbf gp.kernels.RBF(input_dim10, lengthscaletorch.ones(10)) deep_kernel gp.kernels.Warping(rbf, iwarping_fncnn) gpmodule gp.models.VariationalSparseGP( XXu, yNone, kerneldeep_kernel, XuXu, likelihoodlikelihood, latent_shapelatent_shape, num_data60000, whitenTrue, jitter2e-6, )2.1 核的构成RBF WarpingRBF 核又名 Squared Exponential定义在 isotropic.pyk(x,z) σ² · exp(−0.5 · |x−z|² / l²)这里input_dim10对应 CNN 最后一个全连接层输出 10 维特征10 分类时恰好也对应 10 个类别lengthscaletorch.ones(10)表示每个维度初始长度尺度均为 1是可训练参数。Warping 扭曲核Warping(kern, iwarping_fncnn)把 RBF 核包在 CNN 之上。Warping还支持可选的owarping_coef输出扭曲多项式系数必须为非负整数且多项式次数至少为 1本示例未使用。其 docstring 中还演示了如何用pyro.module包装可训练线性层再传入Warping的通用做法见 kernel.py。2.2 分类器 CNN 结构class CNN(nn.Module): def __init__(self): super().__init__() self.conv1 nn.Conv2d(1, 10, kernel_size5) self.conv2 nn.Conv2d(10, 20, kernel_size5) self.fc1 nn.Linear(320, 50) self.fc2 nn.Linear(50, 10) def forward(self, x): x F.relu(F.max_pool2d(self.conv1(x), 2)) x F.relu(F.max_pool2d(self.conv2(x), 2)) x x.view(-1, 320) x F.relu(self.fc1(x)) x self.fc2(x) return x该 CNN 改编自 PyTorch 官方 MNIST 示例脚本 docstring 已注明结构为Conv2d(1→10, k5) → ReLU MaxPool2d(2) → Conv2d(10→20, k5) → ReLU MaxPool2d(2) → view(-1, 320) → Linear(320→50) → ReLU → Linear(50→10)。最终输出 10 维向量恰好匹配RBF(input_dim10)。值得注意这个 CNN 既是分类特征提取器又承担了核的输入扭曲函数角色这是 DKL 中核与网络共享参数的体现——CNN 的权重既影响特征表达也通过核函数影响协方差结构整套参数由 SVI 联合优化。2.3 变分稀疏高斯过程支撑小批量训练VariationalSparseGP的实现位于 vsgp.py。当训练样本 N 很大时直接计算k(X, X)的逆代价高昂该模型引入诱导输入参数Xu把模型写成[f, u] ~ GP(0, k([X, Xu], [X, Xu])) y ~ p(y|f) · p(f)并用变分分布q(f,u) p(f|u) · q(u)逼近后验其中q(u)是均值为u_loc、协方差三角因子为u_scale_tril的多元正态分布二者作为变分参数参与学习。关键特性源码 docstring 明确给出复杂度训练O(NM²)测试O(M³)变分参数规模O(M²)其中 N 为训练样本数、M 为诱导点数——这正是时间复杂度随数据量线性增长docstring 原话为time complexity scales linearly to the number of data points的来源num_data60000告知模型完整训练集大小配合poutine.scale(scalenum_data / X.size(0))见 vsgp.py在 mini-batch 训练时对每个批次的 log-likelihood 进行缩放得到无偏的完整数据集估计whitenTrue将变分参数u_loc、u_scale_tril通过Kuu的 Cholesky 分解Luu的逆做白化变换源码注释与 docstring 均说明开启该标志有助于变分优化收敛jitter2e-6一个小正数加到协方差矩阵对角线上以稳定 Cholesky 分解模型内部通过Kuu.view(-1)[::M1] self.jitter实现见 vsgp.py。三、似然选择Binary 与 MultiClass示例根据是否二分类选择不同似然if args.binary: likelihood gp.likelihoods.Binary() latent_shape torch.Size([]) else: likelihood gp.likelihoods.MultiClass(num_classes10) latent_shape torch.Size([10])Binary似然binary.py基于Bernoulli分布默认响应函数为 sigmoid输出须落在 (0,1)用于二分类此时潜过程是标量故latent_shape为空。MultiClass似然multi_class.py基于Categorical分布默认响应函数为 softmax对输入最右轴做归一化用于多分类。因为它返回每个类别的概率列表所以 GP 模型必须返回latent_shape torch.Size([10])的潜输出——这是示例源码中特别强调的一点见 sv-dkl.py。二分类时训练目标也做了特殊处理target (target % 2).float()把 0–9 的数字标签映射为 0/1 奇偶标签。四、诱导点的初始化从训练数据中采样诱导点是从训练数据中随机抽取的真实图像batches [] for i, (data, _) in enumerate(train_loader): batches.append(data) if i ((args.num_inducing - 1) // args.batch_size): break Xu torch.cat(batches)[: args.num_inducing].clone()默认--num-inducing70、--batch-size64时约取前两个 batch 共 70 张图作为Xu。Xu在模型中注册为torch.nn.Parameter见 vsgp.py会在训练中随梯度更新。因为本示例的诱导点位于原始图像空间与论文不同Xu直接就是图像像素张量。五、训练循环与变分目标5.1 每个 batch 的训练步骤gpmodule.set_data(data, target) optimizer.zero_grad() loss loss_fn(gpmodule.model, gpmodule.guide) loss.backward() optimizer.step()优化器为torch.optim.Adam(gpmodule.parameters(), lrargs.lr)默认学习率 0.01变分目标使用TraceMeanField_ELBO即带 JIT 加速的JitTraceMeanField_ELBO可选取elbo.differentiable_loss作为损失函数elbo infer.JitTraceMeanField_ELBO() if args.jit else infer.TraceMeanField_ELBO() loss_fn elbo.differentiable_loss均场mean-field假设与稀疏 GP 的q(f,u) p(f|u)q(u)结构天然匹配每--log-interval默认 10个 batch 打印一次训练进度与 loss。5.2 测试预测f_loc, f_var gpmodule(data) # GP 后验的均值和方差 pred gpmodule.likelihood(f_loc, f_var) # 经似然响应函数得到类别 correct pred.eq(target).long().cpu().sum().item()gpmodule(data)调用VariationalSparseGP.forward见 vsgp.py基于已学习的Xu、u_loc、u_scale_tril与核参数计算测试点后验随后likelihood(f_loc, f_var)用 sigmoid/softmax 把潜输出转换为预测类别与标签比对统计准确率。六、命令行参数总览脚本通过argparse暴露了以下可调参数见 sv-dkl.py参数默认值说明--data-dir PATHNone自动缓存到脚本旁.data/MNIST 数据缓存目录--num-inducing N70诱导点数量--binaryFalseflag是否做二分类奇偶分类--batch-size N64训练 batch 大小--test-batch-size N1000测试 batch 大小--epochs N10训练轮数达到报告精度需 16 轮--lr LR0.01Adam 学习率--cudaFalseflag启用 CUDA 训练--jitFalseflag启用 PyTorch JIT 编译 ELBO--seed S1随机种子pyro.set_rng_seed--log-interval N10每 N 个 batch 打印训练状态数据加载由 util.py 的get_data_loader完成使用了 MNIST 的标准归一化变换transforms.Normalize((0.1307,), (0.3081,))在 CI 环境下数据目录固定为~/.data否则默认缓存到脚本同级目录.data/见 util.py。七、运行方式与预期结果在仓库根目录执行示例脚本对 Pyro 版本有校验assert pyro.__version__.startswith(1.9.1)# 10 分类MNIST 数字识别默认参数跑 10 轮 python examples/contrib/gp/sv-dkl.py # 二分类奇偶分类 python examples/contrib/gp/sv-dkl.py --binary # GPU 训练 JIT 加速训练 16 轮以复现文档报告的精度 python examples/contrib/gp/sv-dkl.py --cuda --jit --epochs 16脚本 docstring 记录的默认超参数下 16 轮训练结果为10 分类 MNIST 准确率 98.45%二分类准确率 99.41%。每个 epoch 结束后脚本会输出测试集准确率与单轮耗时Amount of time spent for epoch ...便于观察收敛与调参。八、延伸阅读与相关资源核扭曲的完整实现与Warping类的 API 文档kernel.py变分稀疏 GP 的数学推导、复杂度分析与参数说明vsgp.pyGP 模块整体使用指南见 contrib.gp.rst更多 GP 分类/回归示例见 examples/contrib/gp 目录如sv-dkl.py同级的其他模型脚本若想验证Warping、VariationalSparseGP的行为可参考测试目录 tests/contrib/gp 下的test_kernels.py、test_models.py等用例。通过本示例你可以看到 Pyro 把深度特征 高斯过程不确定性的组合压缩成了三行核心配置构造 RBF →Warping包裹 CNN → 传入VariationalSparseGP这为在图像、时序等高维输入上扩展贝叶斯深度模型提供了简洁而可扩展的范式。赞分享人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载相关推荐Android-Password-Store TOTP功能使用如何安全管理双重认证令牌Android Password Store TOTP功能使用如何安全管理双重认证令牌 在数字时代账号安全至关重要而双重认证2FA是保护账号的重要手段Windows 驱动清理完整指南用免费工具 RAPR 安全清理驱动存储库回收 C 盘 10GB 空间Windows 驱动清理完整指南用免费工具 RAPR 安全清理驱动存储库回收 C 盘 10GB 空间 插上 U 盘半天没反应C 盘图标红了大半年每次都要桌面应用运维MXNet 云端部署实战在 AWS 上使用 SageMaker、Deep Learning AMI 与 S3 训练深度学习模型MXNet 云端部署实战在 AWS 上使用 SageMaker、Deep Learning AMI 与 S3 训练深度学习模型 导读 深度学习训练往往需要极其深度学习机器学习人工智能上一篇quick-lru常见问题解答解决90%开发者遇到的缓存难题下一篇Berkeley Document Summarizer常见问题解决从依赖安装到模型训练的完整排错指南 创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
热门专题

继续阅读更多专题内容

围绕企业服务、数字化转型与官网运营的常青话题,持续输出深度内容

企业官网建设指南 企业托管服务模式 财税政策与解读 企业数字化转型 官网SEO与获客 网站安全与运维
配套服务

读完这篇文章,了解更多服务

从整站搭建到SEO布局,17项核心服务助您打造高转化的企业官网

01

企业托管整站搭建

从信息架构到栏目预留,搭建可生长的企业站点骨架,每个页面独立原创设计。...

了解详情
02

规整可信网页设计

雪地靴温暖风原创设计,金属铜线条贯穿全页,拒绝通用模板与AI流水线。...

了解详情
03

企业服务SEO布局

关键词体系与语义化结构,从建站源头为搜索排名而生。...

了解详情
04

业务预约咨询表单

多场景表单与线索收集体系,把访问流量转化为可追踪的销售线索。...

了解详情
05

企业服务站点运维

安全巡检、数据备份与内容更新支持,全年守护网站稳定运行。...

了解详情
06

全终端商务适配

电脑、平板、手机一致呈现,移动端体验与转化同样出色。...

了解详情
需要专业建议?

让专业顾问为您解读行业趋势

关于企业官网建设、SEO获客与数字化转型的任何疑问,欢迎一对一咨询我们的专业顾问。