Actual underlying circular buffer push algorithm.
(self, value: torch.Tensor)
| 229 | self.__push(value) |
| 230 | |
| 231 | def __push(self, value: torch.Tensor) -> None: |
| 232 | """ |
| 233 | Actual underlying circular buffer push algorithm. |
| 234 | """ |
| 235 | orig_full = self.is_full |
| 236 | if self._empty and value.size(self.grow_dim) > 0: self._empty = False |
| 237 | |
| 238 | # Truncate the input if it is greater than the entire buffer |
| 239 | if value.size(self.grow_dim) > self.buf_size: |
| 240 | value = value.narrow(self.grow_dim, start=-self.buf_size, length=self.buf_size) |
| 241 | |
| 242 | # Write the value into buffer. At this time it is guaranteed that the value is |
| 243 | # Smaller than the buffer. |
| 244 | value_size = value.size(self.grow_dim) |
| 245 | assert value_size <= self.buf_size, "Impossible case triggered" |
| 246 | |
| 247 | write_to = 0 |
| 248 | |
| 249 | # Segment 1 - self.end => min(self.start, self.buf_size) |
| 250 | seg1_length = min(self.buf_size - self.q_end, value_size - write_to) |
| 251 | self._buffer.narrow(self.grow_dim, start=self.q_end, length=seg1_length).copy_( |
| 252 | value.narrow(self.grow_dim, start=write_to, length=seg1_length) |
| 253 | ) |
| 254 | self.q_end = (self.q_end + seg1_length) % self.buf_size |
| 255 | write_to += seg1_length |
| 256 | if orig_full: |
| 257 | self.q_start = (self.q_start + seg1_length) % self.buf_size |
| 258 | if write_to == value_size: return |
| 259 | |
| 260 | # Segment 2 - self.end => min(self.buf_size) |
| 261 | seg2_length = min(self.buf_size - self.q_end, value_size - write_to) |
| 262 | self._buffer.narrow(self.grow_dim, start=self.q_end, length=seg2_length).copy_( |
| 263 | value.narrow(self.grow_dim, start=write_to, length=seg2_length) |
| 264 | ) |
| 265 | self.q_end = (self.q_end + seg2_length) % self.buf_size |
| 266 | self.q_start = (self.q_start + seg2_length) % self.buf_size |
| 267 | write_to += seg2_length |
| 268 | |
| 269 | assert write_to == value_size, "Must use up all the input values by this point." |
| 270 | |
| 271 | def push_scalar(self, value: int | float) -> None: |
| 272 | assert self.scalar_array, "Can only push scalar to a scalar array (CircularTensor of 1-dimension)" |
no outgoing calls
no test coverage detected