MCPcopy Create free account
hub / github.com/OpenDriveLab/DriveAdapter / FocalLoss

Class FocalLoss

open_loop_training/code/encoder_decoder_framework.py:48–63  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

46 return self._K_P * error + self._K_I * integral + self._K_D * derivative
47
48class 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
66def init_weights(m):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected