A memory buffer is a contiguous torch tensor that may combine multiple tensors sharing with the underlying memory. It must have a unique type to support this behavior.
| 24 | |
| 25 | |
| 26 | class MemoryBuffer: |
| 27 | """ |
| 28 | A memory buffer is a contiguous torch tensor that may combine multiple tensors sharing with the underlying |
| 29 | memory. It must have a unique type to support this behavior. |
| 30 | """ |
| 31 | |
| 32 | def __init__(self, numel: int, numel_padded: int, dtype: torch.dtype, source: Optional[torch.Tensor] = None): |
| 33 | self.numel = numel |
| 34 | self.numel_padded = numel_padded |
| 35 | self.dtype = dtype |
| 36 | if source is not None: |
| 37 | self.data = source |
| 38 | else: |
| 39 | self.data = torch.zeros(self.numel_padded, dtype=self.dtype, device=get_device_name(), requires_grad=False) |
| 40 | |
| 41 | def zero(self): |
| 42 | """Reset the buffer to zero.""" |
| 43 | self.data.zero_() |
| 44 | |
| 45 | def get(self, shape, start_index): |
| 46 | """Return a tensor with the input `shape` as a view into the |
| 47 | 1-D data starting at `start_index`.""" |
| 48 | end_index = start_index + shape.numel() |
| 49 | assert end_index <= self.numel, "requested tensor is out of the buffer range." |
| 50 | buffer_tensor = self.data[start_index:end_index] |
| 51 | buffer_tensor = buffer_tensor.view(shape) |
| 52 | return buffer_tensor |
| 53 | |
| 54 | |
| 55 | def calc_padded_numel(shape: torch.Size, dtype: torch.dtype): |
no outgoing calls
no test coverage detected