文件最后提交记录最后更新时间
1 个月前
7 个月前
2 个月前
2 个月前
2 个月前
2 个月前
7 个月前
1 个月前
README

CrossEntropyLoss

产品支持情况

产品 是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品
Atlas A2 训练系列产品/Atlas A2 推理系列产品
Atlas 200I/500 A2 推理产品 ×
Atlas 推理系列产品 ×
Atlas 训练系列产品 ×

功能说明

  • 算子功能:计算输入的交叉熵损失。

  • 计算表达式:

    reductionOptional = mean时,交叉熵损失loss的计算公式为:

    ln=−weightyn∗logexp(xn,yn)∑c=1Cexp(xn,c)∗1{yn != ignoreIndex}l_n = -weight_{y_n}*log\frac{exp(x_{n,y_n})}{\sum_{c=1}^Cexp(x_{n,c})}*1\{y_n\ !=\ ignoreIndex \}

    loss={∑n=1N1∑n=1Nweightyn∗1{yn != ignoreIndex}ln,if reductionOptional = ‘mean’∑n=1Nln,if reductionOptional = ‘sum’{l0,l1,...,ln},if reductionOptional = ‘None’loss=\begin{cases}\sum_{n=1}^N\frac{1}{\sum_{n=1}^Nweight_{y_n}*1\{y_n\ !=\ ignoreIndex \}}l_n,&\text{if reductionOptional = ‘mean’} \\\sum_{n=1}^Nl_n,&\text {if reductionOptional = ‘sum’}\\\{l_0,l_1,...,l_n\},&\text{if reductionOptional = ‘None’}\end{cases}

    log_prob计算公式为:

    lsen=log∗∑c=1Cexp(xn,c)lse_n = log*\sum_{c=1}^{C}exp(x_{n,c})

    logProbn,c=xn,c−lsenlogProb_{n,c} = x_{n,c} - lse_n

    zloss计算公式为:

    zlossn=lseSquareScaleForZloss∗(lsen)2zloss_n = lseSquareScaleForZloss *(lse_n)^2

    其中,N为batch数,C为标签数。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
x 输入
    公式中的x。
FLOAT16、BFLOAT16、FLOAT32 ND
y 输入
    表示标签,公式中的y。
INT64,INT32 ND
weight 可选输入
  • 表示为每个类别指定的缩放权重,公式中的weight。
  • 默认为全1。
FLOAT ND
reductionOptional 可选属性
  • 表示loss的归约方式。
  • 默认值为“mean”。
STRING -
ignoreIndex 可选属性
  • 指定被忽略的标签值。
  • 默认值为-100。
INT64 -
labelSmoothing 可选属性
  • 表示计算loss时的平滑量。
  • 默认值为0。
DOUBLE -
lseSquareScaleForZloss 可选属性
  • 表示zloss计算所需的scale。
  • 当前暂不支持。
DOUBLE -
returnZloss 可选属性
  • 控制是否返回zloss输出。Host侧的布尔值。需要输出zLoss时传入True,否则传入False。
  • 当前暂不支持。
BOOL -
lossOut 输出
    表示输出损失,对应公式中的loss。
FLOAT16、BFLOAT16、FLOAT32 ND
logProbOut 输出
    输出给反向计算的输出,对应公式中的logProb。
FLOAT16、BFLOAT16、FLOAT32 ND
zlossOut 输出
  • 表示辅助损失,对应公式中的zlossOut。
  • 当前暂不支持。
FLOAT16、BFLOAT16、FLOAT32 ND
lseForZlossOut 输出
  • 表示zloss场景输出给反向的Tensor,lseSquareScaleForZloss为0时输出为None,对应公式中的lse。
  • 当前暂不支持。
FLOAT16、BFLOAT16、FLOAT32 ND

约束说明

  • target仅支持类标签索引,不支持概率输入。
  • 当前暂不支持zloss相关功能。传入相关输入,即lseSquareScaleForZloss、returnZloss,不会生效。

调用说明

调用方式 样例代码 说明
aclnn接口 test_aclnn_cross_entropy_loss 通过aclnnCrossEntropyLoss接口方式调用CrossEntropyLoss算子。