Reduce loss as specified. Args: loss (Tensor): Elementwise loss tensor. reduction (str): Options are "none", "mean" and "sum". Return: Tensor: Reduced loss tensor.
(loss: Tensor, reduction: str)
| 8 | |
| 9 | |
| 10 | def reduce_loss(loss: Tensor, reduction: str) -> Tensor: |
| 11 | """Reduce loss as specified. |
| 12 | |
| 13 | Args: |
| 14 | loss (Tensor): Elementwise loss tensor. |
| 15 | reduction (str): Options are "none", "mean" and "sum". |
| 16 | |
| 17 | Return: |
| 18 | Tensor: Reduced loss tensor. |
| 19 | """ |
| 20 | reduction_enum = F._Reduction.get_enum(reduction) |
| 21 | # none: 0, elementwise_mean:1, sum: 2 |
| 22 | if reduction_enum == 0: |
| 23 | return loss |
| 24 | elif reduction_enum == 1: |
| 25 | return loss.mean() |
| 26 | elif reduction_enum == 2: |
| 27 | return loss.sum() |
| 28 | |
| 29 | |
| 30 | def weight_reduce_loss(loss: Tensor, |
no outgoing calls
no test coverage detected