ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Sonnet 2 开发指南:基于 TensorFlow 2 的模块化神经网络库核心用法与序列化实战

Sonnet 2 开发指南:基于 TensorFlow 2 的模块化神经网络库核心用法与序列化实战 深度学习机器学习【免费下载链接】sonnetTensorFlow-based neural network library项目地址https://gitcode.com/gh_mirrors/so/sonnet点击查看免费下载Sonnet 是 DeepMind 设计并构建、运行在 TensorFlow 2 之上的神经网络库其核心目标是提供简单、可组合的机器学习研究抽象。本文以 README.md 为主线结合本仓库 sonnet/src 的源码实现完整讲解snt.Module编程模型、内置模块使用、自定义模块开发、参数访问、Checkpoint / SavedModel 序列化以及分布式训练帮助你直接上手用 Sonnet 搭建与训练自己的模型。1. 认识 Sonnet一个少约束的研究向神经网络库Sonnet 构建于 TensorFlow 2 之上被设计用于搭建用于多种目的的神经网络监督学习、非监督学习、强化学习等。与许多框架不同Sonnet 对你如何使用模块极其不加约束unopinionated模块被设计为自包含且彼此完全解耦Sonnet不自带训练框架鼓励用户自建训练循环或采用他人已有的方案Sonnet 的代码力求清晰、聚焦凡是做出了默认选择的地方例如参数初始化默认值都会尽量说明原因。从源码结构看整个库围绕一个核心概念展开snt.Modulesonnet/src/base.py。模块可以持有参数引用、其他模块的引用以及将某种函数应用于用户输入的方法。Sonnet 内置了大量预定义模块如snt.Linear、snt.Conv2D、snt.BatchNorm以及一些预定义的模块网络如snt.nets.MLP同时也强烈鼓励用户自己构建模块。本仓库当前版本在 sonnet/init.py 中声明为2.0.3.dev顶层 API 由 sonnet/init.py 统一导出包括Module、Optimizer、Sequential、Linear、Conv1D/2D/3D、BatchNorm、LSTM/GRU/VanillaRNN、nets、optimizers、distribute、functional等。2. 安装与验证2.1 安装依赖安装 Sonnet 2 前需先安装 TensorFlow 2 与 TensorFlow Probability$ pip install tensorflow tensorflow-probability $ pip install dm-sonnet2.2 验证安装运行下面代码确认两个库都可用import tensorflow as tf import sonnet as snt print(TensorFlow version {}.format(tf.__version__)) print(Sonnet version {}.format(snt.__version__))assert_tf2sonnet/src/base.py会在模块构造时强制校验 TensorFlow 处于 eager 模式Sonnet v2 要求 TensorFlow 2。2.3 仓库内置示例可直接运行的训练脚本examples/simple_mnist.py在 MNIST 上训练一个简易卷积网络5 个 epoch 内即可收敛Colab 风格的 Notebook 示例examples/mlp_on_mnist.ipynb、examples/little_gan_on_mnist.ipynb、examples/distributed_cifar10.ipynb更多说明见 examples/README.md。以simple_mnist.py为例其训练循环展示了 Sonnet 典型的无框架用法用tf.GradientTape计算梯度取model.trainable_variables再交给optimizer.apply(gradients, variables)更新examples/simple_mnist.py。3. 使用内置模块三行代码定义一个 MLPSonnet 内置了大量开箱即用的模块。例如要定义一个 MLP可以用snt.Sequential把一组模块串成链——上一个模块的输出自动成为下一个模块的输入mlp snt.Sequential([ snt.Linear(1024), tf.nn.relu, snt.Linear(10), ])Sequential以及大多数模块都定义了__call__方法因此可以直接像函数一样调用logits mlp(tf.random.normal([batch_size, input_size]))从 sonnet/src/sequential.py 的实现可以看到两个设计要点额外参数只传给第一个层Sequential.__call__会把*args/**kwargs透传给序列中的第一个模块sonnet/src/sequential.py其余层只接收上一层输出。因此它不支持向内部模块如BatchNorm动态传递is_training之类的标志官方给出的替代方案是直接子类化snt.Module并手写__call__例如把Conv2D、BatchNorm、relu组合成带is_training参数的自定义模块见 sonnet/src/sequential.py 文档示例。Linear模块的构造参数sonnet/src/linear.py参数说明默认值output_size输出维度必填with_bias是否包含偏置Truew_init权重初始化器截断正态分布标准差1 / sqrt(input_size)b_init偏置初始化器全零name模块名None自动取类名转小写下划线4. 访问模块参数variables 与 trainable_variablesSonnet 模块的参数在第一次被调用时才创建因为大多数情况下参数形状取决于输入形状随后可以通过两个属性获取all_variables mlp.variables # 该模块引用的【所有】 tf.Variable model_parameters mlp.trainable_variables # 仅可训练变量传给优化器tf.Variable并不只用于模型参数——例如snt.BatchNorm中用于保存状态moving mean / variance的也是变量。TensorFlow 原生通过trainable标记区分模型参数与其他变量因此variables返回全部变量trainable_variables只返回可训练变量这才是通常要传给优化器的列表非可训练变量由其他机制如 moving averages更新。源码层面的两个细节值得注意若在模块尚未被调用还没有变量时就访问这两个属性Sonnet 会抛出带详细提示的ValueError见 sonnet/src/base.py 的NO_VARIABLES_ERROR。如果模块确实是无状态模块可以用snt.allow_empty_variables(module)或snt.allow_empty_variables类装饰器抑制该报错sonnet/src/base.py这两个属性由tf.Module递归收集按属性名排序、BFS 遍历子模块Sonnet 在其上叠加了空结果即报错的防护。5. 自定义模块子类化 snt.ModuleSonnet 强烈鼓励通过子类化snt.Module定义自己的模块。下面从零实现一个简化版Linear层MyLinearclass MyLinear(snt.Module): def __init__(self, output_size, nameNone): super(MyLinear, self).__init__(namename) self.output_size output_size snt.once def _initialize(self, x): initial_w tf.random.normal([x.shape[1], self.output_size]) self.w tf.Variable(initial_w, namew) self.b tf.Variable(tf.zeros([self.output_size]), nameb) def __call__(self, x): self._initialize(x) return tf.matmul(x, self.w) self.b使用方式与内置模块完全一致mod MyLinear(32) mod(tf.ones([batch_size, input_size]))子类化snt.Module可以免费获得以下能力5.1 自动__repr__调试利器默认的__repr__会基于构造参数自动生成非常适合调试与内省 print(repr(mod)) MyLinear(output_size10)该能力由ModuleMetaclass在构造类时注入__repr__sonnet/src/base.pyauto_repr会比对构造参数与默认值只展示非默认参数sonnet/src/base.py。5.2 免费获得variables/trainable_variables mod.variables (tf.Variable my_linear/b:0 shape(10,) ...), tf.Variable my_linear/w:0 shape(1, 10) ...))5.3 自动进入模块 name scope注意上面变量名带有my_linear前缀——Sonnet 模块在调用其方法时会自动进入模块的 name scope从而为 TensorBoard 等工具提供更有意义的计算图分组my_linear内的所有操作会被归入my_linear分组。这是由ModuleMetaclass.__new__用with_name_scope包装类中的每个方法实现的sonnet/src/base.py核心逻辑在wrap_with_name_scopesonnet/src/base.py。几个进阶控制手段若某个方法不希望被自动包上 name scope例如实现clone()/transpose()时可用snt.no_name_scope装饰sonnet/src/base.py调试时可通过环境变量SNT_MODULE_NAME_SCOPES0关闭自动 name scope让堆栈更浅sonnet/src/base.py__init__中必须先调用super().__init__(namename)否则构造时会抛出明确错误sonnet/src/base.py。5.4snt.once惰性初始化的实现机制MyLinear中的_initialize被snt.once装饰确保每个实例上该方法只执行一次sonnet/src/once.py。其实现通过uuid生成的once_id记录在每个实例的_snt_once集合中重复调用直接跳过若被装饰方法有返回值或执行中抛错都会以ValueError反馈给用户。这也解释了为什么首次调用mod(tf.ones([...]))会创建变量__call__内部调用self._initialize(x)首次执行时根据输入形状创建w、b此后不再重建。对比仓库内置的 sonnet/src/linear.py其_initialize同样用once.once装饰并按输入末维inputs.shape[-1]创建形状为[input_size, output_size]的权重。5.5 模块网络snt.nets.MLP仓库还提供了更高层的模块网络 sonnet/src/nets/mlp.pysnt.nets.MLP支持output_sizes各层输出尺寸序列w_init/b_init/with_bias线性层参数activation层间激活函数默认 ReLUdropout_rateNone或0表示不使用 dropout使用 dropout 时必须在__call__传入is_trainingactivate_final是否激活最后一层reverse(activate_final..., name...)返回一个逐层反转的新 MLP可用于自编码器解码器要求模块已被至少调用过一次sonnet/src/nets/mlp.py。注意snt.nets.MLP使用 dropout 时必须显式传is_training且未使用 dropout 时传is_training反而会报错sonnet/src/nets/mlp.py。6. 序列化pickle、Checkpoint 与 SavedModelSonnet 支持多种序列化格式。6.1 pickle不推荐最简单的格式是 Python 的pickle所有内置模块都经过测试确保可在同一 Python 进程内保存/加载。但官方明确不鼓励使用 pickle它不被 TensorFlow 的许多部分良好支持实践中相当脆弱。6.2 TensorFlow Checkpoint训练中断恢复TensorFlow Checkpoint 用于在训练过程中周期性地保存参数值可在程序崩溃或被中断时恢复训练进度。Sonnet 与 TF Checkpoint 配合得天衣无缝checkpoint_root /tmp/checkpoints checkpoint_name example save_prefix os.path.join(checkpoint_root, checkpoint_name) my_module create_my_sonnet_module() # 任何继承自 snt.Module 的对象。 # Checkpoint 对象管理传入对象的 TensorFlow 状态。注意 Checkpoint 支持 # 创建即恢复restore on createmy_module 的变量不需要在 restore 之前 # 创建它们的值会在变量创建时被恢复。 checkpoint tf.train.Checkpoint(modulemy_module) # 大多数训练脚本会在存在 checkpoint 时先恢复例如训练被中断 # 如 GPU 被占用或云环境实例被抢占。 latest tf.train.latest_checkpoint(checkpoint_root) if latest is not None: checkpoint.restore(latest) for step_num in range(num_steps): train(my_module) # 训练过程中偶尔保存权重值。注意这是阻塞调用可能较慢通常 # 写入机器上最慢的存储。如果你的环境更可靠可以降低保存频率。 if step_num and not step_num % 1000: checkpoint.save(save_prefix) # 务必保存最终结果 checkpoint.save(save_prefix)仓库的 sonnet/src/conformance/checkpoints 目录下存放了大量内置模块linear_1x1、conv2d_3x3_2x2、lstm_1、resnet50等的参考 checkpoint配合 sonnet/src/conformance/checkpoint_test.py 可验证模块的 Checkpoint 兼容性。6.3 TensorFlow SavedModel与源码解耦的部署格式SavedModel 保存的是脱离 Python 源码的模型副本一个描述计算的 TensorFlow Graph 加上包含权重值的 checkpoint。第一步创建要保存的snt.Module并先调用一次以创建变量my_module snt.nets.MLP([1024, 1024, 10]) my_module(tf.ones([1, input_size]))第二步新建一个描述要导出哪些部分的模块。官方建议不要原地修改原模型而是包装一个新的导出模块这样可以对导出内容做精细控制——避免 saved model 过大并且可以只分享模型的一部分例如 GAN 中只导出生成器、保留判别器私有tf.function(input_signature[tf.TensorSpec([None, input_size])]) def inference(x): return my_module(x) to_save snt.Module() to_save.inference inference to_save.all_variables list(my_module.variables) tf.saved_model.save(to_save, /tmp/example_saved_model)保存后的目录结构$ ls -lh /tmp/example_saved_model total 24K drwxrwsr-t 2 tomhennigan 154432098 4.0K Apr 28 00:14 assets -rw-rw-r-- 1 tomhennigan 154432098 14K Apr 28 00:15 saved_model.pb drwxrwsr-t 2 tomhennigan 154432098 4.0K Apr 28 00:15 variables加载该模型非常简单可以在另一台没有任何构建代码的机器上完成loaded tf.saved_model.load(/tmp/example_saved_model) # 调用 inference 方法。注意这里运行的并非 to_save 中的 Python 代码 # 而是 saved model 中的 TensorFlow Graph。 loaded.inference(tf.ones([1, input_size])) # all_variables 属性用于获取恢复后的变量。 assert len(loaded.all_variables) 0注意加载得到的对象不是 Sonnet 模块而是一个容器对象只包含前面显式添加的方法inference和属性all_variables。仓库的 sonnet/src/conformance/saved_model_test.py 覆盖了这类导出/加载行为的测试。7. 分布式训练snt.distribute 与 CrossReplicaBatchNormSonnet 通过 sonnet/distribute.py实现位于 sonnet/src/distribute提供对自定义 TensorFlow 分布式策略的支持仓库内置了对应的多 GPU CIFAR-10 示例 examples/distributed_cifar10.ipynb。Sonnet 分布式训练与tf.keras的关键区别在于Sonnet 模块和优化器在分布式策略下不会改变行为——不会替你平均梯度也不会替你同步 batch norm 统计量。官方认为这些方面应由用户完全掌控不应被内建进库中。相应的代价是你需要在训练脚本里自行实现这些特性通常只需两行代码在应用优化器前对梯度做 all-reduce或换用显式感知分布式的模块例如snt.distribute.CrossReplicaBatchNormsonnet/src/distribute/distributed_batch_norm.py它在每一步都使用全部副本的完整 batch来计算 batch 统计量仅在使用snt.(Tpu)Replicator时生效适用于需要跨副本同步归一化统计的场景。此外snt.BatchNorm自身的机制也值得了解sonnet/src/batch_norm.py训练时使用当前 batch 的统计量并更新 moving averages测试时默认使用 moving averagestest_local_statsFalse也可以设为True改用本地统计量常用参数decay_rate默认0.999、eps默认1e-5、create_scale/create_offset是否创建可训练的 scale/offset、data_formatchannels_last/channels_first当输入是 4 维且为NHWC格式时会自动使用融合的FusedBatchNormV2op 提升性能sonnet/src/batch_norm.pyscale/offset也可以不创建而是调用时传入适合条件归一化等场景但二者不可同时使用sonnet/src/batch_norm.py。8. 更多参考资源官方风格文档docs/index.rst、docs/modules.rst、docs/api.rst含完整 API 参考与snt.distribute说明示例代码examples/README.md、examples/simple_mnist.py完整模块清单Linear、Conv1D/2D/3D、BatchNorm、LayerNorm、LSTM/GRU、DeepRNN、MLP、VQ-VAE等见 sonnet/init.py 的__all__导出列表内置模块的参考 checkpoint 与 golden 测试数据位于 sonnet/src/conformance/checkpoints对应的兼容性测试见 sonnet/src/conformance。总而言之Sonnet 的定位是简单、可组合、少约束你只需要掌握snt.Module这一个核心概念就能自由组合内置模块或自定义模块配合 TensorFlow 原生的 GradientTape、Checkpoint 与 SavedModel 完成从研究到部署的完整流程。赞分享深度学习机器学习【免费下载链接】sonnetTensorFlow-based neural network library项目地址https://gitcode.com/gh_mirrors/so/sonnet点击查看免费下载相关推荐Sonnet 官方文档导读基于 TensorFlow 2 的模块化神经网络组件库Sonnet 官方文档导读基于 TensorFlow 2 的模块化神经网络组件库 Sonnet 是由 DeepMind 设计并开源的、构建于 TensorFl深度学习机器学习Sonnet v2 API 参考指南TensorFlow 2 神经网络库的模块化架构与全量公开接口解析Sonnet v2 API 参考指南TensorFlow 2 神经网络库的模块化架构与全量公开接口解析 本文以 Sonnet基于 TensorFlow 2深度学习机器学习深入探索DeepMind SonnetTensorFlow 2上的模块化神经网络库深入探索DeepMind SonnetTensorFlow 2上的模块化神经网络库 DeepMind Sonnet是基于TensorFlow 2构建的模块化神深度学习机器学习上一篇pydantic-ai 生产部署实战用 6 步把 Agent 稳定跑进生产下一篇Nezha监控数据持久化终极指南SQLite与MySQL配置对比创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表