A circular buffer tensor.
| 162 | |
| 163 | |
| 164 | class TensorQueue: |
| 165 | """ |
| 166 | A circular buffer tensor. |
| 167 | """ |
| 168 | def __init__(self, |
| 169 | shape: torch.Size | Sequence[int], |
| 170 | grow_dim: int, |
| 171 | device: torch.device, dtype: torch.dtype |
| 172 | ): |
| 173 | self.device = device |
| 174 | self.grow_dim = grow_dim |
| 175 | self.buf_size = shape[self.grow_dim] |
| 176 | |
| 177 | self.q_start = 0 # Starting point of circular array |
| 178 | self.q_end = 0 # Ending point of circular array |
| 179 | |
| 180 | self._buffer = torch.empty(size=shape, dtype=dtype, device=device) |
| 181 | self._empty = True |
| 182 | |
| 183 | # Special Optimization for scalar write demand |
| 184 | # Explicitly maintein a scalar write cache. When the buffer is read |
| 185 | # or other (non-scalar) write operation is triggered, first write the |
| 186 | # cached scalars in batch operation, then perform the following operations. |
| 187 | self.scalar_array = len(shape) == 1 |
| 188 | self.write_batch = [] |
| 189 | # |
| 190 | |
| 191 | def __repr__(self) -> str: |
| 192 | return f"CircularTensor({self.tensor}, buf_size={self.buf_size}, real_size={len(self)})" |
| 193 | |
| 194 | def __len__(self) -> int: |
| 195 | if self.is_full: return self.buf_size |
| 196 | return (self.q_end - self.q_start) |
| 197 | |
| 198 | def __write_scalar_batch(self) -> None: |
| 199 | if len(self.write_batch) == 0: return |
| 200 | self.__push(torch.tensor(self.write_batch[-self.buf_size:], dtype=self._buffer.dtype, device=self._buffer.device)) |
| 201 | self.write_batch.clear() |
| 202 | |
| 203 | @property |
| 204 | def is_full(self) -> bool: |
| 205 | return self.q_start == self.q_end and (not self._empty) |
| 206 | |
| 207 | @property |
| 208 | def tensor(self) -> torch.Tensor: |
| 209 | self.__write_scalar_batch() |
| 210 | if self._empty: |
| 211 | shape = [_ for _ in self._buffer.shape] |
| 212 | shape[self.grow_dim] = 0 |
| 213 | return torch.zeros(shape, dtype=self._buffer.dtype, device=self._buffer.device) |
| 214 | |
| 215 | if self.q_start < self.q_end: |
| 216 | return self._buffer.narrow_copy(self.grow_dim, self.q_start, self.q_end - self.q_start) |
| 217 | |
| 218 | return torch.cat([ |
| 219 | self._buffer.narrow(self.grow_dim, self.q_start, self.buf_size - self.q_start), |
| 220 | self._buffer.narrow(self.grow_dim, 0, self.q_end) |
| 221 | ], dim=self.grow_dim) |
no outgoing calls