如何加载输入形状含None的3D CNN预训练模型以实现滑动窗口检测?
解决TensorFlow 3D CNN滑动窗口加载模型时通道维度未定义问题
错误原因
训练阶段构建模型时,输入形状指定为(128, None, None, None, 1),将batch维度固定为128,导致模型保存后,加载时无法正确识别固定的通道维度(轴-1),进而触发通道维度为None的错误。
解决方案
方案1:重新训练时修正输入形状定义
修改训练代码中的model.build语句,将batch维度设为可变(None),同时明确固定通道维度为1:
model.build(input_shape=(None, None, None, None, 1))
重新训练并保存模型后,滑动窗口加载时即可正常识别输入形状,无需额外修改加载代码。
方案2:针对已训练完成的模型,加载后重构输入层
如果不想重新训练,可在滑动窗口脚本中加载模型后,手动重构输入层以明确通道维度:
# 加载原模型 model = tf.keras.models.load_model( '/project/cleancode.keras', compile=False, custom_objects={"BalancedAccuracy": BalancedAccuracy(), "F1Score": tfa.metrics.F1Score(num_classes=1, threshold=0.5)} ) # 重新定义输入层:空间维度可变,通道维度固定为1 new_input = tf.keras.Input(shape=(None, None, None, 1)) # 传递新输入到原模型获取输出 new_output = model(new_input) # 创建新的可接受可变尺寸输入的模型 new_model = tf.keras.Model(inputs=new_input, outputs=new_output) # 后续滑动窗口检测使用new_model即可
额外注意事项
确保滑动窗口提取的子体素数据形状符合(batch_size, depth, height, width, 1),通道维度必须明确为1,避免出现维度缺失的情况。
内容的提问来源于stack exchange,提问作者Mustafa Mahmood
相关产品推荐
相关产品推荐

