TensorFlow自定义数据集如何给CNN图像输入添加元数据作为额外输入
解决方案
一、数据集输入逻辑修改
你已经完成了模型结构的调整,只需要补全数据集初始化逻辑,并调整数据处理函数的返回格式即可适配双输入结构:
- 补全初始数据集创建与拆分逻辑,你之前的代码缺少了从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)
- 调整数据处理函数的返回格式,修正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
- 调整数据集的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
相关产品推荐
相关产品推荐

