如何在TFX中按固定属性比例拆分数据集为多折用于K折验证
在TFX中实现分层K折交叉验证(保持类别比例)
问题场景
我有一个含多输入特征、单二分类输出的不平衡数据集(90%样本为0,10%为1),需要拆分为K份用于交叉验证,要求每个折严格保持9:1的类别比例。已知Pandas的实现方式,但想了解TFX中的解决方案。
之前尝试过用SplitConfig直接拆分:
output = tfx.proto.Output( split_config=tfx.proto.SplitConfig(splits=[ tfx.proto.SplitConfig.Split(name='fold_1', hash_buckets=1), tfx.proto.SplitConfig.Split(name='fold_2', hash_buckets=1), tfx.proto.SplitConfig.Split(name='fold_3', hash_buckets=1), tfx.proto.SplitConfig.Split(name='fold_4', hash_buckets=1), tfx.proto.SplitConfig.Split(name='fold_5', hash_buckets=1) ])) example_gen = CsvExampleGen(input_base=input_dir, output_config=output)
但这种随机拆分无法保证各折的类别比例;尝试过partition_feature_name参数,但需要给每个样本手动添加折ID特征,实验中调整折数时太麻烦,不实用。
可行解决方案
方法1:自动生成分层折ID + 利用partition_feature_name
无需手动维护折ID,通过预处理脚本自动生成:
- 用Pandas读取原始CSV,对每个类别单独执行分层K折拆分,给每个样本分配
fold_id字段(取值1-K)。 - 将带
fold_id的数据集保存为新CSV。 - 在TFX中指定按
fold_id拆分:
output = tfx.proto.Output( split_config=tfx.proto.SplitConfig( splits=[ tfx.proto.SplitConfig.Split(name='fold_1', hash_buckets=1), tfx.proto.SplitConfig.Split(name='fold_2', hash_buckets=1), tfx.proto.SplitConfig.Split(name='fold_3', hash_buckets=1), tfx.proto.SplitConfig.Split(name='fold_4', hash_buckets=1), tfx.proto.SplitConfig.Split(name='fold_5', hash_buckets=1) ], partition_feature_name='fold_id' ) ) example_gen = CsvExampleGen(input_base=preprocessed_input_dir, output_config=output)
调整折数时只需修改预处理脚本的K值,重新生成数据集即可,操作成本低。
方法2:自定义ExampleGen组件
如果希望完全在TFX pipeline内完成分层拆分,可以自定义ExampleGen:
- 继承TFX的
BaseExampleGen类。 - 在数据读取阶段按类别分组,对每组执行分层K折拆分,分配对应折标签。
- 将不同折的样本写入对应Split。
核心逻辑示例:
import pandas as pd from sklearn.model_selection import StratifiedKFold from tfx.components.example_gen.base_example_gen import BaseExampleGen from tfx.utils import io_utils import tensorflow as tf class StratifiedKFoldExampleGen(BaseExampleGen): def _generate_examples(self, input_dict): # 读取输入目录下的CSV文件 input_base = input_dict['input_base'] csv_files = io_utils.get_only_files(input_base, '.csv') df = pd.concat([pd.read_csv(os.path.join(input_base, f)) for f in csv_files]) # 初始化分层K折 skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=42) label_col = 'your_label_column_name' # 替换为你的输出标签列名 # 遍历每个折,生成对应样本 for fold_idx, (_, val_idx) in enumerate(skf.split(df, df[label_col])): fold_samples = df.iloc[val_idx] split_name = f'fold_{fold_idx + 1}' for idx, row in fold_samples.iterrows(): # 将行数据转换为tf.train.Example feature = {} for col in df.columns: val = row[col] if pd.api.types.is_integer_dtype(df[col]): feature[col] = tf.train.Feature(int64_list=tf.train.Int64List(value=[val])) elif pd.api.types.is_float_dtype(df[col]): feature[col] = tf.train.Feature(float_list=tf.train.FloatList(value=[val])) else: feature[col] = tf.train.Feature(bytes_list=tf.train.BytesList(value=[str(val).encode('utf-8')])) example = tf.train.Example(features=tf.train.Features(feature=feature)) # 输出到对应Split yield f'{split_name}_{idx}', example
在TFX pipeline中直接使用这个自定义组件即可。
总结
- 优先选择方法1,实现简单、维护成本低,适配大多数场景。
- 若需要完全嵌入TFX pipeline流程,可采用自定义ExampleGen的方式。
内容的提问来源于stack exchange,提问作者Mehran
相关产品推荐
相关产品推荐

