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

如何在tf.data.Dataset中自动检测特征并实现特征堆叠?

自动检测特征实现TensorFlow数据堆叠

原代码中创建main_inputs时硬编码了特征名称,无法适配特征列表变化的场景,可通过以下方式实现自动检测所有输入特征并完成堆叠:

# loading csv
dataset = tf.data.experimental.make_csv_dataset(
        file_pattern=filename,
        num_parallel_reads=2,
        batch_size=128,
        num_epochs=1,
        label_name='streamflow',
        select_columns=keep_columns,
        shuffle_buffer_size=10000,
        header=True,
        field_delim=','
    )

def preprocess_fn(features, label):
    # 自动获取当前所有输入特征的名称
    feature_names = list(features.keys())
    
    # 批量归一化特征(如需不同规则,可自定义映射表)
    for name in feature_names:
        features[name] /= 100.0
    
    # 自动堆叠所有特征生成main_inputs
    feature_tensors = [features[name] for name in feature_names]
    features['main_inputs'] = tf.stack(feature_tensors, axis=-1)
   
    return {'main_inputs': features['main_inputs']}, label
    
dataset = dataset.map(preprocess_fn)

关键说明:

  • 利用list(features.keys())自动获取输入特征名,彻底摆脱硬编码,适配任意特征列表变化
  • 归一化逻辑改为批量处理,若不同特征需要不同归一化规则,可提前定义规则字典(如norm_config = {'feat1': 100, 'feat2': 200}),循环时按名称匹配规则即可
  • Python 3.7+中字典键的顺序与select_columns指定的列顺序一致,tf.stack基于自动生成的特征张量列表完成堆叠,保证顺序符合预期

内容的提问来源于stack exchange,提问作者Jonathan Roy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 00:46:13