MCPcopy Create free account
hub / github.com/MAC-VO/MAC-VO / TensorQueue

Class TensorQueue

Utility/Extensions/TensorExtension.py:164–273  ·  view source on GitHub ↗

A circular buffer tensor.

Source from the content-addressed store, hash-verified

162
163
164class 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)

Callers 4

test_circular_naiveFunction · 0.90
test_circular_randomFunction · 0.90
test_circular_scalarFunction · 0.90

Calls

no outgoing calls

Tested by 3

test_circular_naiveFunction · 0.72
test_circular_randomFunction · 0.72
test_circular_scalarFunction · 0.72