MCPcopy Create free account
hub / github.com/ZhengPeng7/BiRefNet / ClsLoss

Class ClsLoss

loss.py:131–151  ·  view source on GitHub ↗

Auxiliary classification loss for each refined class output.

Source from the content-addressed store, hash-verified

129
130
131class ClsLoss(nn.Module):
132 """
133 Auxiliary classification loss for each refined class output.
134 """
135 def __init__(self):
136 super(ClsLoss, self).__init__()
137 self.config = Config()
138 self.lambdas_cls = self.config.lambdas_cls
139
140 self.criterions_last = {
141 'ce': nn.CrossEntropyLoss()
142 }
143
144 def forward(self, preds, gt):
145 loss = 0.
146 for _, pred_lvl in enumerate(preds):
147 if pred_lvl is None:
148 continue
149 for criterion_name, criterion in self.criterions_last.items():
150 loss += criterion(pred_lvl, gt) * self.lambdas_cls[criterion_name]
151 return loss
152
153
154class PixLoss(nn.Module):

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected