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

基于TensorFlow实现带分组MSE的回归任务自定义损失函数

问题分析与解决方案

核心问题1:TensorFlow中分组依赖损失的构建方式

你的损失函数是全局分组统计量的差值,并非逐样本损失的累加,因此不能用常规逐样本损失的写法。核心是要在计算时保留分组信息,同时保证TensorFlow自动微分能正常追踪梯度。

优化后的损失函数实现

用tf.where替代boolean_mask,避免批量维度不一致问题(比如某批中某组无样本),同时加入除以0保护,提升稳定性:

def custom_loss(y_true, y_pred, group):
    # 统一张量形状
    y_true = tf.reshape(y_true, [-1])
    y_pred = tf.reshape(y_pred, [-1])
    group = tf.reshape(group, [-1])
    
    # 计算每个样本的平方误差
    se = tf.square(y_true - y_pred)
    
    # 生成分组掩码并计算每组样本数
    group_0_mask = tf.equal(group, 0)
    group_1_mask = tf.equal(group, 1)
    
    count_0 = tf.reduce_sum(tf.cast(group_0_mask, tf.float32))
    count_1 = tf.reduce_sum(tf.cast(group_1_mask, tf.float32))
    
    # 避免除以0(某批无对应分组样本时用极小值替代)
    count_0 = tf.maximum(count_0, 1e-6)
    count_1 = tf.maximum(count_1, 1e-6)
    
    # 计算分组MSE
    mse_0 = tf.reduce_sum(tf.where(group_0_mask, se, 0.0)) / count_0
    mse_1 = tf.reduce_sum(tf.where(group_1_mask, se, 0.0)) / count_1
    
    # 用平方差值替代绝对差值,优化梯度稳定性
    return tf.square(mse_0 - mse_1)

核心问题2:简化分组损失的TensorFlow技术

  • tf.math.segment_*系列函数:如果分组是连续整数标识,可直接用tf.math.segment_mean计算分组统计量,示例如下:
    def segment_based_loss(y_true, y_pred, group):
        se = tf.square(y_true - y_pred)
        # 按分组排序,适配segment操作要求
        sorted_indices = tf.argsort(group)
        sorted_group = tf.gather(group, sorted_indices)
        sorted_se = tf.gather(se, sorted_indices)
        # 计算分组MSE
        mse_per_group = tf.math.segment_mean(sorted_se, sorted_group)
        # 兼容分组缺失情况
        mse_0 = mse_per_group[0] if tf.shape(mse_per_group)[0] >=1 else 0.0
        mse_1 = mse_per_group[1] if tf.shape(mse_per_group)[0] >=2 else 0.0
        return tf.square(mse_0 - mse_1)
    
  • 显式传递分组信息:不要将分组作为损失函数的闭包参数,训练时直接传入当前批次的分组,符合TensorFlow数据流范式,也便于结合tf.data管道。

你的训练代码问题及修复

问题1:闭包传递分组的风险

原代码将group作为闭包传入损失函数,若批次分组分布与全局差异大,会导致损失信号不稳定,是泛化差的主要原因之一,需改为显式传入分组。

问题2:验证集损失计算错误

原代码用model.predict(X_val)得到numpy数组,会断开TensorFlow梯度追踪(虽验证集无需梯度,但计算逻辑应与训练一致),需改用model(X_val, training=False)获取张量形式预测值。

问题3:批量处理逻辑不严谨

原代码丢弃最后一批不足batch_size的样本,改用tf.data.Dataset可自动处理批量与剩余样本,提升数据利用率。

修复后的完整训练代码

import numpy as np
import tensorflow as tf
from sklearn.datasets import make_regression
from sklearn.model_selection import train_test_split

