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

如何将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 22:09:03