Premium problem92. Class-Weighted Cross-Entropy

Medium Locked

Compute multi-class cross-entropy where each class carries a weight from class_weights, the usual remedy for an imbalanced training set.

logits is (B, C), targets is (B,). Weight each sample's loss by its target class's weight, then divide by the sum of the weights used, not by the batch size -- that is what PyTorch's own weighted reduction does, and it keeps the result a weighted mean rather than a weighted sum.

Input

logits =
tensor([[2.0000, 1.0000],
        [0.5000, 2.5000],
        [1.0000, 1.0000]])
targets = tensor([0, 1, 0])
class_weights = tensor([3., 1.])

Output

tensor(0.4495)

Premium problem

This one's part of Premium. Unlock the full PyTorch track plus every other premium problem on the site.

Implement solve(...)