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

如何在Keras中实现按组平均MAE的自定义损失函数并适配批量训练?

解决Keras中自定义分组平均MAE损失函数的批量处理问题

问题背景

需在Keras中实现按分组计算MAE后取组平均的自定义损失函数,现有数据集与问题如下:

  • 拆分得到(X_train, y_train), (X_val, y_val):X包含80个数值特征+1个3分类变量的独热编码(2个特征),y为2个数值型输出
  • 原损失函数依赖Pandas全量数据处理,无法适配批量训练逻辑,导致训练失效

原损失函数代码:

def avg_mae(grouping_col_train, grouping_col_val):
    def custom_loss(y_true, y_pred):
        grouping_col = grouping_col_val
        if len(y_true) == len(grouping_col_train):
            grouping_col = grouping_col_train
        df = pd.DataFrame({
            "geid": grouping_col,
            "y_true": y_true,
            "y_pred": y_pred
        })
        return np.mean(df.groupby("geid").apply(lambda x: mean_absolute_error(x["y_true"], x["y_pred"])))
    return custom_loss

训练代码:

model.fit(X_train, 
          y_train, 
          epochs=500, 
          batch_size=2**16, 
          validation_data=(X_val, y_val), 
          callbacks=[callback],
          verbose=1)

全局MAE与分组平均MAE的差异示例:

import numpy as np
import pandas as pd
from sklearn.metrics import mean_absolute_error

df = pd.DataFrame({
    "y_true": np.random.randn(20000),
    "y_pred": np.random.randn(20000),
    "grouping_col": ["a"] * 8000 + ["b"] * 1000 + ["c"] * 11000
})

overall_mae = mean_absolute_error(df["y_true"], df["y_pred"])
print("overall_mae:", overall_mae)

grouped_mae = df.groupby(["grouping_col"]).apply(lambda x: mean_absolute_error(x["y_true"], x["y_pred"]))
print("\ngrouped_mae:")
print(grouped_mae)

avg_grouped_mae = np.mean(grouped_mae)
print("\navg_grouped_mae:", avg_grouped_mae)

输出:

overall_mae: 1.1325261117842

grouped_mae:
grouping_col
a    1.141619
b    1.069323
c    1.131659
dtype: float64

avg_grouped_mae: 1.1142004357897866

解决方案

1. 调整输入结构:拆分分组列作为独立输入

将独热编码的分组特征转回原始类别索引(避免独热编码的维度缺失),并从特征矩阵中拆分出来,作为模型的第二个输入:

# 从独热编码转回类别索引(3个类别对应0/1/2)
def get_group_indices(one_hot_group):
    third_col = 1 - np.sum(one_hot_group, axis=1)
    full_one_hot = np.column_stack([one_hot_group, third_col])
    return np.argmax(full_one_hot, axis=1)

# 提取分组索引并拆分特征矩阵
group_train = get_group_indices(X_train[:, -2:])
group_val = get_group_indices(X_val[:, -2:])
X_train_features = X_train[:, :-2]
X_val_features = X_val[:, :-2]

2. 实现TensorFlow原生的分组平均MAE损失函数

批量训练时数据为Tensor,必须用TensorFlow向量化操作实现分组逻辑,避免破坏计算图:

import tensorflow as tf

def grouped_avg_mae(y_true, y_pred, group_indices):
    # 计算每个样本的平均MAE(适配多输出)
    mae_per_sample = tf.reduce_mean(tf.abs(y_true - y_pred), axis=1)
    # 获取唯一分组索引
    unique_groups = tf.unique(group_indices)[0]
    
    # 定义单组MAE计算逻辑
    def compute_single_group_mae(group):
        group_mask = tf.equal(group_indices, group)
        return tf.reduce_mean(tf.boolean_mask(mae_per_sample, group_mask))
    
    # 计算所有组的MAE并取平均
    group_maes = tf.map_fn(compute_single_group_mae, unique_groups, dtype=tf.float32)
    return tf.reduce_mean(group_maes)

# 封装为Keras兼容的损失函数
def custom_loss_wrapper(group_indices):
    def loss(y_true, y_pred):
        return grouped_avg_mae(y_true, y_pred, group_indices)
    return loss

3. 修改模型结构适配多输入

模型需要同时接收特征矩阵和分组索引两个输入:

from tensorflow.keras.models import Model
from tensorflow.keras.layers import Input, Dense

# 定义输入层
input_features = Input(shape=(X_train_features.shape[1],))
input_groups = Input(shape=(1,), dtype=tf.int32)

# 搭建网络结构示例
x = Dense(64, activation='relu')(input_features)
x = Dense(32, activation='relu')(x)
output = Dense(2)(x)  # 对应2个数值输出

# 构建模型
model = Model(inputs=[input_features, input_groups], outputs=output)
# 编译时传入分组索引输入
model.compile(optimizer='adam', loss=custom_loss_wrapper(input_groups))

4. 调整训练与验证流程

训练时传入双输入,验证阶段单独计算全局分组平均MAE(避免批量计算的误差):

# 训练模型
model.fit(
    [X_train_features, group_train],
    y_train,
    epochs=500,
    batch_size=2**16,
    validation_data=([X_val_features, group_val], y_val),
    callbacks=[callback],
    verbose=1
)

# 验证集上计算全局分组平均MAE
y_val_pred = model.predict([X_val_features, group_val])
val_mae_per_sample = np.mean(np.abs(y_val - y_val_pred), axis=1)
val_grouped_maes = []
for g in np.unique(group_val):
    val_grouped_maes.append(np.mean(val_mae_per_sample[group_val == g]))
val_avg_grouped_mae = np.mean(val_grouped_maes)
print("验证集分组平均MAE:", val_avg_grouped_mae)

关键注意事项

  • 损失函数必须用TensorFlow原生操作:批量数据为Tensor,使用Pandas/Numpy会断裂计算图,无法自动求导
  • 验证阶段需单独计算全局分组MAE:Keras默认验证损失是批量分组平均的总平均,与全量数据的分组平均结果可能不一致
  • 分组索引需为整数类型:确保TensorFlow的tf.equal匹配稳定

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 07:21:33