(self, num_features, num_samples, temp=0.05, momentum=0.2, use_hard=False, args=None)
| 83 | |
| 84 | class 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: |
nothing calls this directly
no outgoing calls
no test coverage detected