
google-research meta_augmentationMAML 元学习与正弦回归/分类实验完整复现指南【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research本指南以meta_augmentation/classification_and_sinusoidal_regression代码库为对象系统讲解如何基于 Model-Agnostic Meta-LearningMAMLFinn et al., ICML 2017在 Omniglot、MiniImagenet、DClaw 三类分类数据集以及正弦函数回归sinusoidal regression上完成元训练meta-training与预训练基线pretraining baseline实验并覆盖测试、随机初始化基线、三种任务构造方式non-exclusive / intrashuffle / intershuffle与 Meta-Augmentation 噪声设置。读完本文你将掌握该仓库从依赖安装、数据准备、训练、评估到结果导出的完整命令行流程并理解 MAML 内部循环inner loop与元更新outer loop在 maml.py 中的具体实现。项目背景Meta-Learning Requires Meta-Augmentation该代码库是论文《Meta-Learning Requires Meta-Augmentation》NeurIPS 2020见 meta_augmentation/README.md的配套实现其代码源自 MAML 官方代码库cbfinn/maml并针对元增强研究做了修改。整体目录结构如下README.md本文所依据的核心使用文档main.py训练与测试的入口集中定义全部命令行参数maml.pyMAML 算法主体与网络定义全连接 / 卷积data_generator.py四类数据源sinusoid / omniglot / miniimagenet / dclaw的任务批次生成utils.py图像读取、卷积块、归一化与损失函数MSE / softmax 交叉熵special_grads.py为 TensorFlow 补齐 MaxPoolGrad 的二阶梯度算子data/DClaw 已预处理数据以及 Omniglot、MiniImagenet 的预处理脚本。环境依赖与安装按照 README.md 的要求运行环境需满足Python 2.x 或 3.xTensorFlow v1.0仓库的 requirements.txt 进一步指定tensorflow1.14.0、numpy1.16.4、Pillow4.0.0。需要注意代码大量使用tensorflow.compat.v1与tf.contrib如tf.contrib.metrics.accuracy、tf.contrib.layers.xavier_initializer因此必须使用 TensorFlow 1.x 环境或基于 2.x 的 compat.v1 模式运行且建议在 Python 3 环境配合 Pillow 处理图像缩放详见下文数据预处理。训练命令MAML 与预训练基线统一入口所有实验统一通过 main.py 启动核心命令为python main.py --logdirlogs/log_dir --expt_nameexpt_name其中--logdir指定摘要TensorBoard summaries与模型检查点checkpoint的输出目录--expt_name指定实验任务构造方式见下文任务设置。模型权重与实验目录名由一个自动生成的exp_string决定见 main.py它编码了类别数、meta batch size、update batch size、内循环步数、内循环学习率、滤波器数、归一化方式等关键超参数因此不同的超参数组合会自动落到不同子目录互不干扰。MAML 训练MAML 通过双层优化学习易初始化的模型参数内循环针对单个任务做若干步梯度下降外循环基于内循环后的损失更新全局参数。启动 MAML 训练只需将预训练迭代数置 0python main.py --metatrain_iterations60000 --pretrain_iterations0预训练基线预训练pretraining等价于把所有任务样本混合做普通监督学习用作对比基线python main.py --metatrain_iterations0 --pretrain_iterations60000从源码看训练循环会先执行pretrain_iterations次model.pretrain_op再执行metatrain_iterations次model.metatrain_op见 main.py两者的优化器都是 Adamtf.train.AdamOptimizer学习率由--meta_lr控制默认 0.001区别在于预训练最小化的是内循环之前的损失total_loss1而元训练最小化的是第num_updates次内循环更新之后的损失total_losses2[num_updates-1]见 maml.py。注意对 miniimagenet / dclaw元训练梯度还会被裁剪到 [-10, 10]。测试命令常规测试与随机初始化基线在测试集上评估python main.py --trainFalse --test_setTrue--trainFalse进入测试分支--test_setTrue表示使用测试集否则使用验证集。测试时程序会把meta_batch_size临时改为 1见 main.py并对每个测试点执行若干次内循环更新后统计精度/损失。测试脚本固定运行NUM_TEST_POINTS 600个任务见 main.py最终将均值、标准差与 95% 置信区间1.96 * std / sqrt(600)写入 logdir 下的 CSV 文件同时把原始结果序列化到同名的.pkl文件文件名包含test_ubsupdate_batch_size_stepsizeupdate_lr见 main.py方便后续画曲线与统计分析。随机初始化基线python main.py --rand_initTrue --trainFalse --test_setTrue--rand_initTrue表示不加载任何检查点直接用随机初始化的网络做测试用于衡量完全不学习的下界表现见 main.py 与加载逻辑if not FLAGS.rand_init:。此外测试时默认加载最新模型若要指定某个迭代步的模型可用--test_iteriter--resumeFalse可关闭断点续训。各数据集的标准训练参数Miniimagenet--datasourceminiimagenet --meta_batch_size4 --update_batch_size1 --update_lr0.01 --num_updates5 --num_classes5 --num_filters32 --max_poolTrue即 5-way 分类、1-shot 支持集、内循环学习率 0.01、内循环 5 步、每轮元更新采样 4 个任务、卷积网络 32 个滤波器并启用 max pooling。对应的网络结构为 4 层 3×3 卷积 全连接输出层输入为 84×84×3 的 RGB 图见 maml.py 与 data_generator.py。Omniglot--datasourceomniglot --meta_batch_size32 --update_batch_size1 --update_lr0.4 --num_updates1Omniglot 使用 28×28 灰度图输入为 1 通道源码中channels1默认num_filters64不强制--max_pool默认 False即使用步长卷积。数据加载时还会对每类数字随机做 0/90/180/270 度旋转以扩充类别见 data_generator.py 与tf.image.rot90调用这也是 Omniglot 任务数量巨大的来源。DClaw--datasourcedclaw --meta_batch_size4 --update_batch_size1 --update_lr0.01 --num_updates5 --num_classes2 --num_filters32 --max_poolTrue --dclaw_pn1DClaw 是仓库自带的机械爪claw正/负样本二分类数据集2-way84×84×3 图像。--dclaw_pn取 1/2/3对应三套不同的 train/val/test 物体划分论文报告的结果是三套划分的平均见 README.md。从 data_generator.py 可看到train/val/test 目录分别由./data/dclaw/train{pn}、./data/dclaw/val{pn}、./data/dclaw/test{pn}决定本仓库中 DClaw 图像已全部预处理为 84×84 并切分完毕无需重复处理。DClaw 正负样本示例来自测试划分 1均为 84×84 已预处理图像正弦函数回归Sinusoidal Regression--datasourcesinusoid --normNone --update_batch_size10 --sine_seed1 --meta_batch_size6正弦回归任务中每个任务是一个y amp * sin(x - phase)函数支持集与查询集各 10 个点update_batch_size10。--sine_seed控制正弦任务池的随机种子。模型使用 40-40 两隐层全连接网络损失为 MSE见 maml.py且由于回归任务不做归一化必须显式设置--normNone。从 data_generator.py 可以看到正弦任务池的构造细节振幅在 [0.1, 5.0] 内均匀采样、相位在 [0, π] 内均匀采样、输入 x 落在 [-5.0, 5.0]代码用sine_seed预先生成 10 组 (amp, phase) 与对应的输入区间起点input_start其中前 6 组用于训练、第 7-8 组用于验证、第 9-10 组用于测试见generate_sinusoid_ne_batchdata_generator.py从而构造出非互斥训练与测试函数不同的实验设定。任务设置三种任务构造与 Meta-Augmentation任务构造方式由--expt_name控制这是该仓库与原始 MAML 代码最核心的差异所在对应论文中互斥/非互斥任务的研究问题。分类任务Omniglot / Miniimagenet / DClaw非互斥Non-mutually-exclusive--expt_namenon_exclusive。训练时按固定顺序将相邻的num_classes个类打包成一个任务不随机抽样类别模型反复见到同一批任务Intrashuffle--expt_nameintrashuffle。同样按相邻类分组但在任务内部对类别做随机洗牌Intershuffle--expt_nameintershuffle默认值。每次从全部类别中随机采样num_classes个类构成任务再对类内顺序做随机洗牌即标准的小样本任务采样方式。这三条分支分别对应 data_generator.py 中make_data_tensor的三种任务打包逻辑non_exclusive直接取相邻类、intrashuffle取相邻类后random.shuffle、其余情况intershuffle用random.sample(folders, num_classes)全局随机采样。正弦回归任务非互斥--expt_namenon_exclusive均匀噪声元增强Meta-augmentation with uniform noise--expt_nameuniform_noise。该设置会在训练时给正弦函数输出叠加一个[-1, 1]的均匀随机平移output_shift见 data_generator.py即在训练阶段引入标签噪声这正是论文 Meta-Augmentation 的核心思想——通过任务层面的数据增强缓解元学习在非互斥任务上的过拟合。其他相关参数--label_smooth训练时对分类标签做标签平滑源码中按均匀分布随机抽取噪声量将 one-hot 标签混合为(1-noise)*label noise/5见 maml.py--stop_grad为加速可关闭元优化中的二阶导数tf.stop_gradient--normbatch_norm默认/layer_norm/None--baselineoracle仅适用于 sinusoid把任务 id振幅与相位作为额外输入通道喂给网络构造已知任务身份的上界基线见 main.py。数据集路径与预处理除 DClaw 外其余数据集需自行下载并放到指定目录详见 README.mdOmniglot路径./data/omniglot_resizedtrain / val / test 划分在代码内完成验证集固定取 100 个字符类预处理先下载images_background与images_evaluation合并到data/omniglot/后执行 resize_images.py该脚本用 PIL 的 LANCZOS 滤波器将所有图像缩放到 28×28cd data/ cp -r omniglot/* omniglot_resized/ cd omniglot_resized/ python resize_images.pyMiniimagenet路径./data/miniImagenet/train或val/test预处理从 Ravi Larochelle17 获取 mini-ImageNet 原始图像与 train/val/test 三个 CSV 文件将 CSV 放入data/miniImagenet/、图像放入data/miniImagenet/images/然后在data/miniImagenet/下执行 proc_images.py。该脚本先将所有图像缩放到 84×84再按 CSV 中的类别标签把每张图移动到train/label/、val/label/、test/label/子目录从而匹配 MAML 的目录式数据加载方式。DClaw路径./data/dclaw/train或val/test预处理数据已在仓库内预处理完毕proc_images.py 将图像统一缩放为 84×84。目录train1, val1, test1对应一套划分dclaw_pn1/2/3选择不同划分。Sinusoid数据完全在代码内生成与处理DataGenerator.generate_sinusoid_ne_batch无需外部下载。训练过程监控与断点机制训练循环main.py内置三类日志节奏摘要间隔SUMMARY_INTERVAL 100每 100 步记录 pre-update / post-update 损失分类任务还包括精度到 TensorBoard打印间隔PRINT_INTERVALsinusoid 为 1000 步、分类为 200 步打印最近一段窗口的平均 pre/post 损失验证间隔TEST_PRINT_INTERVAL PRINT_INTERVAL * 5在验证集上评估。分类任务按最佳验证精度保存model-early-stopiter检查点sinusoid 按最低验证损失保存modeliter检查点见 main.py。由于 checkpoint 名称格式不同--test_iter在两类任务上的解析逻辑也有对应分支main.py。参数速查表参数默认值说明文档推荐取值--datasourcesinusoidsinusoid/omniglot/miniimagenet/dclaw按实验选择--expt_nameintershufflenon_exclusive/intrashuffle/intershuffle/uniform_noise按任务设置选择--metatrain_iterations15000MAML 元训练迭代数MAML 用 60000--pretrain_iterations0预训练迭代数基线用 60000--meta_batch_size25每轮元更新采样的任务数miniimagenet/dclaw 用 4omniglot 用 32sinusoid 用 6--meta_lr0.001外层 Adam 学习率默认--update_batch_size5内循环 K-shot 样本数分类用 1sinusoid 用 10--update_lr1e-3内循环步长 αomniglot 用 0.4其余用 0.01--num_updates1内循环更新步数omniglot 用 1其余用 5--num_classes5分类的 way 数omniglot/miniimagenet 用 5dclaw 用 2--num_filters64卷积滤波器数miniimagenet/dclaw 用 32--max_poolFalse是否用 max pooling 替代步长卷积miniimagenet/dclaw 设为 True--normbatch_normbatch_norm/layer_norm/Nonesinusoid 用None--dclaw_pn1DClaw 划分编号 1/2/31可换 2/3--sine_seed1正弦任务池随机种子1--trainTrueTrue 训练 / False 测试测试时设 False--test_setFalse用测试集还是验证集测试时设 True--rand_initFalse测试时用随机初始化网络随机初始化基线设 True--logdir/tmp/data摘要与检查点目录建议显式指定--test_iter-1加载指定迭代的模型-1 为最新按需设置复现注意事项该代码为 TensorFlow 1.x 生态tf.contrib、tf.InteractiveSession、队列读取器建议使用 TensorFlow 1.14-1.15 的 Python 3 环境DClaw 数据已随仓库提供可直接运行Omniglot 与 MiniImagenet 需先按上文预处理脚本处理处理只需执行一次测试时程序自动使用 meta_batch_size1 并加载对应exp_string目录下的最优检查点若想对比随机初始化下界务必添加--rand_initTrue论文的Meta-Augmentation效果复现关键在于 sinusoid 任务使用--expt_nameuniform_noise、分类任务对比non_exclusive/intrashuffle/intershuffle三种任务构造——这正是该代码库区别于原始 MAML 的扩展点对应的数据生成逻辑均可在 data_generator.py 中逐行验证。【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考