MCPcopy Create free account
hub / github.com/apple/ml-pointersect / forward

Method forward

cdslib/core/nn/modules/focal_loss.py:43–65  ·  view source on GitHub ↗

Args: logits (N, C): tensor of logits. targets (N, ): :math:`targets_i \in {0,1}` Returns: (0,)

(self, logits: torch.Tensor, targets: torch.Tensor)

Source from the content-addressed store, hash-verified

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
68if __name__ == "__main__":

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected