如何将image_dataset_from_directory生成的BatchDataset标签转为浮点型?
解决方案
首先明确tf.keras.utils.image_dataset_from_directory默认返回的标签为int32类型的分类索引,我们可以通过tf.data.Dataset.map()方法对数据集的(图像, 标签)对做批量转换,适配回归任务的浮点标签需求。
步骤1:准备标签映射规则
首先拿到数据集的类索引对应关系,再构建类索引到实际浮点测量值的映射:
import tensorflow as tf # 创建原始BatchDataset,注意保留返回的数据集对象以获取类名对应关系 raw_dataset = tf.keras.utils.image_dataset_from_directory( "./你的数据集根目录", image_size=(224, 224), # 替换为你的图像尺寸 batch_size=32, shuffle=True ) # 获取类名到索引的映射关系,例如class_names[0]对应索引0的子文件夹名 class_names = raw_dataset.class_names # 构建索引到浮点测量值的映射,按你实际的对应关系填写 index_to_measure = { 0: 0.58, 1: 1.23, 2: 2.76, # 补全所有类别的对应浮点值 }
步骤2:构建适配TensorFlow图模式的查找表
直接使用Python字典查询无法适配TensorFlow的图执行模式,需要用StaticHashTable实现查询逻辑:
# 生成哈希表的键值对张量 keys = tf.constant(list(index_to_measure.keys()), dtype=tf.int32) values = tf.constant(list(index_to_measure.values()), dtype=tf.float32) # 初始化静态哈希表 label_table = tf.lookup.StaticHashTable( initializer=tf.lookup.KeyValueTensorInitializer(keys, values), default_value=tf.constant(-1.0, dtype=tf.float32) # 未匹配到索引时的默认返回值 )
步骤3:对数据集应用标签转换
定义转换函数并作用到原始数据集上,同时开启并行处理提升加载效率:
def process_label(image, label): # 查表将int类型索引转换为浮点测量值 float_label = label_table.lookup(label) # 如果回归模型要求输出维度为(batch, 1),可以加下面这行扩展维度 # float_label = tf.expand_dims(float_label, axis=-1) return image, float_label # 应用转换得到适配回归任务的数据集 train_dataset = raw_dataset.map(process_label, num_parallel_calls=tf.data.AUTOTUNE)
简化场景(无需映射直接转类型)
如果你的原始类索引本身就是需要用到的浮点值,不需要额外映射,可以直接做类型转换,省去哈希表步骤:
def process_label(image, label): return image, tf.cast(label, tf.float32) train_dataset = raw_dataset.map(process_label, num_parallel_calls=tf.data.AUTOTUNE)
验证转换结果
可以取一个批次验证标签转换是否符合预期:
for img_batch, label_batch in train_dataset.take(1): print("标签数据类型:", label_batch.dtype) print("前5个标签值:", label_batch[:5])
转换完成的数据集可以直接传入model.fit()训练回归模型,使用MSE、MAE等回归损失函数即可。
内容的提问来源于stack exchange,提问作者pdr0663
相关产品推荐
相关产品推荐

