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

如何在PyTorch中实现支持反向传播的CCC自定义损失函数

PyTorch自定义CCC(一致性相关系数)损失实现

PyTorch的自动微分引擎会自动追踪所有原生张量操作构建计算图,只要自定义损失的计算全程使用PyTorch原生API,不需要手动编写反向传播逻辑,就可以和内置损失函数一样直接调用backward()完成梯度计算。

CCC的数学定义如下:
维基百科CCC定义公式

公式各参数含义:

  • $\rho_c$ 即CCC值,取值范围[-1, 1],越接近1代表预测值与真实值一致性越好
  • $\mu_x$、$\mu_y$ 分别为真实值、预测值的均值
  • $\sigma_x2$、$\sigma_y2$ 分别为真实值、预测值的方差
  • $\sigma_{xy}$ 为真实值与预测值的协方差

注意:CCC本身是越接近1效果越好,作为损失函数使用时需要返回1 - CCC,把优化目标转换为最小化损失,和PyTorch内置损失的优化逻辑对齐。

具体实现

推荐继承nn.Module实现损失类,和PyTorch内置损失(如nn.MSELoss)的调用方式完全一致,实现时全程使用torch原生操作即可自动支持反向传播,同时添加极小值到分母避免除零的数值问题。

import torch
import torch.nn as nn

class CCCLoss(nn.Module):
    def __init__(self, eps=1e-8):
        super().__init__()
        self.eps = eps  # 数值稳定项,避免分母为0报错
    
    def forward(self, y_pred: torch.Tensor, y_true: torch.Tensor) -> torch.Tensor:
        # 拉平输入张量,兼容任意形状的预测/真实值输入
        y_pred = y_pred.flatten()
        y_true = y_true.flatten()

        # 计算均值
        mean_pred = torch.mean(y_pred)
        mean_true = torch.mean(y_true)

        # 计算有偏方差,和原CCC公式定义对齐
        var_pred = torch.var(y_pred, unbiased=False)
        var_true = torch.var(y_true, unbiased=False)

        # 计算有偏协方差
        cov = torch.mean((y_pred - mean_pred) * (y_true - mean_true))

        # 代入公式计算CCC
        ccc = 2 * cov / (var_pred + var_true + (mean_pred - mean_true) ** 2 + self.eps)

        # 返回损失值(最小化1-CCC等价于最大化CCC)
        return 1 - ccc

使用方法

和PyTorch内置损失的调用逻辑完全一致,不需要额外处理反向传播:

# 初始化损失实例
criterion = CCCLoss()

# 构造模拟数据,预测值需要开启梯度记录
y_pred = torch.randn(32, 1, requires_grad=True)  # 批量大小32的预测值
y_true = torch.randn(32, 1)                      # 对应的真实值

# 前向计算损失
loss = criterion(y_pred, y_true)

# 反向传播计算梯度
loss.backward()

# 可以正常获取预测值对应的梯度,接入常规训练循环即可
print(y_pred.grad)

常见踩坑说明

  • 计算过程中不要调用.item()、.numpy()等方法将张量转换为Python原生数值或Numpy数组,这类操作会切断张量和计算图的关联,导致反向传播失效
  • 计算方差、协方差时要设置unbiased=False使用有偏估计,否则计算结果和CCC原公式定义存在偏差
  • 批量训练时不要逐样本计算CCC后取平均,正确做法是将整个batch的样本拉平后整体计算CCC,避免结果有偏
  • 确保预测值和真实值在同一设备(CPU/GPU)上、数据类型一致,避免运行报错

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 05:33:17