Thank you for the awesome repository!
I've noticed torch.nn.CrossEntropyLoss is used for the cross entropy loss and a custom loss from utils.losses is used for the Dice loss, as used as follows:
| ce_loss=CrossEntropyLoss() |
| dice_loss=losses.DiceLoss(2) |
The Dice loss seems to use a 'sum' reduction as follows:
| def_dice_loss(self, score, target): |
| target=target.float() |
| smooth=1e-5 |
| intersect=torch.sum(score*target) |
| y_sum=torch.sum(target*target) |
| z_sum=torch.sum(score*score) |
| loss= (2*intersect+smooth) / (z_sum+y_sum+smooth) |
| loss=1-loss |
| returnloss |
However, the default reduction method for
torch.nn.CrossEntropyLoss is 'mean', so the Dice loss is always roughly about
H*W(*D) times bigger than the CE loss.
So, a direct mean of two losses as used in the following code would not be actually the intended average.
| supervised_loss=0.5* (loss_dice+loss_ce) |
Although I am sure this has minimal effects on most of your SSL methods because it is simply using Dice instead of Dice + CE for the supervisised loss, but still I think it should be checked.
Thank you for the awesome repository!
I've noticed
torch.nn.CrossEntropyLossis used for the cross entropy loss and a custom loss fromutils.lossesis used for the Dice loss, as used as follows:SSL4MIS/code/train_uncertainty_aware_mean_teacher_3D.py
Lines 124 to 125 in 30e05d8
The Dice loss seems to use a 'sum' reduction as follows:
SSL4MIS/code/utils/losses.py
Lines 169 to 177 in 30e05d8
However, the default reduction method for
torch.nn.CrossEntropyLossis 'mean', so the Dice loss is always roughly aboutH*W(*D)times bigger than the CE loss.So, a direct mean of two losses as used in the following code would not be actually the intended average.
SSL4MIS/code/train_uncertainty_aware_mean_teacher_3D.py
Line 171 in 30e05d8
Although I am sure this has minimal effects on most of your SSL methods because it is simply using Dice instead of Dice + CE for the supervisised loss, but still I think it should be checked.