# 优化后的损失函数
def custom_loss(y_true, y_pred, group):
    y_true = tf.reshape(y_true, [-1])
    y_pred = tf.reshape(y_pred, [-1])
    group = tf.reshape(group, [-1])
    
    se = tf.square(y_true - y_pred)
    
    group_0_mask = tf.equal(group, 0)
    group_1_mask = tf.equal(group, 1)
    
    count_0 = tf.reduce_sum(tf.cast(group_0_mask, tf.float32))
    count_1 = tf.reduce_sum(tf.cast(group_1_mask, tf.float32))
    
    count_0 = tf.maximum(count_0, 1e-6)
    count_1 = tf.maximum(count_1, 1e-6)
    
    mse_0 = tf.reduce_sum(tf.where(group_0_mask, se, 0.0)) / count_0
    mse_1 = tf.reduce_sum(tf.where(group_1_mask, se, 0.0)) / count_1
    
    return tf.square(mse_0 - mse_1)

# 用tf.data构建高效数据管道
def create_dataset(X, y, group, batch_size, shuffle=True):
    dataset = tf.data.Dataset.from_tensor_slices((X, y, group))
    if shuffle:
        dataset = dataset.shuffle(buffer_size=len(X))
    dataset = dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)
    return dataset

def train_early_stopping(model, train_dataset, val_dataset,
                         n_epoch=500, patience=10):
    best_val_loss = float('inf')
    wait = 0
    best_epoch = 0
    
    for epoch in range(n_epoch):
        train_losses = []
        # 训练循环
        for X_batch, y_batch, g_batch in train_dataset:
            with tf.GradientTape() as tape:
                y_pred = model(X_batch, training=True)
                loss_value = custom_loss(y_batch, y_pred, g_batch)
            grads = tape.gradient(loss_value, model.trainable_variables)
            model.optimizer.apply_gradients(zip(grads, model.trainable_variables))
            train_losses.append(loss_value.numpy())
        
        # 验证循环
        val_losses = []
        for X_val, y_val, g_val in val_dataset:
            y_pred_val = model(X_val, training=False)
            val_loss = custom_loss(y_val, y_pred_val, g_val)
            val_losses.append(val_loss.numpy())
        avg_val_loss = np.mean(val_losses)
        
        print(f"Epoch {epoch+1}: Train Loss: {np.mean(train_losses):.4f}, Validation Loss: {avg_val_loss:.4f}")
        
        # 早停逻辑
        if avg_val_loss < best_val_loss:
            best_val_loss = avg_val_loss
            best_weights = model.get_weights()
            wait = 0
            best_epoch = epoch
        else:
            wait += 1
            if wait >= patience:
                print(f"Early Stopping triggered at epoch {best_epoch + 1}, Validation Loss: {best_val_loss:.4f}")
                model.set_weights(best_weights)
                break
    if epoch == n_epoch-1:
        print('Not converged.')

# 生成合成数据
X, y = make_regression(n_samples=20000, n_features=10, noise=0.2, random_state=42)
group = np.random.choice([0, 1], size=y.shape)
X_train_full, X_test, y_train_full, y_test, g_train_full, g_test = train_test_split(X, y, group, test_size=0.5, random_state=42)
X_train, X_val, y_train, y_val, g_train, g_val = train_test_split(X_train_full, y_train_full, g_train_full, test_size=0.2, random_state=42)

# 构建模型
num_unit = 64
model_fair = tf.keras.Sequential([
    tf.keras.layers.Dense(num_unit, activation='relu', input_shape=(X.shape[1],)),
    tf.keras.layers.Dense(num_unit, activation='relu'),
    tf.keras.layers.Dense(1)
])
model_fair.compile(optimizer='adam')

# 创建数据集
batch_size = X_train.shape[0]//5
train_dataset = create_dataset(X_train, y_train, g_train, batch_size)
val_dataset = create_dataset(X_val, y_val, g_val, batch_size, shuffle=False)

# 启动训练
train_early_stopping(model_fair, train_dataset, val_dataset, patience=10)

额外建议

  • 批量大小设置:由于损失依赖分组统计,批量不能过小,建议保证每组至少有50个样本,避免批次MSE与全局MSE偏差过大。
  • 损失函数选择:优先用平方差值替代绝对差值,后者在零点处不可导,会导致训练不稳定。
  • 数据分布对齐:确保训练集、验证集的分组分布一致,避免分布偏移导致泛化差。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 00:17:07