Skip to content

[torchlib] Support class probability targets in aten::cross_entropy_loss - #3071

Open
Raashish Aggarwal (raashish1601) wants to merge 1 commit into
microsoft:mainfrom
raashish1601:feature/cross-entropy-prob-target
Open

Raashish Aggarwal (raashish1601) wants to merge 1 commit into
microsoft:mainfrom
raashish1601:feature/cross-entropy-prob-target

Conversation

@raashish1601

Copy link
Copy Markdown

F.cross_entropy also accepts class probabilities as the target (same shape as the input, float dtype), which is common for soft labels, mixup and distillation. aten_cross_entropy_loss always lowered to SoftmaxCrossEntropyLoss, which only accepts integer class indices, so the export produced an invalid model:

x = torch.tensor([[2.0, 0.0, 0.0]]); t = torch.tensor([[0.7, 0.2, 0.1]])
F.cross_entropy(x, t)  # 0.8395
# onnxruntime: INVALID_GRAPH ... Type 'tensor(float)' of input parameter ... of operator (SoftmaxCrossEntropyLoss)

For a floating-point target this now follows PyTorch's cross_entropy_loss_prob_target: -sum(log_softmax(x) * target * weight) over the class dim, with label_smoothing applied to the target (target * (1 - eps) + eps / C). mean divides by the number of elements excluding the class dim, as PyTorch does for probability targets. Unbatched (C,) input is handled too. Integer targets keep the existing SoftmaxCrossEntropyLoss path, and the target annotation changes from IntType to TensorType.

Testing:

  • Removed the xfail for non-int64 targets on nn.functional.cross_entropy. Those 20 samples fail on main and pass with the change. pytest tests/function_libs/torch_lib/ops_test.py -k cross_entropy passes locally (onnxruntime 1.23.0).
  • I also compared torch.onnx.export(..., dynamo=True) + onnxruntime against eager PyTorch for probability targets on (C,), (N, C), (N, C, d1) and (N, C, d1, d2) inputs with every reduction, with and without weight, and with label_smoothing 0 and 0.2 (48 cases). All match within 1e-5.
  • ruff check and format are clean on the changed files.

This touches the same function as #3069 (label_smoothing for index targets), but the two branches merge without conflicts and the combined tests pass.

SoftmaxCrossEntropyLoss only takes class indices, so exporting cross_entropy with float (probability) targets produced an invalid model. Compute it with LogSoftmax like PyTorch's cross_entropy_loss_prob_target, including weight, label_smoothing and all reductions.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

Development

Successfully merging this pull request may close these issues.

1 participant