Reduces the individual weighted loss measurements.
(weighted_losses,
reduction=ReductionV2.SUM_OVER_BATCH_SIZE)
| 56 | |
| 57 | |
| 58 | def reduce_weighted_loss(weighted_losses, |
| 59 | reduction=ReductionV2.SUM_OVER_BATCH_SIZE): |
| 60 | """Reduces the individual weighted loss measurements.""" |
| 61 | if reduction == ReductionV2.NONE: |
| 62 | loss = weighted_losses |
| 63 | else: |
| 64 | loss = math_ops.reduce_sum(weighted_losses) |
| 65 | if reduction == ReductionV2.SUM_OVER_BATCH_SIZE: |
| 66 | loss = _safe_mean(loss, _num_elements(weighted_losses)) |
| 67 | return loss |
| 68 | |
| 69 | |
| 70 | def compute_weighted_loss(losses, |
no test coverage detected