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

TensorFlow自定义数据集如何给CNN图像输入添加元数据作为额外输入

解决方案

一、数据集输入逻辑修改

你已经完成了模型结构的调整,只需要补全数据集初始化逻辑,并调整数据处理函数的返回格式即可适配双输入结构:

  1. 补全初始数据集创建与拆分逻辑,你之前的代码缺少了从CSV读取的数组生成tf.data.Dataset以及拆分的步骤,补充代码如下:
# 从读取到的数组创建全量数据集
ds_full = tf.data.Dataset.from_tensor_slices((file_paths, labels, usages, completions, heights, constructions))
# 拆分数据集
ds_full = ds_full.shuffle(dataset_size, reshuffle_each_iteration=False)
ds_train = ds_full.take(train_size)
ds_remaining = ds_full.skip(train_size)
ds_val = ds_remaining.take(val_size)
ds_test = ds_remaining.skip(val_size)
  1. 调整数据处理函数的返回格式,修正augment函数的参数名错误,将四个元数据拼接为4维张量,适配模型输入要求:
# FUNCTION TO READ AND NORMALIZE THE IMAGES
def read_image(image_file, label, usg, com, hei, con):
    image = tf.io.read_file(image_file)
    image = tf.image.decode_jpeg(image, channels=3)
    image = tf.image.resize(image, (img_width, img_height))
    image = tf.cast(image, tf.float32) / 255.0
    # 拼接四个元数据为shape=(4,)的张量,和你定义的input_dat的shape=(4,)完全匹配
    meta = tf.stack([usg, com, hei, con], axis=-1)
    meta = tf.cast(meta, tf.float32)
    return (image, meta), label

# FUNCTION FOR DATA AUGMENTATION
def augment(image_meta, label):
    image, meta = image_meta
    if tf.random.uniform((), minval=0, maxval=1) < 0.1:
        image = tf.tile(tf.image.rgb_to_grayscale(image), [1, 1, 3])
    image = tf.image.random_brightness(image, max_delta=0.25)
    image = tf.image.random_contrast(image, lower=0.75, upper=1.25)
    image = tf.image.random_saturation(image, lower=0.75, upper=1.25)
    image = tf.image.random_flip_left_right(image)
    return (image, meta), label
  1. 调整数据集的map流程,你原本的训练代码不需要做任何修改,直接可以正常运行:
# SETUP FOR TRAINING, VALIDATION & TEST DATASET
ds_train = ds_train.map(read_image, num_parallel_calls=autotune)
ds_train = ds_train.cache()
ds_train = ds_train.map(augment, num_parallel_calls=autotune)
ds_train = ds_train.batch(batch_size)
ds_train = ds_train.prefetch(autotune)

ds_val = ds_val.map(read_image, num_parallel_calls=autotune)
ds_val = ds_val.batch(batch_size)
ds_val = ds_val.prefetch(autotune)

ds_test = ds_test.map(read_image, num_parallel_calls=autotune)
ds_test = ds_test.batch(batch_size)
ds_test = ds_test.prefetch(autotune)

补充说明:你定义的input_dat = keras.Input(shape=(4,))完全正确,不需要修改。另外建议你把输出层的激活函数从sigmoid改为softmax,更适配多分类场景下的SparseCategoricalCrossentropy损失函数。

二、测试代码适配修改

因为我们已经把ds_test的输出格式调整为((图像batch, 元数据batch), 标签batch),你原来的测试代码几乎不需要修改,直接可以正常运行,调整了转义字符后的代码如下:

y_true = []
y_pred = []

for x, y in ds_test:
    y_true.append(y)
    predicts = model.predict(x)                                         
    y_pred.append(np.argmax(predicts, axis=-1))                                              
true = tf.concat([item for item in y_true], axis=0)
pred = tf.concat([item for item in y_pred], axis=0)

cm = confusion_matrix(true, pred)                                           
testacc = np.trace(cm) / float(np.sum(cm))                                  
cm = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]
fig, ax = plt.subplots(figsize=(10, 10))
color = sns.light_palette("seagreen", as_cmap=False)
sns.heatmap(cm, annot=True, square=True, cmap=color, fmt=".3f",
            linewidths=0.6, linecolor='k', cbar_kws={"shrink": 0.8})
plt.yticks(rotation=0)
plt.xlabel('\nPredicted Labels', fontsize=18)
plt.ylabel('True Labels\n', fontsize=18)
plt.title('Multiclass Model - Confusion Matrix (Test Step)\n', fontsize=24)
plt.text(10, 1.1, 'Accuracy = {:0.4f}'.format(testacc), fontsize=20)
ax.axhline(y=8, color='k', linewidth=1.5)                                   
ax.axvline(x=8, color='k', linewidth=1.5)
plt.show()
print('\naccuracy: {:0.4f}'.format(testacc))

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 19:54:02