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

如何将字典类数据结构传入LSTM层的call()方法并转换为Tensor格式

实现方案

1. 有序字典转固定张量

Keras层的输入仅支持张量类型,无法直接传入Python字典对象,所以需要先提取OrderedDict中的有效数据转换为TensorFlow张量。你的od全程固定不需要更新,直接使用tf.constant定义即可:

import tensorflow as tf
from collections import OrderedDict

od = OrderedDict()
od['a'] = 1
od['b'] = 2
od['c'] = 3
od['d'] = 4

# 提取字典值转为一维张量,可根据后续逻辑调整维度
od_values = list(od.values())
od_tensor = tf.constant(od_values, dtype=tf.float32)
# 扩展维度适配输入格式,变成[1, 1, 4]可直接和输入样本广播
od_tensor = tf.expand_dims(tf.expand_dims(od_tensor, 0), 0)

2. 自定义LSTM层实现

继承原生LSTM层实现自定义类,保留你需要的打印逻辑即可:

class CustomLSTM(tf.keras.layers.LSTM):
    def call(self, inputs, mask=None, training=None, initial_state=None):
        # inputs为传入的输入列表,第0位是原始时序输入,第1位是od张量
        tf.print("inputs[0]: ", inputs[0])
        tf.print("inputs[1]: ", inputs[1])
        inputs_0 = inputs[0]
        self._maybe_reset_cell_dropout_mask(self.cell)
        return super().call(
            inputs_0, mask=mask, training=training, initial_state=initial_state
        )

注:用tf.print替代原生print可以兼容图执行模式,训练过程中也能正常输出内容。

3. 模型构建与训练适配

根据你的使用场景二选一即可:

方案A:od内置在模型中(推荐,适合od固定不变的场景)

直接将od张量和原始输入组合传入自定义LSTM,不需要修改fit的传参逻辑:

inp = Input(shape = (1, look_back), name = "Input")
# 组合原始输入和固定od张量传入层
x = CustomLSTM(4)([inp, od_tensor])
out = Dense(1)(x)

model = Model(inp, out)
model.compile(loss='mean_squared_error', optimizer='adam')
# 训练逻辑和原来完全一致
model.fit(trainX, trainY, epochs=100, batch_size=1, verbose=2)

方案B:od作为独立输入节点(适合后续需要动态修改od的场景)

新增od对应的输入层,训练时传入和样本batch匹配的od张量:

inp = Input(shape = (1, look_back), name = "sequence_input")
od_inp = Input(shape = (1, 4), name = "od_input")
x = CustomLSTM(4)([inp, od_inp])
out = Dense(1)(x)

model = Model([inp, od_inp], out)
model.compile(loss='mean_squared_error', optimizer='adam')

# 将od张量扩展到和训练/测试集相同的batch大小
train_od = tf.repeat(od_tensor, repeats=trainX.shape[0], axis=0)
test_od = tf.repeat(od_tensor, repeats=testX.shape[0], axis=0)

# 训练时传入两个输入
model.fit([trainX, train_od], trainY, epochs=100, batch_size=1, verbose=2)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 10:54:02