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

如何解读Keras中GRU层get_weights()方法的返回权重结果?

Keras GRU层get_weights()返回结果说明

GRU层返回的权重数组结构和门控计算逻辑直接对应,和SimpleRNN的权重拆分逻辑有区别,具体每一项含义如下:

  • 第一个返回数组:输入投影权重,形状为(输入特征维度, 3 * GRU单元数)
    示例代码中输入特征维度为1、GRU单元数为2,因此数组形状为(1,6),和打印结果完全匹配。数组列按顺序均分为3段,每段长度等于GRU单元数:
    • 前2列:更新门(update gate)对应的输入权重
    • 中间2列:重置门(reset gate)对应的输入权重
    • 最后2列:候选隐藏状态计算环节对应的输入权重
  • 第二个返回数组:循环状态权重,形状为(GRU单元数, 3 * GRU单元数)
    示例中形状为(2,6),列的拆分规则和输入投影权重完全一致,分别对应更新门、重置门、候选隐藏状态计算时,乘以上一时刻隐藏状态的权重参数。
  • 第三个返回数组:偏置参数,形状为(2, 3 * GRU单元数)
    示例中形状为(2,6),全0是Keras默认的参数初始化状态。GRU的偏置被拆为两部分存储:第一行是输入侧偏置,第二行是循环状态侧偏置,实际前向计算时会将两行相加后参与运算,并非两组独立偏置。

关于"GRU包含输出门"的命名歧义说明

查阅源码时看到的"GRU包含更新门、重置门、输出门共3个门"的描述属于命名习惯导致的误解:
标准GRU结构仅包含更新门、重置门两个门控单元,不存在LSTM结构中独立的输出门。Keras GRU源码里标注的"output gate"实际就是前述的候选隐藏状态计算环节,只是开发时沿用了LSTM实现的变量命名习惯,并没有给GRU额外新增输出门结构,和经典GRU的公式定义完全一致。

示例实现代码如下:

# 参考SimpleRNN公开示例修改
from pandas import read_csv
import numpy as np
from keras.models import Sequential
from keras.layers import Dense, SimpleRNN, GRU
from sklearn.preprocessing import MinMaxScaler
from sklearn.metrics import mean_squared_error
import math
import matplotlib.pyplot as plt

model = Sequential()
model.add(GRU(units = 2, input_shape = (3,1), activation = 'linear'))
model.add(Dense(units = 1, activation = 'linear'))
model.compile(loss = 'mean_squared_error', optimizer = 'adam')

initial_weights = model.layers[0].get_weights()
print("Shape = ",initial_weights)

运行后打印的初始权重结果:

Shape =  [array([[-0.64266175, -0.0870676 , -0.25356603, -0.03685969,  0.22260845,
        -0.04923642]], dtype=float32), array([[ 0.01929092, -0.4932567 ,  0.3723044 , -0.6559699 , -0.33790302,
         0.27062896],
       [-0.4214194 ,  0.46456426,  0.27233726, -0.00461334, -0.6533575 ,
        -0.32483965]], dtype=float32), array([[0., 0., 0., 0., 0., 0.],
       [0., 0., 0., 0., 0., 0.]], dtype=float32)]

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 20:30:53