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