- 机器学习
- 深度学习
【免费下载链接】jax
Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more
本指南以 JAX 的jax.nn.initializers模块为主体,系统讲解神经网络的参数初始化机制。该模块提供了与 Keras、Sonnet 等主流框架一致的常用初始化器,覆盖常量初始化、随机分布初始化、基于 fan 值的方差缩放初始化(Glorot/Xavier、He/Kaiming、LeCun 系列)以及正交初始化等完整方案。读完本文,你将掌握每个初始化器的接口约定、参数语义、适用场景,并通过仓库源码理解其底层实现与验证逻辑。
初始化器的统一接口:(key, shape, dtype)
jax.nn.initializers是jax.nn模块的子模块(见 docs/jax.nn.rst 中的toctree声明),官方文档将其定位为"与 Keras 和 Sonnet 中定义一致的常用神经网络层初始化器"。
模块中所有初始化器遵循统一的函数签名约定:
一个初始化器是一个接收三个参数
(key, shape, dtype)的函数,返回维度为shape、数据类型为dtype的数组。其中key是 PRNG 密钥(例如由jax.random.key生成),用于产生随机数以初始化数组。
在源码 jax/_src/nn/initializers.py 中,这一约定被形式化为一个运行时可检查的InitializerProtocol 类型:
@export @typing.runtime_checkable class Initializer(Protocol): @staticmethod def __call__(key: KeyArray, shape: core.Shape, dtype: DTypeLikeInexact = jnp.float_) -> Array: raise NotImplementedError关键要点:
key:jax.random.key(typed key)或jax.random.PRNGKey(legacy key)生成的 PRNG 密钥。随机型初始化器用它驱动采样;常量型初始化器(如zeros、ones)会忽略它;shape:要初始化数组的维度形状,如(2, 3)、(32, 128)或卷积核的 4D/5D 形状;dtype:默认值为jnp.float_。源码中所有初始化器都会先调用dtypes.canonicalize_dtype(dtype)做 dtype 规范化,因此传入jnp.float64等弱类型时会被映射到当前配置下的实际类型;- 返回类型:两种 API 形态并存——模块级函数(如
zeros、ones)直接是初始化器,而uniform、normal、variance_scaling等"工厂函数"调用后返回一个闭包式初始化器。
以variance_scaling为代表的工厂式初始化器,其返回值形态可以同时兼容"先建后调"两种用法,这一点在 tests/nn_test.py 的testInitializerProvider测试中有专门验证。
常量初始化器:zeros、ones、constant
zeros与ones是模块级函数,直接返回全零/全一数组,key参数被忽略(源码 jax/_src/nn/initializers.py):
>>> import jax, jax.numpy as jnp >>> jax.nn.initializers.zeros(jax.random.key(42), (2, 3), jnp.float32) Array([[0., 0., 0.], [0., 0., 0.]], dtype=float32) >>> jax.nn.initializers.ones(jax.random.key(42), (3, 2), jnp.float32) Array([[1., 1.], [1., 1.], [1., 1.]], dtype=float32)constant(value, dtype=jnp.float_)则是工厂函数,返回一个以指定常数值填充数组的初始化器(jax/_src/nn/initializers.py):
>>> initializer = jax.nn.initializers.constant(-7) >>> initializer(jax.random.key(42), (2, 3), jnp.float32) Array([[-7., -7., -7.], [-7., -7., -7.]], dtype=float32)应用场景:zeros常用于偏置(bias)向量初始化;constant常用于自定义填充值(如注意力掩码、特定层的固定初值);ones常用于某些归一化层参数的初值。三者底层分别对应jnp.zeros、jnp.ones、jnp.full。
基础随机初始化器:uniform、normal、truncated_normal
这三个工厂函数提供最朴素的高斯/均匀随机初始化,各自带一个尺度参数:
| 初始化器 | 默认参数 | 采样分布 | 返回范围 |
|---|---|---|---|
uniform(scale=1e-2, dtype) | scale=1e-2 | 实均匀分布 | [0, scale) |
normal(stddev=1e-2, dtype) | stddev=1e-2 | 实正态分布 | 均值 0,标准差stddev |
truncated_normal(stddev=1e-2, dtype, lower=-2.0, upper=2.0) | stddev=1e-2 | 截断正态分布 | lower*stddev < x < upper*stddev |
其实现均基于jax.random采样后乘以尺度(jax/_src/nn/initializers.py):
# uniform 的实现:随机均匀采样后乘以 scale return random.uniform(key, shape, dtype) * jnp.array(scale, dtype) # normal 的实现:随机正态采样后乘以 stddev return random.normal(key, shape, dtype) * jnp.array(stddev, dtype) # truncated_normal 的实现:先按 lower/upper 截断,再乘以 stddev return random.truncated_normal( key, lower, upper, shape, dtype) * jnp.array(stddev, dtype)使用truncated_normal时需要注意两点(源码 docstring 中明确强调):
- 不做方差修正:与
variance_scaling系列不同,truncated_normal不会对截断造成的方差损失做修正(即没有除以截断标准差常数),需要用户自行通过stddev参数补偿; - 截断先于缩放:
lower/upper在输出乘以stddev之前应用,因此最终取值范围是lower*stddev到upper*stddev,默认(-2, 2)的截断边界对应"两倍标准差"经验法则。
这一组初始化器适合对分布形态有明确要求的场景,但不像方差缩放族那样随权重张量形状自适应,因此在实际网络初始化中更常用的是下文基于 fan 值计算的初始化器。
核心机制:variance_scaling 与 fan 值计算
variance_scaling是jax.nn.initializers中最重要的自适应初始化器,Glorot/Xavier、He/Kaiming、LeCun 等经典家族都是它的特例。其签名为(jax/_src/nn/initializers.py):
variance_scaling(scale, mode, distribution, in_axis=-2, out_axis=-1, batch_axis=(), dtype=jnp.float_)三个必选参数的含义:
scale:正浮点缩放因子;mode:"fan_in"、"fan_out"、"fan_avg"三者之一,决定方差的分母取输入单元数、输出单元数还是两者平均;distribution:"truncated_normal"、"normal"、"uniform"三者之一。
其数学定义是:无论采用哪种分布,采样值的标准差(截断后,如适用)均为sqrt(scale / n),其中n依mode取 fan_in、fan_out 或二者的平均值。在truncated_normal模式下,采样绝对值会在缩放前截断于 2 个标准差处。
fan 值如何计算:_compute_fans
fan_in/fan_out 由私有函数_compute_fans(shape, in_axis, out_axis, batch_axis)计算(jax/_src/nn/initializers.py):
- 权重张量必须至少 2 维,否则抛出
ValueError("Can't compute input and output sizes of a ... weights tensor. Must be at least 2D."); in_axis、out_axis可以是单个整数轴,也可以是轴序列(此时取各轴尺寸之积);- 不属于
in_axis、out_axis、batch_axis的轴被视作卷积的感受野(receptive field),即核空间维度; - 计算公式为
fan_in = in_size * receptive_field_size,fan_out = out_size * receptive_field_size。
这套设计使variance_scaling能同时适用于全连接层(2D 权重)和卷积层(4D/5D 卷积核)。测试 tests/nn_test.py 中的testVarianceScalingMultiAxis与testVarianceScalingBatchAxis分别验证了多轴输入/输出和 batch 轴忽略的配置,例如:
# 将 0、1 轴视为输入维,-2、-1 轴视为输出维 initializer = nn.initializers.variance_scaling( scale=1.0, mode='fan_avg', distribution='truncated_normal', in_axis=(0, 1), out_axis=(-2, -1)) # 将轴 1 视为 batch 轴忽略 initializer = nn.initializers.variance_scaling( scale=1.0, mode='fan_avg', distribution='truncated_normal', in_axis=0, out_axis=(2, 3), batch_axis=1)三种分布的具体采样方式
根据distribution取值,实现走不同分支(jax/_src/nn/initializers.py):
truncated_normal:实类型时先除以常数0.87962566103423978(标准正态截断到(-2, 2)后的真实标准差,用于补偿截断造成的方差损失)再乘以sqrt(variance);复类型时对应常数0.95311164380491208;normal:直接random.normal(...) * sqrt(variance);uniform:实类型时采样random.uniform(key, shape, dtype, -1)(即[-1, 1))后乘以sqrt(3 * variance);复类型时使用盘状均匀采样。
经典家族:Glorot/Xavier、He/Kaiming、LeCun
模块为经典论文中的初始化方案提供了开箱即用的封装,全部是variance_scaling的参数特化:
| 初始化器 | 别名 | 对应variance_scaling参数 | 出处思想 |
|---|---|---|---|
glorot_uniform() | xavier_uniform | scale=1.0, mode="fan_avg", distribution="uniform" | Glorot & Bengio,Xavier 均匀 |
glorot_normal() | xavier_normal | scale=1.0, mode="fan_avg", distribution="truncated_normal" | Glorot & Bengio,Xavier 正态 |
lecun_uniform() | — | scale=1.0, mode="fan_in", distribution="uniform" | LeCun 等,fan_in 均匀 |
lecun_normal() | — | scale=1.0, mode="fan_in", distribution="truncated_normal" | LeCun 等,fan_in 截断正态 |
he_uniform() | kaiming_uniform | scale=2.0, mode="fan_in", distribution="uniform" | He 等(Kaiming),ReLU 友好 |
he_normal() | kaiming_normal | scale=2.0, mode="fan_in", distribution="truncated_normal" | He 等(Kaiming),ReLU 友好 |
实现上一行即可完成特化(如 jax/_src/nn/initializers.py):
return variance_scaling(1.0, "fan_avg", "uniform", in_axis=in_axis, out_axis=out_axis, batch_axis=batch_axis, dtype=dtype)六个家族初始化器的共同特征:
- 都支持
in_axis、out_axis、batch_axis三个轴参数(默认-2、-1、()),从而天然兼容卷积核形状; glorot_*采用fan_avg,方差更均衡,适合对称激活(如 tanh、sigmoid)场景;he_*与lecun_*采用fan_in,其中he_*的scale=2额外补偿了 ReLU 类激活带来的方差减半效应,是深度 ReLU 网络的常用默认选择;- 各文档字符串中都附带可复现示例,例如:
>>> initializer = jax.nn.initializers.glorot_uniform() >>> initializer(jax.random.key(42), (2, 3), jnp.float32) Array([[ 0.50350785, 0.8088631 , 0.81566876], [-0.6393332 , -0.6865721 , 0.11003882]], dtype=float32)正交初始化器:orthogonal 与 delta_orthogonal
orthogonal
orthogonal(scale=1.0, column_axis=-1, dtype)返回均匀分布的正交矩阵(jax/_src/nn/initializers.py):
- 形状要求:至少 2 维,否则抛
ValueError;非方阵时,较小的一侧保证行或列正交; - 实现流程:生成标准正态矩阵 →
jnp.linalg.qr分解 → 用jnp.sign(jnp.diag(R))修正 Q 的符号(保证正交矩阵的均匀分布性质)→ 重排轴后乘以scale; - column_axis:指定需要正交的列所在的轴,默认最后一轴。
delta_orthogonal
delta_orthogonal(scale=1.0, column_axis=-1, dtype)是为卷积设计的"Delta 正交核"初始化器(jax/_src/nn/initializers.py),出自 2018 年相关研究工作:
- 形状要求:必须是3D、4D 或 5D(对应 1D/2D/3D 卷积核),且要求
shape[-1] >= shape[-2](即fan_in <= fan_out),否则抛ValueError; - 实现方式:先通过
orthogonal生成一个正交矩阵放在核的中心位置(各空间维度取(k-1)//2),其余位置全部置零,例如 3D 核(3, 3, 3)的结果只在中间切片[1, ...]处非零:
>>> initializer = jax.nn.initializers.delta_orthogonal() >>> initializer(jax.random.key(42), (3, 3, 3), jnp.float32) Array([[[ 0. , 0. , 0. ], ... [[ 0.27858758, -0.7949833 , -0.53887904], [ 0.9120717 , 0.04322892, 0.40774566], [-0.30085585, -0.6050892 , 0.73712474]], ... [[ 0. , 0. , 0. ], ...]], dtype=float32)该初始化器适合需要保持卷积层通道间正交性的网络结构(如某些残差卷积架构)。
复数支持与 API 导出细节
复数 dtype 的完整支持
variance_scaling在distribution="normal"、"truncated_normal"、"uniform"三种模式下都支持复数 dtype:
- 均匀分布使用私有函数
_complex_uniform(jax/_src/nn/initializers.py):在复平面单位圆盘内均匀采样,零均值、单位方差——实现方式是半径r = sqrt(2 * U)、角度theta = 2*pi*U的极坐标采样; - 截断正态使用
_complex_truncated_normal(jax/_src/nn/initializers.py):模长截断到upper、截断前方差为 1 的复平面中心正态分布; - 实数分支与复数分支以
jnp.issubdtype(dtype, jnp.floating)区分。
公开 API 的导出方式
模块采用"公开包薄转发、实现集中在_src"的布局:jax.nn.initializers的公开符号由 jax/nn/initializers.py 以import <name> as <name>形式从 jax/_src/nn/initializers.py 转发导出,并通过set_module('jax.nn.initializers')将内部函数的__module__重写为公开路径。
值得注意的是,公开 API 实际上比 API 文档的 autosummary 列表更丰富:除文档列出的 15 个初始化器外,还额外导出了:
- 别名:
kaiming_normal(=he_normal)、kaiming_uniform(=he_uniform)、xavier_normal(=glorot_normal)、xavier_uniform(=glorot_uniform); - 类型:
Initializer(Protocol 类型,可用于类型标注与isinstance检查)。
初始化器的验证与测试
仓库的测试套件 tests/nn_test.py 为初始化器的正确性提供了系统验证,可作为读者理解各初始化器行为边界的参考:
- 测试形状与 dtype 矩阵:
ALL_SHAPES覆盖 1D 到 4D 共 7 种形状,initializer_record为每个初始化器声明允许的维度范围与 dtype 类别(如orthogonal仅限 2D、delta_orthogonal仅限 4D 测试、truncated_normal仅限浮点); testInitializer/testInitializerProvider:验证模块级函数与工厂式两种 API 形态下,输出形状与 dtype 均正确(dtype 会经过canonicalize_dtype规范化);testVarianceScalingError:验证 1D 张量调用variance_scaling会抛出含明确错误信息的ValueError;testAccidentalUpcasting:验证uniform/normal/truncated_normal的标量尺度参数(即使传入float32数组)不会意外把bfloat16权重提升为更高精度。
总结与选型建议
综合全模块,一个实用的选型思路是:
- 偏置与固定值:用
zeros、constant; - 通用全连接/卷积默认:ReLU 系网络优先
he_normal/he_uniform,对称激活网络优先glorot_normal/glorot_uniform; - 需要显式控制分布与 fan 语义:直接使用
variance_scaling(scale, mode, distribution)自由组合,并通过in_axis/out_axis/batch_axis适配任意权重布局; - RNN 与残差卷积的稳定性需求:考虑
orthogonal与delta_orthogonal; - 复值网络:
variance_scaling的三种分布均原生支持复数 dtype。
所有初始化器都遵循统一的(key, shape, dtype)约定,因此可以无差别地传入任意接受initializer(key, shape, dtype)的层构建代码,且天然兼容jax.jit与jax.grad等 JAX 变换。若需深入原理,建议对照阅读 jax/_src/nn/initializers.py(实现)、jax/nn/initializers.py(公开 API 转发)与 tests/nn_test.py(验证用例)。
- 机器学习
- 深度学习
【免费下载链接】jax
Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more
相关推荐
JAX 权重初始化器完全指南:深入解析 `jax.nn.initializers` 的签名、分布与源码实现
JAX 权重初始化器完全指南:深入解析 jax.nn.initializers 的签名、分布与源码实现 jax.nn.initializers 是 JAX 神经
人工智能机器学习深度学习编译器高性能计算如何快速掌握LitGPT模型参数初始化:完整指南与最佳实践
如何快速掌握LitGPT模型参数初始化:完整指南与最佳实践 LitGPT是一个强大的开源项目,允许用户在自己的数据上预训练、微调20多种大型语言模型(LLMs)
大模型预训练微调模型推理服务Candle初始化策略:权重初始化方法与最佳实践
Candle初始化策略:权重初始化方法与最佳实践 引言 权重初始化是深度学习模型训练中的关键环节,直接影响模型的收敛速度和最终性能。在Rust机器学习框架Can
人工智能深度学习机器学习大模型本地部署
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考