Args: logits (N, C): tensor of logits. targets (N, ): :math:`targets_i \in {0,1}` Returns: (0,)
(self, logits: torch.Tensor, targets: torch.Tensor)
| 41 | raise ValueError(f"wrong reduction type {reduction}") |
| 42 | |
| 43 | def forward(self, logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor: |
| 44 | """ |
| 45 | Args: |
| 46 | logits (N, C): |
| 47 | tensor of logits. |
| 48 | targets (N, ): |
| 49 | :math:`targets_i \in {0,1}` |
| 50 | Returns: |
| 51 | (0,) |
| 52 | """ |
| 53 | BCE_loss = F.binary_cross_entropy_with_logits(logits, targets, reduction="none") |
| 54 | pt = torch.exp(-BCE_loss) # prevents nans when probability 0 |
| 55 | if self.dynamic_balance: |
| 56 | pos = torch.tensor(targets, dtype=torch.float, device=logits.device) |
| 57 | prob_pos = torch.mean(pos) |
| 58 | at = targets * prob_pos + (1 - targets) * (1.0 - prob_pos) |
| 59 | else: |
| 60 | at = targets * self.alpha + (1 - targets) * (1.0 - self.alpha) |
| 61 | loss = 2 * at * (1 - pt).pow(self.gamma) * BCE_loss |
| 62 | if self.reduction == "mean": |
| 63 | return loss.mean() |
| 64 | elif self.reduction == "none": |
| 65 | return loss |
| 66 | |
| 67 | |
| 68 | if __name__ == "__main__": |
nothing calls this directly
no outgoing calls
no test coverage detected