| 46 | return self._K_P * error + self._K_I * integral + self._K_D * derivative |
| 47 | |
| 48 | class FocalLoss(nn.Module): |
| 49 | def __init__(self, alpha=0.5, gamma=2, weight=None, ignore_index=255): |
| 50 | super().__init__() |
| 51 | self.alpha = alpha |
| 52 | self.gamma = gamma |
| 53 | self.weight = weight |
| 54 | self.ignore_index = ignore_index |
| 55 | self.ce_fn = nn.CrossEntropyLoss(weight=self.weight, ignore_index=self.ignore_index) |
| 56 | self.fp16_enabled = False |
| 57 | |
| 58 | @force_fp32() |
| 59 | def forward(self, preds, labels): |
| 60 | logpt = -self.ce_fn(preds.float(), labels) |
| 61 | pt = torch.exp(logpt) |
| 62 | loss = -((1 - pt) ** self.gamma) * self.alpha * logpt |
| 63 | return loss |
| 64 | |
| 65 | |
| 66 | def init_weights(m): |