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

ResNet50在Statoil冰山分类任务中输入形状不匹配问题求助

解决Statoil冰山分类挑战赛中ResNet50的输入形状不兼容问题

问题根源

你的错误来自两个核心矛盾:

  1. 定义ResNet50时指定了input_shape=(200,200,3),但训练数据是75x75的3通道图像,形状不匹配导致输入兼容错误。
  2. 直接尝试reshape到200x200失败,因为reshape要求总元素数不变:每个样本的单通道数据是7575=5625个元素,200200=40000,两者数量不相等,无法直接reshape。

解决方案一:修改ResNet50的输入形状适配现有数据

ResNet50支持自定义输入形状(只要尺寸能通过模型的卷积层),直接将其输入形状改为和VGG16/MobileNet一致的(75,75,3)即可,不需要修改数据:

修改factory字典中的ResNet50定义:

factory = {
    'vgg16': lambda: VGG16(include_top=False, input_shape=(75, 75, 3), weights=vgg16_fl),
    'mobilenetv2': lambda: MobileNet(include_top=False, input_shape=(75, 75, 3)),
    # 修改此处的input_shape为75x75x3
    'resnet50': lambda: ResNet50(include_top=False, input_shape=(75, 75, 3), weights='imagenet'),
}

修改后直接运行原有训练代码即可,无需调整数据部分。


解决方案二:缩放原始数据到ResNet50要求的200x200

如果你希望保留ResNet50的200x200输入(贴合预训练权重的原始适配尺寸),需要用图像插值缩放的方式将75x75的图像放大到200x200,而非使用reshape。

以下是TensorFlow版本的实现:

import tensorflow as tf

data = pd.read_json('/content/drive/MyDrive/iceberg/train.json')
b1 = np.array(data["band_1"].values.tolist()).reshape(-1, 75, 75, 1)
b2 = np.array(data["band_2"].values.tolist()).reshape(-1, 75, 75, 1)
b3 = b1 + b2

# 用插值法缩放每个通道到200x200
b1_resized = tf.image.resize(b1, (200, 200)).numpy()
b2_resized = tf.image.resize(b2, (200, 200)).numpy()
b3_resized = tf.image.resize(b3, (200, 200)).numpy()

X = np.concatenate([b1_resized, b2_resized, b3_resized], axis=3)
y = np.array(data['is_iceberg'])
angle = np.array(pd.to_numeric(data['inc_angle'], errors='coerce').fillna(0))

# ResNet50的input_shape保持(200,200,3)即可正常训练
model = get_model('resnet50', train_base=False, use_angle=True)
model.compile(loss='binary_crossentropy', optimizer=Adam(lr=1e-3), metrics=['accuracy'])
history = model.fit([X, angle], y, shuffle=True, verbose=1, epochs=5)

关键说明

  • reshape仅用于调整数组维度,不改变元素总数和排列,不能用来改变图像的空间尺寸;改变图像尺寸必须用插值缩放方法。
  • 方案一更简洁高效,避免额外计算;方案二更贴合预训练权重的原始输入,可能在迁移学习中获得更好效果,可根据需求选择。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 12:24:19