MCPcopy Create free account
hub / github.com/TorchSSL/TorchSSL / consistency_loss

Function consistency_loss

models/fixmatch/fixmatch_utils.py:17–43  ·  view source on GitHub ↗
(logits_s, logits_w, name='ce', T=1.0, p_cutoff=0.0, use_hard_labels=True)

Source from the content-addressed store, hash-verified

15
16
17def consistency_loss(logits_s, logits_w, name='ce', T=1.0, p_cutoff=0.0, use_hard_labels=True):
18 assert name in ['ce', 'L2']
19 logits_w = logits_w.detach()
20 if name == 'L2':
21 assert logits_w.size() == logits_s.size()
22 return F.mse_loss(logits_s, logits_w, reduction='mean')
23
24 elif name == 'L2_mask':
25 pass
26
27 elif name == 'ce':
28 pseudo_label = torch.softmax(logits_w, dim=-1)
29 max_probs, max_idx = torch.max(pseudo_label, dim=-1)
30 mask = max_probs.ge(p_cutoff).float()
31 select = max_probs.ge(p_cutoff).long()
32 # strong_prob, strong_idx = torch.max(torch.softmax(logits_s, dim=-1), dim=-1)
33 # strong_select = strong_prob.ge(p_cutoff).long()
34 # select = select * strong_select * (strong_idx == max_idx)
35 if use_hard_labels:
36 masked_loss = ce_loss(logits_s, max_idx, use_hard_labels, reduction='none') * mask
37 else:
38 pseudo_label = torch.softmax(logits_w / T, dim=-1)
39 masked_loss = ce_loss(logits_s, pseudo_label, use_hard_labels) * mask
40 return masked_loss.mean(), mask.mean(), select, max_idx.long()
41
42 else:
43 assert Exception('Not Implemented consistency_loss')
44

Callers 1

trainMethod · 0.70

Calls 1

ce_lossFunction · 0.90

Tested by

no test coverage detected