| 4 | |
| 5 | |
| 6 | class AverageMeter(nn.Module): |
| 7 | def __init__(self, in_shape, max_size): |
| 8 | super(AverageMeter, self).__init__() |
| 9 | self.max_size = max_size |
| 10 | self.current_size = 0 |
| 11 | self.register_buffer("mean", torch.zeros(in_shape, dtype=torch.float32)) |
| 12 | |
| 13 | def update(self, values): |
| 14 | size = values.size()[0] |
| 15 | if size == 0: |
| 16 | return |
| 17 | new_mean = torch.mean(values.float(), dim=0) |
| 18 | size = np.clip(size, 0, self.max_size) |
| 19 | old_size = min(self.max_size - size, self.current_size) |
| 20 | size_sum = old_size + size |
| 21 | self.current_size = size_sum |
| 22 | self.mean = (self.mean * old_size + new_mean * size) / size_sum |
| 23 | |
| 24 | def clear(self): |
| 25 | self.current_size = 0 |
| 26 | self.mean.fill_(0) |
| 27 | |
| 28 | def __len__(self): |
| 29 | return self.current_size |
| 30 | |
| 31 | def get_mean(self): |
| 32 | return self.mean.squeeze(0).cpu().numpy() |
| 33 | |
| 34 | |
| 35 | class TensorAverageMeter: |
nothing calls this directly
no outgoing calls
no test coverage detected