在Google Colab训练LSTM模型遇Model.fit嵌套tf.function错误求助
问题:Keras LSTM训练时触发RuntimeError错误
错误提示
RuntimeError: Detected a call to `Model.fit` inside a `tf.function`. `Model.fit` is a high-level endpoint that manages its own `tf.function`. Please move the call to `Model.fit` outside of all enclosing `tf.function`s. Note that you can call a `Model` directly on `Tensor`s inside a `tf.function` like: `model(x)`.
原始代码
# Importing required libraries import pandas as pd import numpy as np import seaborn as sns import matplotlib.pyplot as plt from tensorflow.keras.models import Sequential from tensorflow.keras.layers import LSTM, Dropout, Dense from tensorflow.keras.callbacks import EarlyStopping # Define the LSTM model def create_lstm_model(input_size, output_size, lstm_layer_sizes, dropout_rates): lstm_model = Sequential() for size, rate in zip(lstm_layer_sizes, dropout_rates): lstm_model.add(LSTM(units=size, return_sequences=True)) lstm_model.add(Dropout(rate=rate)) lstm_model.add(Dense(units=output_size)) return lstm_model # Set the parameters input_size = 6 output_size = 3 unit = 'unit1' outage = 'moh' lstm_layer_sizes = (64,128,256,128,64) dropout_rates = (0.05,0.05,0.05,0.05,0.05) # Prepare the data (omitting data retrieval steps for brevity) y = kinerja_df_extended_nanremoved_standardized[f'{unit}_{outage}s'] current_dates = kinerja_df_extended_nanremoved_standardized['date'] x = np.array([current_dates[i:i+input_size] for i in range(len(current_dates)-input_size+1)]) y = np.array([y[i:i+output_size] for i in range(len(y)-output_size+1)]) # Instantiate and compile the model lstm_model = create_lstm_model(input_size=input_size, output_size=output_size, lstm_layer_sizes=lstm_layer_sizes, dropout_rates=dropout_rates) lstm_model.compile(optimizer='adam', loss='mean_squared_error') # The following line causes the error history = lstm_model.fit(x=x, y=y, batch_size=1, epochs=128, validation_split=0.1, shuffle=False) # Plot the training and validation loss plt.plot(history.history['loss'], label='Training Loss') plt.plot(history.history['val_loss'], label='Validation Loss') plt.legend() plt.show()
修复方案
1. 确保Model.fit未被包裹在tf.function中
检查你的Colab笔记本,确认调用lstm_model.fit的代码没有处于任何被@tf.function装饰的函数或代码块内部。如果之前误将fit调用放在了tf.function里面,直接将其移到全局代码区域即可。
2. 修正LSTM输入数据的维度与类型
LSTM模型要求输入张量形状为(样本数, 时间步长, 特征数),且输入必须为数值型数据,无法直接处理日期类型:
- 将日期转换为数值特征(比如时间戳、日期编码)
- 扩展输入数据为三维数组
- 确保输入与输出的样本数量匹配
修改后的数据准备代码:
# 修正数据准备步骤 # 将日期转换为时间戳数值(或其他数值化编码方式) kinerja_df_extended_nanremoved_standardized['timestamp'] = kinerja_df_extended_nanremoved_standardized['date'].apply(lambda x: x.timestamp()) # 选择用于预测的特征(可添加设备运行参数等其他相关特征) feature_cols = ['timestamp'] X_data = kinerja_df_extended_nanremoved_standardized[feature_cols].values # 构造LSTM输入:形状为(样本数, input_size, 特征数) # 确保x和y的样本数一致 total_samples = len(X_data) - input_size - output_size + 1 x = np.array([X_data[i:i+input_size] for i in range(total_samples)]) y = np.array([y[i+input_size:i+input_size+output_size] for i in range(total_samples)])
3. 修正LSTM模型结构
原始模型中所有LSTM层都设置了return_sequences=True,导致最后一层LSTM输出序列数据,与Dense层的输出形状不匹配。修改模型定义,将最后一个LSTM层的return_sequences设为False:
def create_lstm_model(input_size, output_size, lstm_layer_sizes, dropout_rates): lstm_model = Sequential() for idx, (size, rate) in enumerate(zip(lstm_layer_sizes, dropout_rates)): # 最后一个LSTM层不返回序列 if idx == len(lstm_layer_sizes) - 1: lstm_model.add(LSTM(units=size, return_sequences=False)) else: lstm_model.add(LSTM(units=size, return_sequences=True)) lstm_model.add(Dropout(rate=rate)) lstm_model.add(Dense(units=output_size)) return lstm_model
4. 验证并重新训练
完成上述修改后,重新运行代码即可正常训练LSTM模型。
内容的提问来源于stack exchange,提问作者Muhammad Ikhwan Perwira
相关产品推荐
相关产品推荐

