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

如何在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:

  1. 继承TFX的BaseExampleGen类。
  2. 在数据读取阶段按类别分组,对每组执行分层K折拆分,分配对应折标签。
  3. 将不同折的样本写入对应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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 20:01:14