| 1721 | } |
| 1722 | |
| 1723 | void ComputeCrossEntropy(const Tensor &p, const Tensor &t, Tensor *loss) { |
| 1724 | CHECK_LE(p.nDim(), 2u); |
| 1725 | CHECK_LE(t.nDim(), 2u); |
| 1726 | size_t batchsize = 1; |
| 1727 | if (p.nDim() == 2u) batchsize = p.shape(0); |
| 1728 | size_t dim = p.Size() / batchsize; |
| 1729 | TYPE_LANG_SWITCH(p.data_type(), DType, p.device()->lang(), Lang, { |
| 1730 | Tensor &lossRef = *loss; |
| 1731 | p.device()->Exec( |
| 1732 | [batchsize, dim, t, p, lossRef](Context *ctx) mutable { |
| 1733 | bool int_target = t.Size() == batchsize; |
| 1734 | ComputeCrossEntropy<DType, Lang>(int_target, batchsize, dim, p, t, |
| 1735 | &lossRef, ctx); |
| 1736 | }, |
| 1737 | {p.block(), t.block()}, {loss->block()}, "ComputeCrossEntropy"); |
| 1738 | }); |
| 1739 | } |
| 1740 | |
| 1741 | void SoftmaxCrossEntropyBwd(const Tensor &t, Tensor *p) { |
| 1742 | CHECK_LE(p->nDim(), 2u); |