@@ -225,12 +225,8 @@ def __init__(
225225 self .gamma = gamma
226226 self .delta = delta
227227 self .weight : float = weight
228- self .asy_focal_loss = AsymmetricFocalLoss (
229- to_onehot_y = False , gamma = self .gamma , delta = self .delta , ignore_index = ignore_index
230- )
231- self .asy_focal_tversky_loss = AsymmetricFocalTverskyLoss (
232- to_onehot_y = False , gamma = self .gamma , delta = self .delta , ignore_index = ignore_index
233- )
228+ self .asy_focal_loss = AsymmetricFocalLoss (to_onehot_y = False , gamma = self .gamma , delta = self .delta )
229+ self .asy_focal_tversky_loss = AsymmetricFocalTverskyLoss (to_onehot_y = False , gamma = self .gamma , delta = self .delta )
234230 self .ignore_index = ignore_index
235231
236232 def forward (self , y_pred : torch .Tensor , y_true : torch .Tensor ) -> torch .Tensor :
@@ -283,14 +279,11 @@ def forward(self, y_pred: torch.Tensor, y_true: torch.Tensor) -> torch.Tensor:
283279
284280 mask = create_ignore_mask (original_y_true , self .ignore_index )
285281
286- use_mask = self .ignore_index is not None and (self .ignore_index < 0 or self .ignore_index >= self .num_classes )
287- if use_mask :
288- y_pred_masked , y_true_masked = mask_loss_inputs (y_pred , y_true , self .ignore_index , mask = mask )
289- else :
290- y_pred_masked , y_true_masked = y_pred , y_true
282+ if self .ignore_index is not None :
283+ y_pred , y_true = mask_loss_inputs (y_pred , y_true , self .ignore_index , mask = mask )
291284
292- asy_focal_loss = self .asy_focal_loss (y_pred_masked , y_true_masked )
293- asy_focal_tversky_loss = self .asy_focal_tversky_loss (y_pred_masked , y_true_masked )
285+ asy_focal_loss = self .asy_focal_loss (y_pred , y_true )
286+ asy_focal_tversky_loss = self .asy_focal_tversky_loss (y_pred , y_true )
294287
295288 loss : torch .Tensor = self .weight * asy_focal_loss + (1 - self .weight ) * asy_focal_tversky_loss
296289
0 commit comments