MCPcopy Create free account
hub / github.com/TPCD/DCCL / ClusterMemory

Class ClusterMemory

project_utils/cluster_memory_utils.py:84–114  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

82
83
84class ClusterMemory(nn.Module, ABC):
85 def __init__(self, num_features, num_samples, temp=0.05, momentum=0.2, use_hard=False, args=None):
86 super(ClusterMemory, self).__init__()
87 self.num_features = num_features
88 self.num_samples = num_samples
89
90 self.momentum = momentum
91 self.temp = temp
92 self.use_hard = use_hard
93 if args is not None:
94 self.device = args.device
95 else:
96 self.device = None
97
98 self.register_buffer('features', torch.zeros(num_samples, num_features))
99
100 def forward(self, inputs, targets):
101 if self.device is not None:
102 inputs = F.normalize(inputs, dim=1).to(self.device)
103 targets = targets.to(self.device)
104 else:
105 inputs = F.normalize(inputs, dim=1).cuda()
106 targets = targets.cuda()
107 if self.use_hard:
108 outputs = cm_hard(inputs, targets, self.features, self.momentum, self.device)
109 else:
110 outputs = cm(inputs, targets, self.features, self.momentum, self.device)
111
112 outputs /= self.temp
113 loss = F.cross_entropy(outputs, targets)
114 return loss

Callers 1

trainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected