| 文件 | 最后提交记录 | 最后更新时间 |
|---|---|---|
| 1 个月前 | ||
| 7 个月前 | ||
| 2 个月前 | ||
| 2 个月前 | ||
| 2 个月前 | ||
| 2 个月前 | ||
| 7 个月前 | ||
| 1 个月前 |
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 | 输入 |
|
FLOAT16、BFLOAT16、FLOAT32 | ND |
| y | 输入 |
|
INT64,INT32 | ND |
| weight | 可选输入 |
|
FLOAT | ND |
| reductionOptional | 可选属性 |
|
STRING | - |
| ignoreIndex | 可选属性 |
|
INT64 | - |
| labelSmoothing | 可选属性 |
|
DOUBLE | - |
| lseSquareScaleForZloss | 可选属性 |
|
DOUBLE | - |
| returnZloss | 可选属性 |
|
BOOL | - |
| lossOut | 输出 |
|
FLOAT16、BFLOAT16、FLOAT32 | ND |
| logProbOut | 输出 |
|
FLOAT16、BFLOAT16、FLOAT32 | ND |
| zlossOut | 输出 |
|
FLOAT16、BFLOAT16、FLOAT32 | ND |
| lseForZlossOut | 输出 |
|
FLOAT16、BFLOAT16、FLOAT32 | ND |
约束说明
- target仅支持类标签索引,不支持概率输入。
- 当前暂不支持zloss相关功能。传入相关输入,即lseSquareScaleForZloss、returnZloss,不会生效。
调用说明
| 调用方式 | 样例代码 | 说明 |
|---|---|---|
| aclnn接口 | test_aclnn_cross_entropy_loss | 通过aclnnCrossEntropyLoss接口方式调用CrossEntropyLoss算子。 |