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

TensorFlow是否支持非二进制/字符串类型的多浮点标签模型?

解决TensorFlow多浮点标签的多输出回归问题

Hey there! I see you're working on a multi-output regression task with 6 float64 labels, and running into issues with the default DNNLinearCombinedRegressor/DNNLinearCombinedClassifier—let's fix this together.

问题根源

The default setup for DNNLinearCombinedRegressor is designed for single-output regression, which is why you're getting errors when passing multiple float labels. For multi-output tasks (like your 6 continuous labels), we need to use tf.estimator.MultiHead to wrap individual regression heads for each target.

具体解决步骤

1. 调整输入函数的标签格式

First, modify your input_fn to return labels as a tuple (or list) instead of a dictionary. This ensures each label maps correctly to its corresponding regression head later:

def input_fn(data_file, num_epochs, shuffle, batch_size):
    """Generate an input function for the Estimator."""
    assert tf.gfile.Exists(data_file), (
        '%s not found. Please make sure you have run data_download.py and '
        'set the --data_dir argument to the correct path.' % data_file)

    def parse_csv(value):
        columns = tf.decode_csv(value, record_defaults=_CSV_COLUMN_DEFAULTS)
        feature_columns = columns[6:10]
        features = dict(zip(_CSV_FEATURE_COLUMNS, feature_columns))
        label_columns = columns[0:6]
        # 把标签从字典改为元组,顺序要和后续创建的head一一对应
        labels = tuple(label_columns)
        return features, labels

    dataset = tf.data.TextLineDataset(data_file)
    dataset = dataset.map(parse_csv, num_parallel_calls=5)
    dataset = dataset.repeat(num_epochs)
    dataset = dataset.batch(batch_size)
    return dataset

2. 修正默认值的数据类型(匹配float64)

Your current _CSV_COLUMN_DEFAULTS uses Python floats (which map to tf.float32). To ensure consistency with your float64 requirement, update it to use explicit tf.float64 defaults:

_CSV_FEATURE_COLUMNS = ['vgs', 'vbs', 'vds', 'current']
_CSV_LABEL_COLUMNS = ['plo_tox', 'plo_dxl', 'plo_dxw', 'parl1', 'parl2', 'random_fn']
# 用tf.float64作为默认值,确保输入输出类型一致
_CSV_COLUMN_DEFAULTS = [
    [tf.constant(0.0, dtype=tf.float64)] for _ in range(10)
]

3. 构建带MultiHead的DNNLinearCombinedRegressor

Create a MultiHead that wraps a regression head for each of your 6 labels, then pass this to the estimator. Here's how your build_estimator function should look:

def build_estimator(model_dir, model_type):
    # 定义4个float64输入特征列
    vgs = tf.feature_column.numeric_column('vgs', dtype=tf.float64)
    vbs = tf.feature_column.numeric_column('vbs', dtype=tf.float64)
    vds = tf.feature_column.numeric_column('vds', dtype=tf.float64)
    current = tf.feature_column.numeric_column('current', dtype=tf.float64)
    
    wide_columns = [vgs, vbs, vds, current]
    deep_columns = [vgs, vbs, vds, current]

    # 为每个标签创建独立的回归head
    heads = []
    for label_name in _CSV_LABEL_COLUMNS:
        head = tf.estimator.RegressionHead(
            label_dimension=1,
            name=label_name,
            dtype=tf.float64
        )
        heads.append(head)
    
    # 组合成MultiHead,支持多输出回归
    multi_head = tf.estimator.MultiHead(heads)

    # 根据model_type创建对应的宽深模型
    if model_type == 'wide':
        return tf.estimator.DNNLinearCombinedRegressor(
            model_dir=model_dir,
            linear_feature_columns=wide_columns,
            head=multi_head
        )
    elif model_type == 'deep':
        return tf.estimator.DNNLinearCombinedRegressor(
            model_dir=model_dir,
            dnn_feature_columns=deep_columns,
            dnn_hidden_units=[128, 64],  # 可根据任务需求调整隐藏层大小
            head=multi_head
        )
    else:
        return tf.estimator.DNNLinearCombinedRegressor(
            model_dir=model_dir,
            linear_feature_columns=wide_columns,
            dnn_feature_columns=deep_columns,
            dnn_hidden_units=[128, 64],
            head=multi_head
        )

为什么这样有效?

  • MultiHead lets you train multiple regression tasks simultaneously, with each head handling one of your 6 float labels.
  • By matching the dtype across features, labels, and heads, you avoid type mismatch errors.
  • The tuple format for labels ensures each target is correctly associated with its head.

评估时的注意事项

When you run model.evaluate(), the results will include metrics for each individual label (e.g., plo_tox_loss, plo_dxl_root_mean_squared_error), so you can monitor the performance of each output separately.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 07:32:09