You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何通过Flax的model.init将GRUCell隐藏状态设为可学习参数

解决方案

一、让model.init包含可学习初始隐藏状态

要让初始隐藏状态成为可学习参数,直接在Module中通过self.param定义即可,同时可指定自定义初始化函数。修改后的代码如下:

import jax.numpy as np
from jax import random
import flax.linen as nn
from jax.nn import initializers

class RNN(nn.Module):
    n_RNN_units: int
    # 可选:允许传入初始状态的自定义初始化函数
    carry_init_fn: callable = initializers.zeros

    @nn.compact
    def __call__(self, inputs):
        # 将初始隐藏状态注册为可学习参数
        initial_carry = self.param('initial_carry', self.carry_init_fn, (self.n_RNN_units,))
        carry = initial_carry
        outputs = []
        # 遍历输入序列(适配(seq_len, data_dim)或(batch, seq_len, data_dim)格式)
        for x in np.moveaxis(inputs, 0, -1):
            carry, out = nn.GRUCell()(carry, x)
            outputs.append(out)
        return carry, np.stack(outputs, axis=0)

# 实例化模型,可自定义初始状态初始化逻辑(比如正态分布初始化)
n_RNN_units = 200
model = RNN(n_RNN_units=n_RNN_units, carry_init_fn=initializers.normal(stddev=0.1))

# 初始化参数,此时params会包含GRU权重和可学习初始carry
data_dim = 20
dummy_inputs = np.empty((10, data_dim))  # 示例序列长度为10
params = model.init(random.PRNGKey(1), dummy_inputs)

# 查看参数结构,会看到'GRUCell_0'和'initial_carry'
print(params['params'].keys())

关键说明

  • self.param('initial_carry', ...)会将初始隐藏状态注册为模型参数,model.init会自动将其纳入参数字典。
  • 可通过类参数carry_init_fn指定初始化逻辑,支持Jax内置初始化器(如initializers.ones)或自定义函数(输入rng和shape返回对应数组)。
  • 修改__call__接口为直接接收输入序列,内部处理初始状态,更贴合常规使用流程;若需保留原carry输入接口,可添加逻辑:当carry为None时使用可学习初始状态。

二、不使用model.init的替代方案

若不想依赖model.init,可手动构建参数字典,合并GRU权重与手动初始化的可学习初始状态:

import jax.numpy as np
from jax import random
import flax.linen as nn
from jax.nn import initializers

class RNN(nn.Module):
    n_RNN_units: int

    @nn.compact
    def __call__(self, carry, inputs):
        carry, outputs = nn.GRUCell()(carry, inputs)
        return carry, outputs

n_RNN_units = 200
data_dim = 20
model = RNN(n_RNN_units=n_RNN_units)

# 1. 初始化GRU的权重参数
dummy_carry = np.empty((n_RNN_units,))
dummy_inputs = np.empty((data_dim,))
gru_params = model.init(random.PRNGKey(1), dummy_carry, dummy_inputs)

# 2. 手动初始化可学习的初始隐藏状态
rng = random.PRNGKey(2)
carry_init_fn = initializers.normal(stddev=0.1)
initial_carry = carry_init_fn(rng, (n_RNN_units,))

# 3. 合并参数字典,保持Flax参数结构格式
params = {
    'params': {
        **gru_params['params'],
        'initial_carry': initial_carry
    }
}

# 使用时取出初始carry并运行模型
initial_carry = params['params']['initial_carry']
carry, outputs = model.apply(params, initial_carry, dummy_inputs)

关键说明

  • 先通过model.init获取GRU核心权重,再单独初始化可学习初始状态。
  • 合并参数时需遵循Flax的嵌套结构(参数都在'params'键下)。
  • 该方式灵活性更高,适合需要对参数做额外自定义处理的场景。

内容的提问来源于stack exchange,提问作者Jabby

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.21 05:06:28