如何将字典类数据结构传入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
相关产品推荐
相关产品推荐

