在GCP AI平台使用tf.dataset时如何实现One hot encoding
问题:TensorFlow Dataset中对分类特征进行One-Hot编码(GCP AI平台场景)
我在GCP AI平台运行TensorFlow模型时,因为数据集太大没法全部放进内存,所以用下面的代码把数据读取成tf.data.Dataset:
def read_dataset(filepattern): def decode_csv(value_column): cols = tf.io.decode_csv(value_column, record_defaults=[[0.0],[0],[0.0]]) features=[cols[1],cols[2]] label = cols[0] return features, label # Create list of files that match pattern file_list = tf.io.gfile.glob(filepattern) # Create dataset from file list dataset = tf.data.TextLineDataset(file_list).map(decode_csv) return dataset training_data=read_dataset(<filepattern>)
现在的问题是,数据的第二列(也就是cols[1])是分类特征,需要做One-Hot编码。请问我该在decode_csv函数里处理,还是后续对tf.dataset做处理?具体怎么实现?
解决方案
针对这个需求,有两种实用的实现方式,你可以根据自己的代码结构偏好来选择:
方式一:在decode_csv函数内直接完成One-Hot编码
这种方式适合你已经明确知道分类特征所有可能类别数量的场景。只需要修改decode_csv里的特征处理逻辑,用tf.one_hot把整数类型的分类特征转换成One-Hot向量。
假设你的第二列分类特征总共有NUM_CLASSES个类别(比如如果是0-4的5个类别,NUM_CLASSES=5),修改后的代码如下:
def read_dataset(filepattern, num_classes=5): # 传入分类类别数 def decode_csv(value_column): cols = tf.io.decode_csv(value_column, record_defaults=[[0.0],[0],[0.0]]) # 对第二列分类特征执行One-Hot编码 one_hot_feature = tf.one_hot(cols[1], depth=num_classes, axis=-1) # 替换原特征为One-Hot向量 features = [one_hot_feature, cols[2]] label = cols[0] return features, label file_list = tf.io.gfile.glob(filepattern) dataset = tf.data.TextLineDataset(file_list).map(decode_csv) return dataset # 使用时传入实际的类别数量 training_data = read_dataset(<filepattern>, num_classes=5)
注意点:
depth参数必须准确对应分类特征的所有可能取值数量,比如如果你的分类值是0、1、2,那depth=3,否则会出现编码不全或者维度错误的问题。- 你的原始代码里第二列的
record_defaults设为[0],已经是整数类型,刚好符合tf.one_hot的输入要求(输入必须是整数张量)。
方式二:后续对Dataset单独做map处理(更灵活)
如果不想把数据读取和特征编码耦合在一起,或者后续可能需要调整编码逻辑,推荐把One-Hot编码放在Dataset读取完成后单独处理。这样代码逻辑更清晰,也方便后续修改。
实现代码如下:
# 保留原始的read_dataset函数,不做修改 def read_dataset(filepattern): def decode_csv(value_column): cols = tf.io.decode_csv(value_column, record_defaults=[[0.0],[0],[0.0]]) features=[cols[1],cols[2]] label = cols[0] return features, label file_list = tf.io.gfile.glob(filepattern) dataset = tf.data.TextLineDataset(file_list).map(decode_csv) return dataset # 定义单独的One-Hot编码函数 def one_hot_encode_feature(features, label): num_classes = 5 # 替换为实际的分类类别数 # 对第一个特征(原数据第二列)执行One-Hot编码 one_hot_feature = tf.one_hot(features[0], depth=num_classes, axis=-1) # 替换原特征,保持其他特征不变 new_features = [one_hot_feature, features[1]] return new_features, label # 读取数据集后,应用编码逻辑 training_data = read_dataset(<filepattern>) training_data = training_data.map(one_hot_encode_feature)
这种方式的优势:
- 数据读取和特征工程逻辑分离,代码更易维护。
- 如果后续需要调整分类类别数,或者更换其他编码方式(比如embedding),只需要修改
one_hot_encode_feature函数即可,不用改动数据读取的核心代码。
额外提示
如果你的分类特征类别数是动态的(比如需要从数据集中统计),可以先遍历一次数据集统计所有可能的分类值,再确定depth参数。不过在GCP AI平台的大规模数据集场景下,建议提前统计好类别数,避免额外的遍历开销。
内容的提问来源于stack exchange,提问作者Helga Holmestad
相关产品推荐
相关产品推荐

