From dfdd0348b158d2399dc83882cb772bebca64c1ba Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Perceval=20Wajsb=C3=BCrt?= Date: Wed, 10 Sep 2025 17:33:13 +0200 Subject: [PATCH 1/2] feat: add indice and indice ranges mapping functions --- README.md | 53 ++++ changelog.md | 7 + docs/benchmark.md | 97 ++++--- foldedtensor/__init__.py | 366 +++++++++++++++++++++--- foldedtensor/functions.cpp | 535 ++++++++++++++++++++++++++++++++++++ scripts/benchmark.py | 78 +++--- tests/test_folded_tensor.py | 45 +++ tests/test_indices.py | 249 +++++++++++++++++ 8 files changed, 1327 insertions(+), 103 deletions(-) create mode 100644 tests/test_indices.py diff --git a/README.md b/README.md index d6bf614..ceb9d3d 100644 --- a/README.md +++ b/README.md @@ -106,6 +106,59 @@ print(refolded_embedding.shape) # torch.Size([2, 5, 16]) # 2 samples, 5 words max, 16 dims ``` +### Pooling spans + +You can pool variable length spans directly on a refolded view without padding by building flat indices and offsets and then using `embedding_bag`. + +The helper `lengths.make_indices_ranges` expands ranges defined over one or more variable dimensions. + +- `indices` are the flat positions in the refolded tensor viewed as a single dimension +- `offsets` are the start positions of each span within `indices` +- `spans` gives the span id for every expanded position, which can be useful for functions like `torch.index_add` or `torch.index_reduce` + +Example that sums over word spans to produce one vector per span + +```python +import torch +import foldedtensor as ft + +# Build a 4 level tensor with names: first word of the first context is split into three tokens, etc +input_ids = ft.as_folded_tensor( + [ + [ + [[0, 2, 3], [10], [4]], + [[0, 1, 2], [2, 3], [10, 11], [100, 101]], + ], + ], + full_names=("sample", "context", "word", "token"), +).refold( + "token" +) # any refolding is fine + +# Create embeddings from the input ids +embedding = torch.nn.Embedding(2048, 16) +weight = embedding(input_ids) + +# Pool two word spans per the test +# span 1 covers words 0 to 2 -> mean pool over 4 tokens [0, 2, 3, 10] +# span 2 covers words 5 to 7 -> mean pool over 4 tokens [10, 11, 100, 101] +indices, offsets, spans = input_ids.lengths.make_indices_ranges( + begins=(torch.tensor([0, 5]),), + ends=(torch.tensor([2, 7]),), + indice_dims=("word",), +) + +# Sum embeddings over each span +pooled = torch.nn.functional.embedding_bag( + input=indices, + # Flatten embeddings so rows align with flattened token positions + weight=weight.view(-1, weight.size(-1)), + offsets=offsets, + mode="mean", +) +print(pooled) +``` + ## Benchmarks View the comparisons of `foldedtensor` against various alternatives here: [docs/benchmarks](https://github.com/aphp/foldedtensor/blob/main/docs/benchmark.md). diff --git a/changelog.md b/changelog.md index 5450eaa..2737058 100644 --- a/changelog.md +++ b/changelog.md @@ -1,5 +1,12 @@ # Changelog +## Unreleased + +- Add `map_indices` and `make_indices_ranges` with C++ backends and expose `lengths.map_indices` and `lengths.make_indices_ranges` with boundary handling and flat indices with offsets and span ids for pooling with `embedding_bag`. +- Introduce `FoldedTensorLayout` to store `full_names` and `data_dims` with named dimension resolution and helper methods and use it as the `lengths` container for `FoldedTensor` +- Improve `as_folded_tensor` to better infer dims and dtype from nested data and to accept named `data_dims` and better handle names and empty structures +- Benchmark script adds `--cases` to run selected cases and a new case for range based pooling and adjusts outputs + ## v0.4.0 - Fix `storage` torch warning diff --git a/docs/benchmark.md b/docs/benchmark.md index da4ec8a..903d9be 100644 --- a/docs/benchmark.md +++ b/docs/benchmark.md @@ -8,9 +8,9 @@ It compares the performance of `foldedtensor` with various alternatives for padd and working with nested lists and tensors. Environment: -- `torch.__version__ == '2.6.0'` +- `torch.__version__ == '2.8.0'` - `foldedtensor.__version__ == '0.4.0'` -- `python == 3.9.20` +- `python == 3.11.3` - `sys.platform == 'darwin'` @@ -22,13 +22,13 @@ nested_list = make_nested_list(32, (50, 100), (25, 30), value=1) Comparisons: %timeit python_padding(nested_list) -# 100 loops, best of 5: 15.09 ms per loop +# 100 loops, best of 5: 19.02 ms per loop %timeit foldedtensor.as_folded_tensor(nested_list) -# 100 loops, best of 5: 0.73 ms per loop +# 100 loops, best of 5: 0.82 ms per loop ``` -Speedup against best alternative: **20.67x** :rocket: +Speedup against best alternative: **23.24x** :rocket: ## Case 2 (same lengths nested lists) @@ -36,22 +36,22 @@ Speedup against best alternative: **20.67x** :rocket: nested_list = make_nested_list(32, 100, 30, value=1) %timeit torch.tensor(nested_list) -# 100 loops, best of 5: 6.51 ms per loop +# 100 loops, best of 5: 7.86 ms per loop %timeit torch.LongTensor(nested_list) -# 100 loops, best of 5: 2.78 ms per loop +# 100 loops, best of 5: 3.69 ms per loop %timeit python_padding(nested_list) -# 100 loops, best of 5: 18.38 ms per loop +# 100 loops, best of 5: 23.35 ms per loop %timeit torch.nested.nested_tensor([torch.LongTensor(sub) for sub in nested_list]).to_padded_tensor(0) -# 100 loops, best of 5: 3.00 ms per loop +# 100 loops, best of 5: 3.94 ms per loop %timeit foldedtensor.as_folded_tensor(nested_list) -# 100 loops, best of 5: 1.08 ms per loop +# 100 loops, best of 5: 1.18 ms per loop ``` -Speedup against best alternative: **2.58x** :rocket: +Speedup against best alternative: **3.12x** :rocket: ## Case 3 (simple list) @@ -59,19 +59,19 @@ Speedup against best alternative: **2.58x** :rocket: simple_list = make_nested_list(10000, value=1) %timeit torch.tensor(simple_list) -# 100 loops, best of 5: 0.63 ms per loop +# 100 loops, best of 5: 0.77 ms per loop %timeit torch.LongTensor(simple_list) -# 100 loops, best of 5: 0.27 ms per loop +# 100 loops, best of 5: 0.37 ms per loop %timeit python_padding(simple_list) -# 100 loops, best of 5: 0.28 ms per loop +# 100 loops, best of 5: 0.37 ms per loop %timeit foldedtensor.as_folded_tensor(simple_list) -# 100 loops, best of 5: 0.08 ms per loop +# 100 loops, best of 5: 0.10 ms per loop ``` -Speedup against best alternative: **3.32x** :rocket: +Speedup against best alternative: **3.59x** :rocket: ## Case 4 (same lengths nested lists to flat tensor) @@ -79,22 +79,22 @@ Speedup against best alternative: **3.32x** :rocket: nested_list = make_nested_list(32, 100, 30, value=1) %timeit torch.tensor(nested_list).view(-1) -# 100 loops, best of 5: 6.52 ms per loop +# 100 loops, best of 5: 7.83 ms per loop %timeit torch.LongTensor(nested_list).view(-1) -# 100 loops, best of 5: 2.76 ms per loop +# 100 loops, best of 5: 3.68 ms per loop %timeit python_padding(nested_list).view(-1) -# 100 loops, best of 5: 18.62 ms per loop +# 100 loops, best of 5: 23.17 ms per loop %timeit foldedtensor.as_folded_tensor(nested_list).view(-1) -# 100 loops, best of 5: 1.12 ms per loop +# 100 loops, best of 5: 1.19 ms per loop %timeit foldedtensor.as_folded_tensor(nested_list, data_dims=(2,)) -# 100 loops, best of 5: 1.08 ms per loop +# 100 loops, best of 5: 1.16 ms per loop ``` -Speedup against best alternative: **2.47x** :rocket: +Speedup against best alternative: **3.10x** :rocket: ## Case 5 (variable lengths nested lists) to padded embeddings Nested lists with different lengths (second level lists have lengths between 50 and 150). We compare `foldedtensor` with `torch.nested`. @@ -104,24 +104,24 @@ nested_list = make_nested_list(32, (50, 150), 30, value=1) # Padding with 0 %timeit torch.nested.nested_tensor([torch.LongTensor(sub) for sub in nested_list]).to_padded_tensor(0) -# 100 loops, best of 5: 3.02 ms per loop +# 100 loops, best of 5: 4.40 ms per loop %timeit foldedtensor.as_folded_tensor(nested_list).as_tensor() -# 100 loops, best of 5: 1.03 ms per loop +# 100 loops, best of 5: 1.29 ms per loop ``` -Speedup against best alternative: **2.95x** :rocket: +Speedup against best alternative: **3.41x** :rocket: ```python # Padding with 1 %timeit torch.nested.nested_tensor([torch.FloatTensor(sub) for sub in nested_list]).to_padded_tensor(1) -# 100 loops, best of 5: 3.72 ms per loop +# 100 loops, best of 5: 4.77 ms per loop %timeit x = foldedtensor.as_folded_tensor(nested_list); x.masked_fill_(x.mask, 1) -# 100 loops, best of 5: 1.62 ms per loop +# 100 loops, best of 5: 1.65 ms per loop ``` -Speedup against best alternative: **2.30x** :rocket: +Speedup against best alternative: **2.89x** :rocket: ## Case 6 (2d padding) @@ -129,16 +129,47 @@ Speedup against best alternative: **2.30x** :rocket: nested_list = make_nested_list(160, (50, 150), value=1) %timeit python_padding(nested_list) -# 100 loops, best of 5: 1.33 ms per loop +# 100 loops, best of 5: 1.73 ms per loop %timeit torch.nested.nested_tensor([torch.LongTensor(sub) for sub in nested_list]).to_padded_tensor(0) -# 100 loops, best of 5: 1.14 ms per loop +# 100 loops, best of 5: 1.48 ms per loop %timeit torch.nn.utils.rnn.pad_sequence([torch.LongTensor(sub) for sub in nested_list], batch_first=True, padding_value=0) -# 100 loops, best of 5: 0.86 ms per loop +# 100 loops, best of 5: 1.22 ms per loop %timeit foldedtensor.as_folded_tensor(nested_list) -# 100 loops, best of 5: 0.15 ms per loop +# 100 loops, best of 5: 0.18 ms per loop ``` -Speedup against best alternative: **5.88x** :rocket: +Speedup against best alternative: **6.68x** :rocket: + +## Case 7 (summing vectors inside each differently-sized sequence, all concatenated) + +```python +def sum_all_words_per_sample(t): + begins = torch.arange(len(t.lengths[1])) + ends = begins + 1 + indices, offsets, spans = t.lengths.make_indices_ranges( + begins=(begins,), ends=(ends,), indice_dims=(0,) + ) + return torch.nn.functional.embedding_bag( + input=indices, + weight=t.view(-1, t.size(-1)), + offsets=offsets, + mode="sum", + ) + +embedder = torch.nn.Embedding(500, 128) +nested_list = make_nested_list(320, (150, 250), value=1) +ft = foldedtensor.as_folded_tensor(nested_list).refold(1) +ft = embedder(ft) + + +%timeit ft.refold(0, 1).sum(-2) +# 100 loops, best of 5: 3.54 ms per loop + +%timeit sum_all_words_per_sample(ft) +# 100 loops, best of 5: 1.01 ms per loop + +``` +Speedup against pad-then-sum: **3.52x** :rocket: diff --git a/foldedtensor/__init__.py b/foldedtensor/__init__.py index 5499ab8..abe5a49 100644 --- a/foldedtensor/__init__.py +++ b/foldedtensor/__init__.py @@ -8,7 +8,183 @@ import torch from torch.autograd import Function -from . import _C +from . import _C # type: ignore[import] + +Dim = Union[int, str] + + +def map_indices( + indices: Tuple[Sequence[int], ...], + indice_dims: Tuple[int, ...], + lengths: Sequence[Sequence[int]], + data_dims: Tuple[int, ...], + *, + return_tensors: Optional[str] = None, +): + """ + Compute leaf (last-dim) flat indices given indices in other dimensions. + + Parameters + ---------- + indices: Tuple[Sequence[int], ...] + Tuple of index sequences (broadcasted together) describing positions + along `indice_dims`. + indice_dims: Tuple[int, ...] + Names or indices of the addressing dims. + lengths: Sequence[Sequence[int]] + Nested lengths describing the folded structure. + data_dims: Tuple[int, ...] + Names or indices describing the padded layout used for flattening. + return_tensors: Optional[str], optional (default=None) + Return type: "pt" for torch, "np" for numpy, "list" for python list. + + Returns + ------- + Union[List[int], np.ndarray, torch.Tensor] + Returns a list of flat indices compatible with `.view(-1)` of a tensor + refolded with `data_dims`. + """ + D = len(lengths) + if data_dims[-1] != D - 1: + raise ValueError( + "data_dims must end with the last variable dimension (e.g., 'token')" + ) + + orig_shape = None + saw_pt = False + saw_np = False + np_indices: Tuple[np.ndarray, ...] = tuple( + ( + ( + lambda a: ( + (lambda arr: arr.reshape(-1))( + a.detach().cpu().numpy() + if isinstance(a, torch.Tensor) + else (np.asarray(a)) + ) + ) + )(arr) + ) + for arr in indices + ) # type: ignore[arg-type] + + # Track types and original shape from the first array + first = indices[0] + if isinstance(first, torch.Tensor): + saw_pt = True + orig_shape = tuple(first.shape) + else: + arr0 = np.asarray(first) + if arr0.ndim > 1: + orig_shape = tuple(arr0.shape) + saw_np = isinstance(first, np.ndarray) or saw_np + + if len(indice_dims) != len(np_indices): + raise ValueError("indices and indice_dims must have the same length") + + res = _C.map_indices( + lengths, + list(data_dims), + list(indice_dims), + np_indices, + ) + out_np = np.asarray(res) + # Reshape if needed + if orig_shape is not None: + out_np = out_np.reshape(orig_shape) + + if return_tensors == "pt" or return_tensors is None and saw_pt: + return torch.from_numpy(out_np) + if return_tensors == "np" or return_tensors is None and saw_np: + return out_np + return out_np.tolist() + + +def make_indices_ranges( + *, + begins, + ends, + indice_dims, + lengths, + data_dims, + return_tensors: Union[typing.Optional[str], bool] = None, +): + """ + Expand multiple ranges specified along indice_dims into: + - flat indices (compatible with `.view(-1)` of a tensor refolded with `data_dims`), + - start offsets per span, + - and span indices (the span id for each expanded position). + + Parameters use the same conventions as map_indices. `begins` and `ends` are + tuples of 1D tensors or lists corresponding to each dimension in `indice_dims`. + Ranges are half-open: [begin, end), with boundary support when the last + coordinate equals the number of children of its parent. + """ + if not isinstance(begins, (list, tuple)) or not isinstance(ends, (list, tuple)): + raise TypeError("begins and ends must be tuples/lists of arrays") + if len(begins) != len(indice_dims) or len(ends) != len(indice_dims): + raise ValueError("begins/ends must match indice_dims length") + + saw_pt = False + saw_np = False + # Determine original shape from the first begins entry + first_b = begins[0] + if isinstance(first_b, torch.Tensor): + orig_shape = tuple(first_b.shape) + saw_pt = True + else: + arr0 = np.asarray(first_b) + orig_shape = tuple(arr0.shape) if arr0.ndim > 1 else None + saw_np = isinstance(first_b, np.ndarray) or saw_np + + def _to_np1d(x): + nonlocal saw_pt, saw_np + if isinstance(x, torch.Tensor): + saw_pt = True + return x.detach().cpu().numpy().reshape(-1) + a = np.asarray(x) + if isinstance(x, np.ndarray): + saw_np = True + return a.reshape(-1) + + begins_np = [_to_np1d(b) for b in begins] + ends_np = [_to_np1d(e) for e in ends] + + res = _C.make_indices_ranges( + lengths, + list(data_dims), + list(indice_dims), + begins_np, + ends_np, + ) + + indices, offsets, span_indices = res + indices_np = np.asarray(indices) + offsets_np = np.asarray(offsets) + span_indices_np = np.asarray(span_indices) + + # Reshape offsets to original input shape if multi-dimensional + if orig_shape is not None: + offsets_np = offsets_np.reshape(orig_shape) + + if return_tensors == "pt" or return_tensors is None and saw_pt: + return ( + torch.from_numpy(indices_np.astype(np.int64, copy=False)), + torch.from_numpy(offsets_np.astype(np.int64, copy=False)), + torch.from_numpy(span_indices_np.astype(np.int64, copy=False)), + ) + if return_tensors == "np" or return_tensors is None and saw_np: + return ( + indices_np, + offsets_np, + span_indices_np, + ) + return ( + indices_np.astype(np.int64, copy=False).tolist(), + offsets_np.astype(np.int64, copy=False).tolist(), + span_indices_np.astype(np.int64, copy=False).tolist(), + ) + np_to_torch_dtype = { torch.bool: bool, @@ -49,13 +225,115 @@ __version__ = "0.4.0" -class FoldedTensorLengths(UserList): +class FoldedTensorLayout(UserList): + """ + Folded tensor layout information. + """ + + def __init__( + self, + initlist: Optional[Sequence[Sequence[int]]] = None, + *, + data_dims: Optional[Sequence[Union[int, str]]], + full_names: Optional[Sequence[str]], + ) -> None: + super().__init__(initlist or []) + self._full_names: Optional[Tuple[str, ...]] = ( + tuple(full_names) if full_names is not None else None + ) + if self._full_names is not None: + dd = tuple( + d if isinstance(d, int) else self._full_names.index(d) + for d in data_dims + ) + else: + # Accept ints only when no names are provided + dd = tuple(int(d) for d in data_dims) + self._data_dims: Optional[Tuple[int, ...]] = dd + def __hash__(self): return id(self) + @property + def full_names(self) -> Optional[Tuple[str, ...]]: + return self._full_names + + @property + def data_dims(self) -> Optional[Tuple[int, ...]]: + return self._data_dims + + def __getitem__(self, index: Union[int, str]) -> typing.Any: + if isinstance(index, str): + if self._full_names is None: + raise ValueError( + "Cannot resolve named index without full_names in the layout" + ) + try: + index = self._full_names.index(index) + except ValueError as exc: # pragma: no cover + raise ValueError(f"Unknown dimension name {index!r}") from exc + if not isinstance(index, int): # pragma: no cover + raise TypeError("Index must be an int or a str") + return super().__getitem__(index) + + def resolve_dim(self, dim): + if isinstance(dim, tuple): + return tuple(self.resolve_dim(d) for d in dim) + if isinstance(dim, str): + if self._full_names is None: + raise ValueError( + "Cannot resolve named dim without full_names in the layout" + ) + try: + dim = self._full_names.index(dim) + except ValueError as exc: # pragma: no cover + raise ValueError(f"Unknown dimension name {dim!r}") from exc + return int(dim) + + def map_indices( + self, + indices: Tuple[Sequence[int], ...], + indice_dims: Tuple[Union[int, str], ...], + *, + data_dims: Optional[Sequence[Union[int, str]]] = None, + return_tensors: Optional[str] = None, + ): + indice_dims = self.resolve_dim(indice_dims) + data_dims = self.resolve_dim(data_dims or self.data_dims) + + return map_indices( + indices=indices, + indice_dims=indice_dims, + lengths=self, + data_dims=data_dims, + return_tensors=return_tensors, + ) + + def make_indices_ranges( + self, + *, + begins, + ends, + indice_dims, + data_dims: Optional[Sequence[Union[int, str]]] = None, + return_tensors: Optional[str] = None, + ): + # Resolve indice_dims against this layout's names if provided + indice_dims = self.resolve_dim(indice_dims) + data_dims = self.resolve_dim(data_dims or self.data_dims) + + return make_indices_ranges( + begins=begins, + ends=ends, + indice_dims=indice_dims, + lengths=self, + data_dims=data_dims, + return_tensors=return_tensors, + ) -if typing.TYPE_CHECKING: - FoldedTensorLengths = List[List[int]] # noqa: F811 + +# Backward-compatibility alias +FoldedTensorLengths = FoldedTensorLayout # noinspection PyMethodOverriding @@ -88,11 +366,12 @@ def forward( refolded_data.view(-1, *shape_suffix)[indexer] = data.view( -1, *shape_suffix ).index_select(0, self.indexer) + lengths = FoldedTensorLayout( + self.lengths, data_dims=dims, full_names=self.full_names + ) return FoldedTensor( data=refolded_data, - lengths=self.lengths, - data_dims=dims, - full_names=self.full_names, + lengths=lengths, indexer=indexer, ) @@ -146,7 +425,7 @@ def as_folded_tensor( data_dims: Optional[Sequence[Union[int, str]]] = None, full_names: Optional[Sequence[str]] = None, dtype: Optional[torch.dtype] = None, - lengths: Optional[List[List[int]]] = None, + lengths: Optional[Union[FoldedTensorLayout, List[List[int]]]] = None, device: Optional[Union[str, torch.device]] = None, ): """ @@ -169,6 +448,9 @@ def as_folded_tensor( device: Optional[Unit[str, torch.device]] The device of the output tensor """ + if isinstance(lengths, FoldedTensorLayout): + data_dims = lengths.data_dims or data_dims + full_names = lengths.full_names or full_names if full_names is not None: if data_dims is not None: data_dims = tuple( @@ -189,11 +471,10 @@ def as_folded_tensor( f"Shape inferred from lengths is not compatible with data dims: {shape}, " f"{data.shape}, {len(data_dims)}" ) + layout = FoldedTensorLayout(lengths, data_dims=data_dims, full_names=full_names) result = FoldedTensor( data=data, - lengths=FoldedTensorLengths(lengths), - data_dims=data_dims, - full_names=full_names, + lengths=layout, indexer=torch.from_numpy(np_indexer).to(data.device), ) elif isinstance(data, Sequence): @@ -217,11 +498,10 @@ def as_folded_tensor( padded = torch.from_numpy(padded) # In case of empty sequences, lengths are not computed correctly lengths = (list(lengths) + [[0]] * deepness)[:deepness] + layout = FoldedTensorLayout(lengths, data_dims=data_dims, full_names=full_names) result = FoldedTensor( data=padded, - lengths=FoldedTensorLengths(lengths), - data_dims=data_dims, - full_names=full_names, + lengths=layout, indexer=indexer, ) else: @@ -246,8 +526,6 @@ def _postprocess_func_result(result, input): return FoldedTensor( data=result, lengths=input.lengths, - data_dims=input.data_dims, - full_names=input.full_names, indexer=input.indexer, mask=input._mask, ) @@ -281,18 +559,12 @@ class FoldedTensor(torch.Tensor): def __new__( cls, data: torch.Tensor, - lengths: FoldedTensorLengths, - data_dims: Sequence[int], - full_names: Sequence[str], + lengths: FoldedTensorLayout, indexer: torch.Tensor, mask: Optional[torch.Tensor] = None, ): - data_dims = data_dims - full_names = full_names instance = data.as_subclass(cls) instance.lengths = lengths - instance.data_dims = data_dims - instance.full_names = full_names instance.indexer = indexer instance._mask = mask return instance @@ -301,12 +573,18 @@ def with_data(self, data: torch.Tensor): return FoldedTensor( data=data, lengths=self.lengths, - data_dims=self.data_dims, - full_names=self.full_names, indexer=self.indexer, mask=self._mask, ) + @property + def data_dims(self) -> Tuple[int, ...]: + return self.lengths.data_dims + + @property + def full_names(self) -> Optional[Tuple[str, ...]]: + return self.lengths.full_names + @property def mask(self): if self._mask is None: @@ -323,18 +601,14 @@ def as_tensor(self): def to(self, *args, **kwargs): with torch._C.DisableTorchFunction(): - result = super().to(*args, **kwargs) + res = super().to(*args, **kwargs) copy = kwargs.get("copy", False) - non_blocking = kwargs.get("non_blocking", False) + nb = kwargs.get("non_blocking", False) return FoldedTensor( - data=result, + data=res, lengths=self.lengths, - data_dims=self.data_dims, - full_names=self.full_names, - indexer=self.indexer.to( - result.device, copy=copy, non_blocking=non_blocking - ), - mask=self._mask.to(result.device, copy=copy, non_blocking=non_blocking) + indexer=self.indexer.to(res.device, copy=copy, non_blocking=nb), + mask=self._mask.to(res.device, copy=copy, non_blocking=nb) if self._mask is not None else None, ) @@ -361,14 +635,14 @@ def __torch_function__(cls, func, types, args=(), kwargs=None): if isinstance(arg, FoldedTensor): assert ( ft is None or ft.data_dims == arg.data_dims - ), "Cannot perform operation on FoldedTensors with different structures" + ), "Cannot perform operation on FoldedTensors with different layouts" ft = arg elif isinstance(arg, (list, tuple)): for item in arg: if isinstance(item, FoldedTensor): assert ft is None or ft.data_dims == item.data_dims, ( "Cannot perform operation on FoldedTensors with " - "different structures" + "different layouts" ) ft = item @@ -408,12 +682,28 @@ def refold(self, *dims: Union[Sequence[Union[int, str]], int, str]): dim if isinstance(dim, int) else self.full_names.index(dim) for dim in dims ) - except ValueError: + except ValueError: # pragma: no cover raise ValueError( f"Folded tensor with available dimensions {self.full_names} " f"could not be refolded with dimensions {list(dims)}" ) + # Ensure the leaf (last variable) dimension is last in the refolded layout + leaf = len(self.lengths) - 1 + if dims[-1] != leaf: + leaf_name = ( + self.full_names[leaf] if self.full_names is not None else str(leaf) + ) + dim_names = tuple( + self.full_names[d] if self.full_names is not None else str(d) + for d in dims + ) + raise ValueError( + "The last dimension of data_dims must be the last variable " + f"dimension {leaf_name!r} (ie. {leaf}); got data_dims={dim_names} " + f"(ie. {tuple(dims)}" + ) + if dims == self.data_dims: return self @@ -431,8 +721,6 @@ def reduce_foldedtensor(self: FoldedTensor): ( self.data.as_tensor(), self.lengths, - self.data_dims, - self.full_names, self.indexer.clone() if self.indexer.is_shared() and self.indexer.storage().is_cuda else self.indexer, diff --git a/foldedtensor/functions.cpp b/foldedtensor/functions.cpp index 718c69b..d8730b9 100644 --- a/foldedtensor/functions.cpp +++ b/foldedtensor/functions.cpp @@ -306,12 +306,547 @@ static bool init_numpy() { return true; } +static std::vector cumsum(const std::vector &v) { + std::vector out; + out.reserve(v.size() + 1); + out.push_back(0); + int64_t total = 0; + for (auto x : v) { + total += x; + out.push_back(total); + } + return out; +} + +/** + * Compute per-dimension child start offsets (exclusive prefix sums). + * + * For every variable dimension j>0, `lengths[j]` contains, for each parent + * entity at dimension j-1, the number of children at dimension j. The + * exclusive prefix-sum of this array maps a parent global id to the global id + * of its first child at the next dimension. + * + * - starts[j].size() == lengths[j].size() + 1 + * - For a parent global id g at dimension j-1, the first child global id at + * dimension j is `starts[j][g]`, and the number of children is + * `lengths[j][g]`. + * - starts[0] is left empty (unused), since there is no dimension -1. + * + * @param lengths Variable lengths per dimension. For j>0, lengths[j][g] is the + * number of children in dim j for parent g in dim j-1. + * @return For each j, starts[j] = cumsum(lengths[j]) (exclusive prefix sum). + */ +static std::vector> child_start_offsets( + const std::vector> &lengths) { + const size_t D = lengths.size(); + std::vector> starts(D); + for (size_t j = 1; j < D; ++j) { + starts[j] = cumsum(lengths[j]); + } + return starts; +} + +static std::vector> leaf_offsets_per_dim( + const std::vector> &lengths +) { + const size_t D = lengths.size(); + if (D < 2) { + return std::vector>(); + } + // Start with one leaf per word + size_t n_words = 0; + for (auto x : lengths[D - 2]) n_words += x; + std::vector counts(n_words, 1); + + std::vector> offsets(D); + // For words (D-2): [0,1,2,...,n_words] + offsets[D - 2].resize(n_words + 1); + for (size_t i = 0; i < n_words + 1; ++i) offsets[D - 2][i] = (int64_t)i; + + for (int d = (int)D - 3; d >= 0; --d) { + std::vector new_counts; + new_counts.reserve(lengths[d + 1].size()); + auto it = counts.begin(); + for (auto n_children : lengths[d + 1]) { + int64_t s = 0; + for (int64_t k = 0; k < n_children; ++k) { + if (it == counts.end()) break; + s += *it; + ++it; + } + new_counts.push_back(s); + } + counts.swap(new_counts); + offsets[d] = cumsum(counts); + } + return offsets; +} + +/** + * Compute the flat begin index for every leaf under a refolded layout. + * + * Given the nested `lengths` description and the list of data dimensions + * `data_dims` (which must end at the leaf dimension D-1), this function + * simulates iterating leaves (tokens) while incrementing the multi-dimensional + * index over the data layout. It returns, for each leaf (global leaf id), the + * flat index at which that leaf begins in the contiguous, refolded array. + * + * The resulting flat indices are computed using strides derived from the + * maximum extents observed during the simulated iteration of `data_dims`. + * + * @param lengths Variable lengths per dimension + * @param data_dims Contiguous data dimensions in order, must end at D-1. + * @return Vector `begins[leaf_gid]` giving the flat begin offset of each leaf. + * @throws std::invalid_argument if `data_dims` does not end with D-1. + */ +static std::vector begin_idx_per_leaf( + std::vector> lengths, + const std::vector &data_dims) { + const size_t D = lengths.size(); + const size_t n_new = data_dims.size(); + if (n_new == 0) return {}; + if ((size_t)data_dims.back() != D - 1) { + throw std::invalid_argument("data_dims must end with last variable dimension"); + } + + std::vector new_dim_map(D, -1); + for (size_t i = 0; i < n_new; ++i) new_dim_map[data_dims[i]] = (int8_t)i; + + std::vector new_idx(n_new, 0); + std::vector new_shape(n_new, 0); + std::vector offsets(D - 1, 0); + + std::vector, int64_t>> ops; // (idx snapshot, leaf length) + ops.reserve(lengths.back().size()); + + for (auto leaf_len : lengths.back()) { + ops.emplace_back(new_idx, leaf_len); + + new_idx.back() += leaf_len; + if (new_idx.back() > new_shape.back()) new_shape.back() = new_idx.back(); + + int dim = (int)D - 2; + int8_t mapped = new_dim_map[dim]; + if (mapped >= 0) { + new_idx[mapped] += 1; + if (new_idx[mapped] > new_shape[mapped]) new_shape[mapped] = new_idx[mapped]; + for (size_t i = mapped + 1; i < n_new; ++i) new_idx[i] = 0; + } + + for (dim = (int)D - 2; dim >= 0; --dim) { + lengths[dim][offsets[dim]] -= 1; + if (lengths[dim][offsets[dim]] > 0) { + break; + } + offsets[dim] += 1; + if (dim == 0) break; + int next_dim = dim - 1; + int8_t next_mapped = new_dim_map[next_dim]; + if (next_mapped >= 0) { + new_idx[next_mapped] += 1; + if (new_idx[next_mapped] > new_shape[next_mapped]) new_shape[next_mapped] = new_idx[next_mapped]; + for (int8_t i = next_mapped + 1; i < (int8_t)n_new; ++i) new_idx[i] = 0; + } + } + } + + // strides + std::vector strides(n_new, 1); + for (int i = (int)n_new - 2; i >= 0; --i) { + int64_t s = new_shape[i + 1]; + if (s <= 0) s = 1; + strides[i] = strides[i + 1] * s; + } + + std::vector begins; + begins.reserve(ops.size()); + for (auto &op : ops) { + auto &idx = op.first; + int64_t base = 0; + for (size_t i = 0; i + 1 < n_new; ++i) base += idx[i] * strides[i]; + base += idx.back(); + begins.push_back(base); + } + return begins; +} + +/** + * Resolve a (possibly multi-dimensional) coordinate into a flat token index. + * + * The coordinate spans the contiguous variable dimensions given by + * `indice_dims`. Depending on the last addressed dimension, the function + * supports boundary indices (equal to the size) and maps them to the logical + * end position after the last token of the addressed entity/leaf. + * + * Single-dimension addressing rules: + * - If d == D-1 (token dimension): idx in [0, total_tokens] -> begin_of_leaf + offset. + * - If d == D-2 (leaf/word id): idx in [0, total_words] -> begin_of_leaf. + * - Else (higher level): idx in [0, leaf_offs[d].size()-1] -> first token of entity. + * In all cases, idx == size selects the end position after the last token. + * + * Multi-dimension addressing (contiguous): interpret `coord` as offsets within + * the subtree rooted at `indice_dims[0]`, descend using `starts` to compute the + * parent global id, and resolve the last coordinate either to a token offset or + * to the first token of the targeted child and boundary at the last dimension is + * supported analogously. + * + * @param lengths Variable lengths per dimension. + * @param data_dims Data dimensions (must end at D-1). + * @param indice_dims Contiguous addressed variable dimensions. + * @param starts Per-dimension child start offsets: starts[j] = cumsum(lengths[j]). + * @param leaf_offs For each dimension, offsets into the leaf (token) axis. + * @param begins_per_leaf Flat begin index per leaf (from begin_idx_per_leaf). + * @param token_starts Global token cumsum across leaves. + * @param coord Coordinate values aligned with `indice_dims`. + * @return Flat token index (or end position) in the refolded layout. + * @throws std::out_of_range on invalid coordinates beyond the allowed boundary. + */ +// Helper: memoized count of descendants at a target dimension under an entity. +// cache[target_dim][from_dim] is a vector of size = number of entities at from_dim, +// storing the count of target_dim entities under each entity at from_dim. +static int64_t count_descendants_memo( + const std::vector> &lengths, + const std::vector> &starts, + int from_dim, + int64_t gid, + int target_dim, + std::vector>> &cache) { + if (from_dim == target_dim) return 1; // the entity itself counts as 1 at its own dimension + auto &level_cache = cache[target_dim][from_dim]; + if (gid < 0 || gid >= (int64_t)level_cache.size()) return 0; + int64_t val = level_cache[gid]; + if (val >= 0) return val; + // Sum descendant counts over immediate children + int next_dim = from_dim + 1; + int64_t n_children = lengths[next_dim][gid]; + int64_t start = starts[next_dim][gid]; + int64_t total = 0; + for (int64_t i = 0; i < n_children; ++i) { + total += count_descendants_memo(lengths, starts, next_dim, start + i, target_dim, cache); + } + level_cache[gid] = total; + return total; +} + +// Map a flattened offset within descendants at target_dim to a concrete child gid at target_dim. +static int64_t descendant_gid_by_flat_offset( + const std::vector> &lengths, + const std::vector> &starts, + int from_dim, + int64_t gid, + int target_dim, + int64_t offset, + std::vector>> &cache) { + if (from_dim == target_dim) return gid; + int next_dim = from_dim + 1; + int64_t n_children = lengths[next_dim][gid]; + int64_t start = starts[next_dim][gid]; + for (int64_t i = 0; i < n_children; ++i) { + int64_t child_gid = start + i; + int64_t cnt = count_descendants_memo(lengths, starts, next_dim, child_gid, target_dim, cache); + if (offset < cnt) { + return descendant_gid_by_flat_offset(lengths, starts, next_dim, child_gid, target_dim, offset, cache); + } + offset -= cnt; + } + // Should not reach here if offset < total descendants + throw std::out_of_range("Offset beyond descendant count"); +} + +static int64_t compute_flat_index( + const std::vector> &lengths, + const std::vector &data_dims, + const std::vector &indice_dims, + const std::vector> &starts, + const std::vector> &leaf_offs, + const std::vector &begins_per_leaf, + const std::vector &token_starts, + const std::vector &coord) { + const size_t D = lengths.size(); + const int dend = indice_dims.back(); + + // Single-dimension convenience addressing + if (indice_dims.size() == 1) { + const int d = indice_dims[0]; + int64_t idx = coord[0]; + if (d == (int)D - 1) { + // Global token index (pooled across leaves), boundary allowed + int64_t total_tokens = token_starts.back(); + if (idx < 0 || idx > total_tokens) throw std::out_of_range("Token index out of bounds"); + if (idx == total_tokens) { + int64_t last_leaf = (int64_t)lengths.back().size() - 1; + return begins_per_leaf[last_leaf] + lengths.back()[last_leaf]; + } + auto it = std::upper_bound(token_starts.begin(), token_starts.end(), idx); + int64_t leaf = (int64_t)(it - token_starts.begin()) - 1; + int64_t offset = idx - token_starts[leaf]; + return begins_per_leaf[leaf] + offset; + } else if (d == (int)D - 2) { + // Global word index (leaf id), boundary allowed + int64_t total_words = 0; for (auto x : lengths[D - 2]) total_words += x; + if (idx < 0 || idx > total_words) throw std::out_of_range("Word index out of bounds"); + if (idx == total_words) { + int64_t last_leaf = total_words - 1; + return begins_per_leaf[last_leaf] + lengths.back()[last_leaf]; + } + return begins_per_leaf[idx]; + } else { + // Higher-level entity: map to first token of its first leaf, boundary allowed + // Bounds are based on the total number of entities at this level across parents, + // which corresponds to leaf_offs[d].size() - 1, not lengths[d].size(). + int64_t total_entities = (int64_t)leaf_offs[d].size() - 1; + if (idx < 0 || idx > total_entities) throw std::out_of_range("Index out of bounds"); + if (idx == total_entities) { + int64_t leaf_end = leaf_offs[d].back(); + if (leaf_end == 0) return 0; // empty + int64_t last_leaf = leaf_end - 1; + return begins_per_leaf[last_leaf] + lengths.back()[last_leaf]; + } + int64_t leaf_idx = leaf_offs[d][idx]; + return begins_per_leaf[leaf_idx]; + } + } + + // General non-contiguous multi-dimension addressing (flatten intermediate dims) + // Validate strictly increasing dims + for (size_t i = 1; i < indice_dims.size(); ++i) { + if (indice_dims[i] <= indice_dims[i - 1]) { + throw std::invalid_argument("indice_dims must be strictly increasing"); + } + } + + // Build a cache for descendant counts for all target dims that may be addressed + // Prepare cache[target_dim][from_dim][gid] = count, initialized to -1 + std::vector>> cache; + cache.resize(D); + for (size_t t = 0; t < D; ++t) { + cache[t].resize(D); + for (size_t fd = 0; fd < D; ++fd) { + size_t n_entities = 0; + if (fd == D - 1) { + n_entities = lengths.back().size(); + } else if (fd + 1 < D) { + n_entities = lengths[fd + 1].size(); + } + cache[t][fd] = std::vector(n_entities, -1); + } + } + + // Resolve the first coordinate to a global entity id at its dim + int d0 = indice_dims[0]; + int64_t gid = coord[0]; + // Bounds for first coordinate (no boundary allowed except if last dim only, already handled) + if (d0 == (int)D - 1) { + throw std::invalid_argument("First indice_dim cannot be the leaf/token dimension when multiple dims are provided"); + } else if (d0 == (int)D - 2) { + int64_t total_words = 0; for (auto x : lengths[D - 2]) total_words += x; + if (gid < 0 || gid >= total_words) throw std::out_of_range("Index out of bounds at first dimension"); + } else { + int64_t total_entities = (int64_t)leaf_offs[d0].size() - 1; + if (gid < 0 || gid >= total_entities) throw std::out_of_range("Index out of bounds at first dimension"); + } + + if (indice_dims.size() == 2 && dend == (int)D - 1) { + // Common case: [d_parent, token] with flattening across intermediates + int parent_dim = d0; + int64_t last_idx = coord[1]; + // Number of tokens under this parent entity + int64_t token_count = count_descendants_memo(lengths, starts, parent_dim, gid, (int)D - 1, cache); + if (last_idx < 0 || last_idx > token_count) throw std::out_of_range("Token index out of bounds"); + int64_t leaf_begin = leaf_offs[parent_dim][gid]; + int64_t leaf_end = leaf_offs[parent_dim][gid + 1]; + if (last_idx == token_count) { + if (leaf_end == leaf_begin) return begins_per_leaf[leaf_begin]; + int64_t last_leaf = leaf_end - 1; + return begins_per_leaf[last_leaf] + lengths.back()[last_leaf]; + } + // Map token offset to absolute token index across global token_starts + int64_t base_tokens = token_starts[leaf_begin]; + int64_t abs_token = base_tokens + last_idx; + auto it = std::upper_bound(token_starts.begin(), token_starts.end(), abs_token); + int64_t leaf = (int64_t)(it - token_starts.begin()) - 1; + int64_t offset = abs_token - token_starts[leaf]; + return begins_per_leaf[leaf] + offset; + } + + // Traverse successive addressed dims, skipping/flattening intermediates + for (size_t i = 1; i + 1 < indice_dims.size(); ++i) { + int target_dim = indice_dims[i]; + int64_t off = coord[i]; + if (off < 0) throw std::out_of_range("Negative index not allowed"); + int64_t cnt = count_descendants_memo(lengths, starts, indice_dims[i - 1], gid, target_dim, cache); + if (off >= cnt) throw std::out_of_range("Index out of bounds at intermediate dimension"); + gid = descendant_gid_by_flat_offset(lengths, starts, indice_dims[i - 1], gid, target_dim, off, cache); + } + + // Handle last dimension + int prev_dim = indice_dims[indice_dims.size() - 2]; + int64_t last_idx = coord.back(); + if (dend == (int)D - 1) { + // last is token, gid is entity at prev_dim + int64_t token_count = count_descendants_memo(lengths, starts, prev_dim, gid, (int)D - 1, cache); + if (last_idx < 0 || last_idx > token_count) throw std::out_of_range("Token index out of bounds"); + int64_t leaf_begin = leaf_offs[prev_dim][gid]; + int64_t leaf_end = leaf_offs[prev_dim][gid + 1]; + if (last_idx == token_count) { + if (leaf_end == leaf_begin) return begins_per_leaf[leaf_begin]; + int64_t last_leaf = leaf_end - 1; + return begins_per_leaf[last_leaf] + lengths.back()[last_leaf]; + } + int64_t base_tokens = token_starts[leaf_begin]; + int64_t abs_token = base_tokens + last_idx; + auto it = std::upper_bound(token_starts.begin(), token_starts.end(), abs_token); + int64_t leaf = (int64_t)(it - token_starts.begin()) - 1; + int64_t offset = abs_token - token_starts[leaf]; + return begins_per_leaf[leaf] + offset; + } else { + // last is an addressed non-leaf level, select descendant at dend with boundary allowed + int64_t cnt = count_descendants_memo(lengths, starts, prev_dim, gid, dend, cache); + if (last_idx < 0 || last_idx > cnt) throw std::out_of_range("Index out of bounds at last dimension"); + if (last_idx == cnt) { + int64_t leaf_begin = leaf_offs[prev_dim][gid]; + int64_t leaf_end = leaf_offs[prev_dim][gid + 1]; + if (leaf_end == leaf_begin) return begins_per_leaf[leaf_begin]; + int64_t last_leaf = leaf_end - 1; + return begins_per_leaf[last_leaf] + lengths.back()[last_leaf]; + } + int64_t child_gid = descendant_gid_by_flat_offset(lengths, starts, prev_dim, gid, dend, last_idx, cache); + int64_t leaf_idx = (dend == (int)D - 2) ? child_gid : leaf_offs[dend][child_gid]; + return begins_per_leaf[leaf_idx]; + } +} + +static py::array_t map_indices_cpp( + const std::vector> &lengths, + const std::vector &data_dims, + const std::vector &indice_dims, + const std::vector> &indices) { + const size_t D = lengths.size(); + if (D == 0) return py::array_t(0); + + if (data_dims.empty() || (size_t)data_dims.back() != D - 1) { + throw std::invalid_argument("data_dims must end with the last variable dimension"); + } + if (indices.size() != indice_dims.size()) { + throw std::invalid_argument("indices and indice_dims must have the same length"); + } + size_t n = indices.empty() ? 0 : indices[0].size(); + for (auto &v : indices) if (v.size() != n) throw std::invalid_argument("indices must be same length"); + + // Precompute helpers + std::vector begins = begin_idx_per_leaf(lengths, data_dims); + std::vector> leaf_offs = leaf_offsets_per_dim(lengths); + std::vector> starts = child_start_offsets(lengths); + std::vector token_starts = cumsum(lengths.back()); + + // Validate monotonic increasing dims and bounds + if (!indice_dims.empty()) { + for (size_t i = 1; i < indice_dims.size(); ++i) { + if (indice_dims[i] <= indice_dims[i - 1]) { + throw std::invalid_argument("indice_dims must be strictly increasing"); + } + } + if (indice_dims.back() > (int)D - 1) { + throw std::invalid_argument("Final indice_dim must be <= leaf dimension"); + } + } + + py::array_t out(n); + auto *out_ptr = (int64_t *) out.mutable_data(); + std::vector coord(indice_dims.size()); + for (size_t i = 0; i < n; ++i) { + for (size_t j = 0; j < indice_dims.size(); ++j) coord[j] = indices[j][i]; + out_ptr[i] = compute_flat_index(lengths, data_dims, indice_dims, starts, leaf_offs, begins, token_starts, coord); + } + return out; +} + +// Extracted from inline binding: build flat indices for spans between begins and ends +static py::tuple make_indices_ranges_cpp( + const std::vector> &lengths, + const std::vector &data_dims, + const std::vector &indice_dims, + const std::vector> &begins, + const std::vector> &ends) { + const size_t D = lengths.size(); + if (data_dims.empty() || (size_t)data_dims.back() != D - 1) { + throw std::invalid_argument("data_dims must end with the last variable dimension"); + } + if (begins.size() != ends.size() || begins.size() != indice_dims.size()) { + throw std::invalid_argument("begins/ends must match indice_dims length"); + } + size_t n = begins.empty() ? 0 : begins[0].size(); + for (auto &v : begins) if (v.size() != n) throw std::invalid_argument("begins arrays must be same length"); + for (auto &v : ends) if (v.size() != n) throw std::invalid_argument("ends arrays must be same length as begins"); + + // Precompute helpers + std::vector begins_per_leaf = begin_idx_per_leaf(lengths, data_dims); + std::vector> leaf_offs = leaf_offsets_per_dim(lengths); + std::vector> starts = child_start_offsets(lengths); + std::vector token_starts = cumsum(lengths.back()); + + // Validate monotonic increasing dims + if (!indice_dims.empty()) { + for (size_t i = 1; i < indice_dims.size(); ++i) { + if (indice_dims[i] <= indice_dims[i - 1]) { + throw std::invalid_argument("indice_dims must be strictly increasing"); + } + } + if (indice_dims.back() > (int)D - 1) { + throw std::invalid_argument("Final indice_dim must be <= leaf dimension"); + } + } + + // First pass: compute starts and total length + std::vector starts_vec; + starts_vec.reserve(n); + std::vector> be_pairs; + be_pairs.reserve(n); + int64_t total = 0; + for (size_t i = 0; i < n; ++i) { + std::vector bcoord(indice_dims.size()); + std::vector ecoord(indice_dims.size()); + for (size_t j = 0; j < indice_dims.size(); ++j) { + bcoord[j] = begins[j][i]; + ecoord[j] = ends[j][i]; + } + int64_t b = compute_flat_index(lengths, data_dims, indice_dims, starts, leaf_offs, begins_per_leaf, token_starts, bcoord); + int64_t e = compute_flat_index(lengths, data_dims, indice_dims, starts, leaf_offs, begins_per_leaf, token_starts, ecoord); + if (e < b) throw std::invalid_argument("Range end before begin"); + starts_vec.push_back(total); + be_pairs.emplace_back(b, e); + total += (e - b); + } + + // Build outputs + py::array_t indices(total); + auto *ind_ptr = (int64_t *) indices.mutable_data(); + // Also build span indices: the span number for each expanded position + py::array_t span_indices(total); + auto *span_ptr = (int64_t *) span_indices.mutable_data(); + for (size_t i = 0; i < be_pairs.size(); ++i) { + auto &p = be_pairs[i]; + for (int64_t x = p.first; x < p.second; ++x) { + *ind_ptr++ = x; + *span_ptr++ = (int64_t) i; + } + } + py::array_t offsets(starts_vec.size()); + auto *off_ptr = (int64_t *) offsets.mutable_data(); + for (auto s : starts_vec) *off_ptr++ = s; + return py::make_tuple(indices, offsets, span_indices); +} + PYBIND11_MODULE(_C, m) { // Initialize the NumPy API. init_numpy(); m.def("make_refolding_indexer", &make_refolding_indexer, "Build an indexer to refold data into a different shape"); m.def("nested_py_list_to_padded_array", &nested_py_list_to_padded_np_array, "Converts a nested Python list to a padded array"); + m.def("map_indices", &map_indices_cpp, "Maps indices to flat leaf starts with boundary support"); + m.def("make_indices_ranges", &make_indices_ranges_cpp, "Expand ranges between begins and ends into flat indices, start offsets, and span indices"); } +// PARTS TO SIMPLIFY -- END + #pragma clang diagnostic pop diff --git a/scripts/benchmark.py b/scripts/benchmark.py index c03564f..d160dbe 100644 --- a/scripts/benchmark.py +++ b/scripts/benchmark.py @@ -1,4 +1,5 @@ # ruff: noqa: F401, E501 +import argparse import contextlib import random import subprocess @@ -132,7 +133,16 @@ def format_time(dt): if __name__ == "__main__": # fmt: off - cases = [1, 2, 3, 4, 5, 6] + parser = argparse.ArgumentParser(description="Run foldedtensor benchmarks.") + parser.add_argument( + "-c", + "--cases", + type=int, + nargs="*", + help="Space-separated case IDs to run (1-8). Default: all.", + ) + args = parser.parse_args() + cases = args.cases or list(range(1, 9)) if 1 in cases: print("\n## Case 1 (pad variable lengths nested list)\n") @@ -228,44 +238,50 @@ def format_time(dt): print(f"Speedup against best alternative: **{min(alt) / ft_time:.2f}x** :rocket:") if 7 in cases: - # Test case not working yet - - def sum_all_words_per_sample(ft): - lengths = ft.lengths - ids = torch.arange(lengths[0][0]) - for i in range(1, len(lengths)): - ids = torch.repeat_interleave( - ids, - lengths[i], - output_size=len(lengths[i + 1]) - if i < len(lengths) - 1 - else ft.size(len(ft.data_dims) - 1), - ) - - out = torch.zeros(lengths[0][0], ft.shape[-1]) - out.index_add_(source=ft.as_tensor(), dim=0, index=ids) - - return out - - - print("\n## Case 7 (flat sums)\n") + print("\n## Case 7 (summing vectors inside each differently-sized sequence, all concatenated)\n") with block_code(): exec_and_print( + 'def sum_all_words_per_sample(t):\n' + ' begins = torch.arange(len(t.lengths[1]))\n' + ' ends = begins + 1\n' + ' indices, offsets, spans = t.lengths.make_indices_ranges(\n' + ' begins=(begins,), ends=(ends,), indice_dims=(0,)\n' + ' )\n' + ' return torch.nn.functional.embedding_bag(\n' + ' input=indices,\n' + ' weight=t.view(-1, t.size(-1)),\n' + ' offsets=offsets,\n' + ' mode="sum",\n' + ' )\n\n' "embedder = torch.nn.Embedding(500, 128)\n" "nested_list = make_nested_list(320, (150, 250), value=1)\n" - "ft = foldedtensor.as_folded_tensor(nested_list).refold(2)\n" - "nt = torch.nested.nested_tensor([torch.LongTensor(sub) for sub in nested_list])\n" + "ft = foldedtensor.as_folded_tensor(nested_list).refold(1)\n" + #"nt = torch.nested.nested_tensor([torch.LongTensor(sub) for sub in nested_list])\n" "ft = embedder(ft)\n" - "nt = embedder(nt)\n" + #"nt = embedder(nt)\n" ) - - nt_time = timeit("nt.sum(dim=1)") + pd_time = timeit("ft.refold(0, 1).sum(-2)") ft_time = timeit("sum_all_words_per_sample(ft)") - print(f"Speedup against best alternative: **{nt_time / ft_time:.2f}x** :rocket:") + print(f"Speedup against pad-then-sum: **{pd_time / ft_time:.2f}x** :rocket:") + + if 8 in cases: + print("\n## Case 8 (CamemBERT tokenization with padding)\n") + + with block_code(): + exec_and_print( + "from transformers import AutoTokenizer\n" + "tokenizer = AutoTokenizer.from_pretrained('camembert-base')\n" + "texts = [\n" + " ('Le chat est sur le tapis. ' * random.randint(50, 150)).strip()\n" + " for _ in range(64)\n" + "]" + ) + hf_time = timeit("tokenizer(texts, return_tensors='pt', padding=True)") + ft_time = timeit("foldedtensor.as_folded_tensor(tokenizer(texts)['input_ids'])") - # timeit("embedder(ft)") - # timeit("embedder(ft).refold(0, 1)") - # timeit("embedder(nt)") + print( + f"Speedup against baseline: **{hf_time / ft_time:.2f}x** :rocket:" + ) # fmt: on diff --git a/tests/test_folded_tensor.py b/tests/test_folded_tensor.py index a4d2e24..67139a4 100644 --- a/tests/test_folded_tensor.py +++ b/tests/test_folded_tensor.py @@ -446,3 +446,48 @@ def test_missing_dims(): tensor.refold("line", "token") assert "line" in str(e.value) + + +def test_get_lengths(): + tensor = as_folded_tensor( + [ + [0, 1, 2], + [3, 4], + ], + full_names=("sample", "token"), + dtype=torch.long, + ) + assert tensor.lengths == [[2], [3, 2]] + assert tensor.lengths["token"] == [3, 2] + + +def test_recreate_folded_tensor_manually(): + tensor = as_folded_tensor( + [ + [0, 1, 2], + [3, 4], + ], + full_names=("sample", "token"), + dtype=torch.long, + ) + as_folded_tensor( + data=tensor.data, + lengths=tensor.lengths, + data_dims=tensor.data_dims, + full_names=("sample_bis", "token_bis"), + ) + + +def test_fail_on_refold_missing_last_dim(): + tensor = as_folded_tensor( + [ + [[0], [1, 2]], + [[3, 4], [8, 9, 10, 11]], + ], + full_names=("sample", "sent", "word"), + dtype=torch.long, + ) + with pytest.raises(ValueError) as e: + tensor.refold("sent") + + assert "The last dimension" in str(e.value) diff --git a/tests/test_indices.py b/tests/test_indices.py new file mode 100644 index 0000000..94749f8 --- /dev/null +++ b/tests/test_indices.py @@ -0,0 +1,249 @@ +import numpy as np +import torch + +import foldedtensor as ft + + +def build_tensor(): + # Two samples total. First sample mirrors the example in the prompt + # and totals 14 tokens (contexts: 5 and 9). Second sample has one + # small context to ensure strides are unchanged for the (context, token) view. + data = [ + [ + [ + [0, 2, 3], + [10], + [4], + ], + [ + [0, 1, 2], + [2, 3], + [10, 11], + [100, 101], + ], + ], + [ + [ + [7], + [8, 9], + ], + ], + ] + return ft.as_folded_tensor(data, full_names=("sample", "context", "word", "token")) + + +def build_tensor_single_sample(): + # Single sample as in the prompt example + data = [ + [ + [ + [0, 2, 3], + [10], + [4], + ], + [ + [0, 1, 2], + [2, 3], + [10, 11], + [100, 101], + ], + ], + ] + return ft.as_folded_tensor(data, full_names=("sample", "context", "word", "token")) + + +def test_map_indices_flat_unpadded_tokens_by_token(): + t = build_tensor() + + assert t.refold("token").lengths.map_indices( + indices=([0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13],), + indice_dims=("token",), + ) == [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13] + + +def test_map_indices_flat_unpadded_tokens_by_word(): + t = build_tensor() + + assert t.refold("token").lengths.map_indices( + indices=([0, 1, 2, 3, 4, 5, 6, 7],), + indice_dims=("word",), + ) == [0, 3, 4, 5, 8, 10, 12, 14] + + +def test_map_indices_flat_unpadded_tokens_by_context(): + t = build_tensor() + + assert t.refold("token").lengths.map_indices( + indices=([0, 1, 2],), + indice_dims=("context",), + ) == [0, 5, 14] + + +def test_map_indices_flat_unpadded_tokens_by_sample(): + t = build_tensor() + + assert t.refold("token").lengths.map_indices( + indices=([0, 1],), + indice_dims=("sample",), + ) == [0, 14] + + +def test_map_indices_subset_words(): + t = build_tensor() + + assert t.refold("token").lengths.map_indices( + indices=([0, 1, 2, 4, 6],), + indice_dims=("word",), + ) == [0, 3, 4, 8, 12] + + +def test_map_indices_context_word_to_padded_context_token(): + t = build_tensor() + + assert t.refold("context", "token").lengths.map_indices( + indices=([0, 0, 0, 0, 1, 1, 1, 1], [0, 1, 2, 3, 0, 1, 2, 3]), + indice_dims=("context", "word"), + ) == [0, 3, 4, 5, 9, 12, 14, 16] + + +def test_make_indices_ranges_flat_tokens(): + t = build_tensor_single_sample() + + indices, offsets, spans = t.refold("token").lengths.make_indices_ranges( + begins=(torch.as_tensor([0, 0, 1]), torch.as_tensor([0, 1, 2])), + ends=(torch.as_tensor([0, 1, 1]), torch.as_tensor([1, 3, 4])), + indice_dims=("context", "word"), + return_tensors=False, + ) + + assert indices == [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 10, 11, 12, 13] + assert offsets == [0, 3, 12] + assert spans == [0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2] + + +def test_make_indices_ranges_one_dim_token(): + t = build_tensor_single_sample() + indices, offsets, spans = t.refold("token").lengths.make_indices_ranges( + begins=(torch.tensor([0, 3, 12]),), + ends=(torch.tensor([3, 12, 14]),), + indice_dims=("token",), + return_tensors=False, + ) + assert indices == list(range(0, 3)) + list(range(3, 12)) + list(range(12, 14)) + assert offsets == [0, 3, 12] + assert spans == [0] * 3 + [1] * 9 + [2] * 2 + + +def test_make_indices_ranges_one_dim_word(): + t = build_tensor_single_sample() + indices, offsets, spans = t.refold("token").lengths.make_indices_ranges( + begins=(torch.tensor([0, 1, 3]),), + ends=(torch.tensor([1, 3, 7]),), + indice_dims=("word",), + return_tensors=False, + ) + assert indices == list(range(0, 3)) + list(range(3, 5)) + list(range(5, 14)) + assert offsets == [0, 3, 5] + assert spans == [0] * 3 + [1] * 2 + [2] * 9 + + +def test_word_span_mean_pooler_with_embedding_bag_flat_indices(): + t = build_tensor().refold("context", "token") + # 0 -> 2: [[0, 2, 3], [10]] + # 5 -> 7: [[10, 11], [100, 101]] + indices, offsets, spans = t.lengths.make_indices_ranges( + begins=(torch.tensor([0, 5]),), + ends=(torch.tensor([2, 7]),), + indice_dims=("word",), + ) + embeds = t.unsqueeze(-1).expand(-1, -1, 2).float() + res = torch.nn.functional.embedding_bag( + input=indices, + weight=embeds.view(-1, 2), + offsets=offsets, + mode="mean", + ) + assert res.tolist() == [[3.75, 3.75], [55.5, 55.5]] + + +def test_word_span_mean_pooler_with_embedding_bag_multidim_indices(): + t = build_tensor().refold("context", "token") + # 0 -> 2: [[0, 2, 3], [10]] + # 5 -> 7: [[10, 11], [100, 101]] + indices, offsets, spans = t.lengths.make_indices_ranges( + begins=( + torch.tensor([0, 1]), + torch.tensor([0, 2]), + ), + ends=( + torch.tensor([0, 1]), + torch.tensor([2, 4]), + ), + indice_dims=( + "context", + "word", + ), + ) + embeds = t.unsqueeze(-1).expand(-1, -1, 2).float() + res = torch.nn.functional.embedding_bag( + input=indices, + weight=embeds.view(-1, 2), + offsets=offsets, + mode="mean", + ) + assert res.tolist() == [[3.75, 3.75], [55.5, 55.5]] + + +def test_map_indices_format_torch_multidimensional(): + t = build_tensor() + + assert torch.allclose( + t.lengths.map_indices( + indices=( + torch.as_tensor([[0, 1, 2, 3, 4, 5, 6], [7, 8, 9, 10, 11, 12, 13]]), + ), + indice_dims=("token",), + ), + torch.tensor([[0, 1, 2, 3, 6, 12, 13], [14, 15, 16, 18, 19, 21, 22]]), + ) + + +def test_map_indices_format_numpy_multidimensional(): + t = build_tensor() + + assert np.allclose( + t.lengths.map_indices( + indices=(np.asarray([[0, 1, 2, 3, 4, 5, 6], [7, 8, 9, 10, 11, 12, 13]]),), + indice_dims=("token",), + ), + np.asarray([[0, 1, 2, 3, 6, 12, 13], [14, 15, 16, 18, 19, 21, 22]]), + ) + + +def test_make_indices_ranges_format_torch_multidimensional(): + t = build_tensor_single_sample() + + indices, offsets, spans = t.lengths.make_indices_ranges( + begins=(torch.as_tensor([[0, 0, 1]]), torch.as_tensor([[0, 1, 2]])), + ends=(torch.as_tensor([[0, 1, 1]]), torch.as_tensor([[1, 3, 4]])), + indice_dims=("context", "word"), + ) + + assert isinstance(indices, torch.Tensor) + assert isinstance(offsets, torch.Tensor) + assert isinstance(spans, torch.Tensor) + assert offsets.shape == (1, 3) + + +def test_make_indices_ranges_format_numpy_multidimensional(): + t = build_tensor_single_sample() + + indices, offsets, spans = t.lengths.make_indices_ranges( + begins=(np.asarray([[0, 0, 1]]), np.asarray([[0, 1, 2]])), + ends=(np.asarray([[0, 1, 1]]), np.asarray([[1, 3, 4]])), + indice_dims=("context", "word"), + ) + assert isinstance(indices, np.ndarray) + assert isinstance(offsets, np.ndarray) + assert isinstance(spans, np.ndarray) + assert offsets.shape == (1, 3) From 173f50f1e3597b55696259eedfc8cdd6892ad815 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Perceval=20Wajsb=C3=BCrt?= Date: Tue, 8 Sep 2026 16:06:24 +0200 Subject: [PATCH 2/2] feat: add indice and indice ranges mapping functions --- README.md | 55 +-- changelog.md | 8 +- docs/benchmark.md | 97 ++--- foldedtensor/__init__.py | 495 +++++++++++-------------- foldedtensor/functions.cpp | 712 +++++++----------------------------- scripts/benchmark.py | 78 ++-- tests/test_folded_tensor.py | 100 +++-- tests/test_indices.py | 315 ++++++---------- 8 files changed, 596 insertions(+), 1264 deletions(-) diff --git a/README.md b/README.md index ceb9d3d..185ea92 100644 --- a/README.md +++ b/README.md @@ -108,57 +108,34 @@ print(refolded_embedding.shape) ### Pooling spans -You can pool variable length spans directly on a refolded view without padding by building flat indices and offsets and then using `embedding_bag`. - -The helper `lengths.make_indices_ranges` expands ranges defined over one or more variable dimensions. - -- `indices` are the flat positions in the refolded tensor viewed as a single dimension -- `offsets` are the start positions of each span within `indices` -- `spans` gives the span id for every expanded position, which can be useful for functions like `torch.index_add` or `torch.index_reduce` - -Example that sums over word spans to produce one vector per span +`lengths.make_indices_ranges` maps half open spans to storage positions, +excluding padding even when a span crosses rows. It returns the expanded indices, +start offsets and the span id of each selected item. ```python import torch import foldedtensor as ft -# Build a 4 level tensor with names: first word of the first context is split into three tokens, etc -input_ids = ft.as_folded_tensor( - [ - [ - [[0, 2, 3], [10], [4]], - [[0, 1, 2], [2, 3], [10, 11], [100, 101]], - ], - ], - full_names=("sample", "context", "word", "token"), -).refold( - "token" -) # any refolding is fine - -# Create embeddings from the input ids -embedding = torch.nn.Embedding(2048, 16) -weight = embedding(input_ids) - -# Pool two word spans per the test -# span 1 covers words 0 to 2 -> mean pool over 4 tokens [0, 2, 3, 10] -# span 2 covers words 5 to 7 -> mean pool over 4 tokens [10, 11, 100, 101] -indices, offsets, spans = input_ids.lengths.make_indices_ranges( - begins=(torch.tensor([0, 5]),), - ends=(torch.tensor([2, 7]),), +tensor = ft.as_folded_tensor([[1.0, 2.0], [3.0]], full_names=("sample", "word")) +indices, offsets, spans = tensor.lengths.make_indices_ranges( + begins=(torch.tensor([0, 1]),), + ends=(torch.tensor([2, 3]),), indice_dims=("word",), ) - -# Sum embeddings over each span pooled = torch.nn.functional.embedding_bag( - input=indices, - # Flatten embeddings so rows align with flattened token positions - weight=weight.view(-1, weight.size(-1)), - offsets=offsets, + indices, + tensor.as_tensor().reshape(-1, 1), + offsets, mode="mean", ) -print(pooled) +assert pooled.tolist() == [[1.5], [2.5]] ``` +Boundary mapping and expansion run in C++, using prefix offsets from the +sequence lengths and the refolding indexer for padded layouts. Range indices are +computed on CPU and returned on the input device. Embedding gathering and pooling +run on the embedding tensor's device. + ## Benchmarks View the comparisons of `foldedtensor` against various alternatives here: [docs/benchmarks](https://github.com/aphp/foldedtensor/blob/main/docs/benchmark.md). diff --git a/changelog.md b/changelog.md index 2737058..69642b5 100644 --- a/changelog.md +++ b/changelog.md @@ -2,10 +2,10 @@ ## Unreleased -- Add `map_indices` and `make_indices_ranges` with C++ backends and expose `lengths.map_indices` and `lengths.make_indices_ranges` with boundary handling and flat indices with offsets and span ids for pooling with `embedding_bag`. -- Introduce `FoldedTensorLayout` to store `full_names` and `data_dims` with named dimension resolution and helper methods and use it as the `lengths` container for `FoldedTensor` -- Improve `as_folded_tensor` to better infer dims and dtype from nested data and to accept named `data_dims` and better handle names and empty structures -- Benchmark script adds `--cases` to run selected cases and a new case for range based pooling and adjusts outputs +- Add range expansion that returns storage indices without padding +- Preserve empty contexts and words when refolding padded layouts +- Store dimension names and tensor axes in `FoldedTensorLayout` for span pooling before the forward pass +- Keep explicit names and dimensions when recreating a tensor from a layout ## v0.4.0 diff --git a/docs/benchmark.md b/docs/benchmark.md index 903d9be..da4ec8a 100644 --- a/docs/benchmark.md +++ b/docs/benchmark.md @@ -8,9 +8,9 @@ It compares the performance of `foldedtensor` with various alternatives for padd and working with nested lists and tensors. Environment: -- `torch.__version__ == '2.8.0'` +- `torch.__version__ == '2.6.0'` - `foldedtensor.__version__ == '0.4.0'` -- `python == 3.11.3` +- `python == 3.9.20` - `sys.platform == 'darwin'` @@ -22,13 +22,13 @@ nested_list = make_nested_list(32, (50, 100), (25, 30), value=1) Comparisons: %timeit python_padding(nested_list) -# 100 loops, best of 5: 19.02 ms per loop +# 100 loops, best of 5: 15.09 ms per loop %timeit foldedtensor.as_folded_tensor(nested_list) -# 100 loops, best of 5: 0.82 ms per loop +# 100 loops, best of 5: 0.73 ms per loop ``` -Speedup against best alternative: **23.24x** :rocket: +Speedup against best alternative: **20.67x** :rocket: ## Case 2 (same lengths nested lists) @@ -36,22 +36,22 @@ Speedup against best alternative: **23.24x** :rocket: nested_list = make_nested_list(32, 100, 30, value=1) %timeit torch.tensor(nested_list) -# 100 loops, best of 5: 7.86 ms per loop +# 100 loops, best of 5: 6.51 ms per loop %timeit torch.LongTensor(nested_list) -# 100 loops, best of 5: 3.69 ms per loop +# 100 loops, best of 5: 2.78 ms per loop %timeit python_padding(nested_list) -# 100 loops, best of 5: 23.35 ms per loop +# 100 loops, best of 5: 18.38 ms per loop %timeit torch.nested.nested_tensor([torch.LongTensor(sub) for sub in nested_list]).to_padded_tensor(0) -# 100 loops, best of 5: 3.94 ms per loop +# 100 loops, best of 5: 3.00 ms per loop %timeit foldedtensor.as_folded_tensor(nested_list) -# 100 loops, best of 5: 1.18 ms per loop +# 100 loops, best of 5: 1.08 ms per loop ``` -Speedup against best alternative: **3.12x** :rocket: +Speedup against best alternative: **2.58x** :rocket: ## Case 3 (simple list) @@ -59,19 +59,19 @@ Speedup against best alternative: **3.12x** :rocket: simple_list = make_nested_list(10000, value=1) %timeit torch.tensor(simple_list) -# 100 loops, best of 5: 0.77 ms per loop +# 100 loops, best of 5: 0.63 ms per loop %timeit torch.LongTensor(simple_list) -# 100 loops, best of 5: 0.37 ms per loop +# 100 loops, best of 5: 0.27 ms per loop %timeit python_padding(simple_list) -# 100 loops, best of 5: 0.37 ms per loop +# 100 loops, best of 5: 0.28 ms per loop %timeit foldedtensor.as_folded_tensor(simple_list) -# 100 loops, best of 5: 0.10 ms per loop +# 100 loops, best of 5: 0.08 ms per loop ``` -Speedup against best alternative: **3.59x** :rocket: +Speedup against best alternative: **3.32x** :rocket: ## Case 4 (same lengths nested lists to flat tensor) @@ -79,22 +79,22 @@ Speedup against best alternative: **3.59x** :rocket: nested_list = make_nested_list(32, 100, 30, value=1) %timeit torch.tensor(nested_list).view(-1) -# 100 loops, best of 5: 7.83 ms per loop +# 100 loops, best of 5: 6.52 ms per loop %timeit torch.LongTensor(nested_list).view(-1) -# 100 loops, best of 5: 3.68 ms per loop +# 100 loops, best of 5: 2.76 ms per loop %timeit python_padding(nested_list).view(-1) -# 100 loops, best of 5: 23.17 ms per loop +# 100 loops, best of 5: 18.62 ms per loop %timeit foldedtensor.as_folded_tensor(nested_list).view(-1) -# 100 loops, best of 5: 1.19 ms per loop +# 100 loops, best of 5: 1.12 ms per loop %timeit foldedtensor.as_folded_tensor(nested_list, data_dims=(2,)) -# 100 loops, best of 5: 1.16 ms per loop +# 100 loops, best of 5: 1.08 ms per loop ``` -Speedup against best alternative: **3.10x** :rocket: +Speedup against best alternative: **2.47x** :rocket: ## Case 5 (variable lengths nested lists) to padded embeddings Nested lists with different lengths (second level lists have lengths between 50 and 150). We compare `foldedtensor` with `torch.nested`. @@ -104,24 +104,24 @@ nested_list = make_nested_list(32, (50, 150), 30, value=1) # Padding with 0 %timeit torch.nested.nested_tensor([torch.LongTensor(sub) for sub in nested_list]).to_padded_tensor(0) -# 100 loops, best of 5: 4.40 ms per loop +# 100 loops, best of 5: 3.02 ms per loop %timeit foldedtensor.as_folded_tensor(nested_list).as_tensor() -# 100 loops, best of 5: 1.29 ms per loop +# 100 loops, best of 5: 1.03 ms per loop ``` -Speedup against best alternative: **3.41x** :rocket: +Speedup against best alternative: **2.95x** :rocket: ```python # Padding with 1 %timeit torch.nested.nested_tensor([torch.FloatTensor(sub) for sub in nested_list]).to_padded_tensor(1) -# 100 loops, best of 5: 4.77 ms per loop +# 100 loops, best of 5: 3.72 ms per loop %timeit x = foldedtensor.as_folded_tensor(nested_list); x.masked_fill_(x.mask, 1) -# 100 loops, best of 5: 1.65 ms per loop +# 100 loops, best of 5: 1.62 ms per loop ``` -Speedup against best alternative: **2.89x** :rocket: +Speedup against best alternative: **2.30x** :rocket: ## Case 6 (2d padding) @@ -129,47 +129,16 @@ Speedup against best alternative: **2.89x** :rocket: nested_list = make_nested_list(160, (50, 150), value=1) %timeit python_padding(nested_list) -# 100 loops, best of 5: 1.73 ms per loop +# 100 loops, best of 5: 1.33 ms per loop %timeit torch.nested.nested_tensor([torch.LongTensor(sub) for sub in nested_list]).to_padded_tensor(0) -# 100 loops, best of 5: 1.48 ms per loop +# 100 loops, best of 5: 1.14 ms per loop %timeit torch.nn.utils.rnn.pad_sequence([torch.LongTensor(sub) for sub in nested_list], batch_first=True, padding_value=0) -# 100 loops, best of 5: 1.22 ms per loop +# 100 loops, best of 5: 0.86 ms per loop %timeit foldedtensor.as_folded_tensor(nested_list) -# 100 loops, best of 5: 0.18 ms per loop +# 100 loops, best of 5: 0.15 ms per loop ``` -Speedup against best alternative: **6.68x** :rocket: - -## Case 7 (summing vectors inside each differently-sized sequence, all concatenated) - -```python -def sum_all_words_per_sample(t): - begins = torch.arange(len(t.lengths[1])) - ends = begins + 1 - indices, offsets, spans = t.lengths.make_indices_ranges( - begins=(begins,), ends=(ends,), indice_dims=(0,) - ) - return torch.nn.functional.embedding_bag( - input=indices, - weight=t.view(-1, t.size(-1)), - offsets=offsets, - mode="sum", - ) - -embedder = torch.nn.Embedding(500, 128) -nested_list = make_nested_list(320, (150, 250), value=1) -ft = foldedtensor.as_folded_tensor(nested_list).refold(1) -ft = embedder(ft) - - -%timeit ft.refold(0, 1).sum(-2) -# 100 loops, best of 5: 3.54 ms per loop - -%timeit sum_all_words_per_sample(ft) -# 100 loops, best of 5: 1.01 ms per loop - -``` -Speedup against pad-then-sum: **3.52x** :rocket: +Speedup against best alternative: **5.88x** :rocket: diff --git a/foldedtensor/__init__.py b/foldedtensor/__init__.py index abe5a49..a636fef 100644 --- a/foldedtensor/__init__.py +++ b/foldedtensor/__init__.py @@ -1,4 +1,3 @@ -import typing import warnings from collections import UserList from multiprocessing.reduction import ForkingPickler @@ -10,180 +9,96 @@ from . import _C # type: ignore[import] -Dim = Union[int, str] - -def map_indices( - indices: Tuple[Sequence[int], ...], - indice_dims: Tuple[int, ...], - lengths: Sequence[Sequence[int]], - data_dims: Tuple[int, ...], - *, - return_tensors: Optional[str] = None, +def make_indices_ranges( + *, begins, ends, indice_dims, lengths, data_dims, return_tensors=None ): """ - Compute leaf (last-dim) flat indices given indices in other dimensions. + Expand half open ranges into storage indices for span pooling, excluding padding Parameters ---------- - indices: Tuple[Sequence[int], ...] - Tuple of index sequences (broadcasted together) describing positions - along `indice_dims`. - indice_dims: Tuple[int, ...] - Names or indices of the addressing dims. - lengths: Sequence[Sequence[int]] - Nested lengths describing the folded structure. - data_dims: Tuple[int, ...] - Names or indices describing the padded layout used for flattening. - return_tensors: Optional[str], optional (default=None) - Return type: "pt" for torch, "np" for numpy, "list" for python list. + begins, ends + Coordinate arrays broadcast together, one per addressed dimension + indice_dims + Increasing dimension indices used to address range boundaries + lengths + Lengths of nested sequences at each dimension + data_dims + Dimensions represented as tensor axes, ending with the innermost dimension + return_tensors + Output format, pt, np or list, inferred from the input when omitted Returns ------- - Union[List[int], np.ndarray, torch.Tensor] - Returns a list of flat indices compatible with `.view(-1)` of a tensor - refolded with `data_dims`. - """ - D = len(lengths) - if data_dims[-1] != D - 1: - raise ValueError( - "data_dims must end with the last variable dimension (e.g., 'token')" - ) - - orig_shape = None - saw_pt = False - saw_np = False - np_indices: Tuple[np.ndarray, ...] = tuple( - ( - ( - lambda a: ( - (lambda arr: arr.reshape(-1))( - a.detach().cpu().numpy() - if isinstance(a, torch.Tensor) - else (np.asarray(a)) - ) - ) - )(arr) - ) - for arr in indices - ) # type: ignore[arg-type] - - # Track types and original shape from the first array - first = indices[0] - if isinstance(first, torch.Tensor): - saw_pt = True - orig_shape = tuple(first.shape) - else: - arr0 = np.asarray(first) - if arr0.ndim > 1: - orig_shape = tuple(arr0.shape) - saw_np = isinstance(first, np.ndarray) or saw_np - - if len(indice_dims) != len(np_indices): - raise ValueError("indices and indice_dims must have the same length") - - res = _C.map_indices( - lengths, - list(data_dims), - list(indice_dims), - np_indices, + indices, offsets, span_indices + Flat storage indices excluding padding, start offsets with the broadcast + input shape and the span index of each selected element + + Examples + -------- + Expand two overlapping ranges and an empty range in a sequence of six elements + + ```python + indices, offsets, span_indices = make_indices_ranges( + begins=([0, 2, 5],), + ends=([3, 4, 5],), + indice_dims=(0,), + lengths=[[6]], + data_dims=(0,), ) - out_np = np.asarray(res) - # Reshape if needed - if orig_shape is not None: - out_np = out_np.reshape(orig_shape) - - if return_tensors == "pt" or return_tensors is None and saw_pt: - return torch.from_numpy(out_np) - if return_tensors == "np" or return_tensors is None and saw_np: - return out_np - return out_np.tolist() - - -def make_indices_ranges( - *, - begins, - ends, - indice_dims, - lengths, - data_dims, - return_tensors: Union[typing.Optional[str], bool] = None, -): - """ - Expand multiple ranges specified along indice_dims into: - - flat indices (compatible with `.view(-1)` of a tensor refolded with `data_dims`), - - start offsets per span, - - and span indices (the span id for each expanded position). - - Parameters use the same conventions as map_indices. `begins` and `ends` are - tuples of 1D tensors or lists corresponding to each dimension in `indice_dims`. - Ranges are half-open: [begin, end), with boundary support when the last - coordinate equals the number of children of its parent. + # indices: [0, 1, 2, 2, 3] + # ------- ---- + # offsets: [0, 3, 5] + # span_indices: [0, 0, 0, 1, 1] + ``` """ - if not isinstance(begins, (list, tuple)) or not isinstance(ends, (list, tuple)): - raise TypeError("begins and ends must be tuples/lists of arrays") - if len(begins) != len(indice_dims) or len(ends) != len(indice_dims): - raise ValueError("begins/ends must match indice_dims length") - - saw_pt = False - saw_np = False - # Determine original shape from the first begins entry - first_b = begins[0] - if isinstance(first_b, torch.Tensor): - orig_shape = tuple(first_b.shape) - saw_pt = True - else: - arr0 = np.asarray(first_b) - orig_shape = tuple(arr0.shape) if arr0.ndim > 1 else None - saw_np = isinstance(first_b, np.ndarray) or saw_np - - def _to_np1d(x): - nonlocal saw_pt, saw_np - if isinstance(x, torch.Tensor): - saw_pt = True - return x.detach().cpu().numpy().reshape(-1) - a = np.asarray(x) - if isinstance(x, np.ndarray): - saw_np = True - return a.reshape(-1) - - begins_np = [_to_np1d(b) for b in begins] - ends_np = [_to_np1d(e) for e in ends] - - res = _C.make_indices_ranges( + if ( + not indice_dims + or len(begins) != len(indice_dims) + or len(ends) != len(indice_dims) + ): + raise ValueError("begins and ends must match the nonempty indice_dims") + coords = np.broadcast_arrays( + *[ + np.asarray( + x.detach().cpu().numpy() if isinstance(x, torch.Tensor) else x, + dtype=np.int64, + ) + for x in (*begins, *ends) + ] + ) + shape = coords[0].shape + coords = np.stack(coords).reshape(2 * len(indice_dims), -1) + indices, offsets, spans = _C.make_indices_ranges( + coords[: len(indice_dims)], + coords[len(indice_dims) :], + indice_dims, lengths, - list(data_dims), - list(indice_dims), - begins_np, - ends_np, + data_dims, ) - - indices, offsets, span_indices = res - indices_np = np.asarray(indices) - offsets_np = np.asarray(offsets) - span_indices_np = np.asarray(span_indices) - - # Reshape offsets to original input shape if multi-dimensional - if orig_shape is not None: - offsets_np = offsets_np.reshape(orig_shape) - - if return_tensors == "pt" or return_tensors is None and saw_pt: - return ( - torch.from_numpy(indices_np.astype(np.int64, copy=False)), - torch.from_numpy(offsets_np.astype(np.int64, copy=False)), - torch.from_numpy(span_indices_np.astype(np.int64, copy=False)), - ) - if return_tensors == "np" or return_tensors is None and saw_np: - return ( - indices_np, - offsets_np, - span_indices_np, + result = (indices, offsets.reshape(shape), spans) + if ( + return_tensors == "pt" + or return_tensors is None + and isinstance(begins[0], torch.Tensor) + ): + return tuple( + torch.as_tensor( + x, + device=begins[0].device + if isinstance(begins[0], torch.Tensor) + else None, + ) + for x in result ) - return ( - indices_np.astype(np.int64, copy=False).tolist(), - offsets_np.astype(np.int64, copy=False).tolist(), - span_indices_np.astype(np.int64, copy=False).tolist(), - ) + if ( + return_tensors == "np" + or return_tensors is None + and isinstance(begins[0], np.ndarray) + ): + return result + return tuple(x.tolist() for x in result) np_to_torch_dtype = { @@ -227,112 +142,97 @@ def _to_np1d(x): class FoldedTensorLayout(UserList): """ - Folded tensor layout information. + Sequence lengths and tensor dimensions used by refolding and range expansion + + Parameters + ---------- + initlist: Sequence[Sequence[int]] + Lengths of nested sequences at each dimension, ordered by their parents + data_dims: Optional[Sequence[Union[int, str]]] + Dimensions represented as tensor axes, by index or name, defaults to all + dimensions, omitted dimensions are flattened into the following axis + full_names: Optional[Sequence[str]] + Names of all nested dimensions, used to address lengths and ranges by name """ - def __init__( - self, - initlist: Optional[Sequence[Sequence[int]]] = None, - *, - data_dims: Optional[Sequence[Union[int, str]]], - full_names: Optional[Sequence[str]], - ) -> None: - super().__init__(initlist or []) - self._full_names: Optional[Tuple[str, ...]] = ( - tuple(full_names) if full_names is not None else None + def __init__(self, initlist=(), *, data_dims=None, full_names=None): + super().__init__(initlist) + self.full_names = tuple(full_names) if full_names is not None else None + self.data_dims = ( + self.resolve_dims(data_dims) + if data_dims is not None + else tuple(range(len(self))) ) - if self._full_names is not None: - dd = tuple( - d if isinstance(d, int) else self._full_names.index(d) - for d in data_dims - ) - else: - # Accept ints only when no names are provided - dd = tuple(int(d) for d in data_dims) - self._data_dims: Optional[Tuple[int, ...]] = dd + + def __setstate__(self, state): + """ + Restore sequence lengths with defaults for omitted layout metadata + """ + self.__dict__.update( + full_names=None, data_dims=tuple(range(len(state["data"]))) + ) + self.__dict__.update(state) def __hash__(self): return id(self) - @property - def full_names(self) -> Optional[Tuple[str, ...]]: - return self._full_names - - @property - def data_dims(self) -> Optional[Tuple[int, ...]]: - return self._data_dims - - def __getitem__(self, index: Union[int, str]) -> typing.Any: - if isinstance(index, str): - if self._full_names is None: - raise ValueError( - "Cannot resolve named index without full_names in the layout" - ) - try: - index = self._full_names.index(index) - except ValueError as exc: # pragma: no cover - raise ValueError(f"Unknown dimension name {index!r}") from exc - if not isinstance(index, int): # pragma: no cover - raise TypeError("Index must be an int or a str") - return super().__getitem__(index) - - def resolve_dim(self, dim): - if isinstance(dim, tuple): - return tuple(self.resolve_dim(d) for d in dim) - if isinstance(dim, str): - if self._full_names is None: - raise ValueError( - "Cannot resolve named dim without full_names in the layout" - ) - try: - dim = self._full_names.index(dim) - except ValueError as exc: # pragma: no cover - raise ValueError(f"Unknown dimension name {dim!r}") from exc - return int(dim) - - def map_indices( - self, - indices: Tuple[Sequence[int], ...], - indice_dims: Tuple[Union[int, str], ...], - *, - data_dims: Optional[Sequence[Union[int, str]]] = None, - return_tensors: Optional[str] = None, - ): - indice_dims = self.resolve_dim(indice_dims) - data_dims = self.resolve_dim(data_dims or self.data_dims) - - return map_indices( - indices=indices, - indice_dims=indice_dims, - lengths=self, - data_dims=data_dims, - return_tensors=return_tensors, + def __getitem__(self, index): + return self.data[ + self.full_names.index(index) if isinstance(index, str) else index + ] + + def resolve_dims(self, dims): + """ + Resolve dimension names for layout construction and index mapping + """ + return tuple( + self.full_names.index(d) if isinstance(d, str) else d for d in dims ) def make_indices_ranges( - self, - *, - begins, - ends, - indice_dims, - data_dims: Optional[Sequence[Union[int, str]]] = None, - return_tensors: Optional[str] = None, + self, *, begins, ends, indice_dims, data_dims=None, return_tensors=None ): - # Resolve indice_dims against this layout's names if provided - indice_dims = self.resolve_dim(indice_dims) - data_dims = self.resolve_dim(data_dims or self.data_dims) - + """ + Expand ranges into storage indices and span offsets for pooling + + Parameters + ---------- + begins, ends: Sequence[array-like] + Start and end coordinates, one array per addressed dimension, broadcast + together, end coordinates are excluded + indice_dims: Sequence[Union[int, str]] + Increasing dimension indices or names used to address range boundaries, + coordinates in omitted dimensions are flattened + data_dims: Optional[Sequence[Union[int, str]]] + Dimensions represented as tensor axes, defaults to this layout, + must end with the innermost dimension + return_tensors: Optional[str] + Output format, pt, np or list, inferred from the first begin coordinate + array when omitted, tensor outputs use the device of the first begin + coordinate if it is a tensor, or CPU otherwise + + Returns + ------- + indices + Flat storage indices excluding padding, repeated for overlapping ranges + offsets + Start offset of each range in indices, with the broadcast coordinate shape, + empty ranges have repeated offsets + span_indices + Range index for each selected element, using flattened broadcast order + """ return make_indices_ranges( begins=begins, ends=ends, - indice_dims=indice_dims, + indice_dims=self.resolve_dims(indice_dims), lengths=self, - data_dims=data_dims, + data_dims=self.data_dims + if data_dims is None + else self.resolve_dims(data_dims), return_tensors=return_tensors, ) -# Backward-compatibility alias FoldedTensorLengths = FoldedTensorLayout @@ -449,8 +349,8 @@ def as_folded_tensor( The device of the output tensor """ if isinstance(lengths, FoldedTensorLayout): - data_dims = lengths.data_dims or data_dims - full_names = lengths.full_names or full_names + data_dims = lengths.data_dims if data_dims is None else data_dims + full_names = lengths.full_names if full_names is None else full_names if full_names is not None: if data_dims is not None: data_dims = tuple( @@ -544,25 +444,46 @@ class FoldedTensor(torch.Tensor): Parameters ---------- data: torch.Tensor - The data tensor. - lengths: List[List[int]] - The lengths of the sequences of variable size, one list for each dimension. - data_dims: Sequence[int] - The flattened dimensions of the data tensor. The last dim must be the last - variable dimension. - full_names: Sequence[str] - The names of the variable dimensions. + Embedding values or scalar data + lengths: Union[FoldedTensorLayout, List[List[int]]] + Lengths of nested sequences at each dimension with optional layout metadata + data_dims: Optional[Sequence[int]] + Tensor axes ending with the innermost dimension, defaults to the + layout dimensions or all dimensions for plain lengths + full_names: Optional[Sequence[str]] + Names of the nested dimensions, defaults to the layout names + indexer: torch.Tensor + One flat storage row index per element of the innermost dimension, + excluding padding mask: Optional[torch.Tensor] - A mask tensor that indicates which elements of the data tensor are not padded. + Boolean mask identifying storage rows without padding, computed lazily """ def __new__( cls, data: torch.Tensor, - lengths: FoldedTensorLayout, - indexer: torch.Tensor, + lengths: Union[FoldedTensorLayout, List[List[int]]], + data_dims: Optional[Sequence[int]] = None, + full_names: Optional[Sequence[str]] = None, + indexer: Optional[torch.Tensor] = None, mask: Optional[torch.Tensor] = None, ): + if indexer is None: + raise TypeError("FoldedTensor requires an indexer") + if ( + not isinstance(lengths, FoldedTensorLayout) + or data_dims is not None + or full_names is not None + ): + lengths = FoldedTensorLayout( + lengths, + data_dims=data_dims + if data_dims is not None + else getattr(lengths, "data_dims", None), + full_names=full_names + if full_names is not None + else getattr(lengths, "full_names", None), + ) instance = data.as_subclass(cls) instance.lengths = lengths instance.indexer = indexer @@ -581,10 +502,22 @@ def with_data(self, data: torch.Tensor): def data_dims(self) -> Tuple[int, ...]: return self.lengths.data_dims + @data_dims.setter + def data_dims(self, dims): + self.lengths = FoldedTensorLayout( + self.lengths, data_dims=dims, full_names=self.full_names + ) + @property def full_names(self) -> Optional[Tuple[str, ...]]: return self.lengths.full_names + @full_names.setter + def full_names(self, names): + self.lengths = FoldedTensorLayout( + self.lengths, data_dims=self.data_dims, full_names=names + ) + @property def mask(self): if self._mask is None: @@ -601,14 +534,16 @@ def as_tensor(self): def to(self, *args, **kwargs): with torch._C.DisableTorchFunction(): - res = super().to(*args, **kwargs) + result = super().to(*args, **kwargs) copy = kwargs.get("copy", False) - nb = kwargs.get("non_blocking", False) + non_blocking = kwargs.get("non_blocking", False) return FoldedTensor( - data=res, + data=result, lengths=self.lengths, - indexer=self.indexer.to(res.device, copy=copy, non_blocking=nb), - mask=self._mask.to(res.device, copy=copy, non_blocking=nb) + indexer=self.indexer.to( + result.device, copy=copy, non_blocking=non_blocking + ), + mask=self._mask.to(result.device, copy=copy, non_blocking=non_blocking) if self._mask is not None else None, ) @@ -633,16 +568,17 @@ def __torch_function__(cls, func, types, args=(), kwargs=None): ft = None for arg in (*args, *kwargs.values()): if isinstance(arg, FoldedTensor): - assert ( - ft is None or ft.data_dims == arg.data_dims - ), "Cannot perform operation on FoldedTensors with different layouts" + assert ft is None or ft.data_dims == arg.data_dims, ( + "Cannot perform operation on FoldedTensors with " + "different structures" + ) ft = arg elif isinstance(arg, (list, tuple)): for item in arg: if isinstance(item, FoldedTensor): assert ft is None or ft.data_dims == item.data_dims, ( "Cannot perform operation on FoldedTensors with " - "different layouts" + "different structures" ) ft = item @@ -682,26 +618,15 @@ def refold(self, *dims: Union[Sequence[Union[int, str]], int, str]): dim if isinstance(dim, int) else self.full_names.index(dim) for dim in dims ) - except ValueError: # pragma: no cover + except ValueError: raise ValueError( f"Folded tensor with available dimensions {self.full_names} " f"could not be refolded with dimensions {list(dims)}" ) - # Ensure the leaf (last variable) dimension is last in the refolded layout - leaf = len(self.lengths) - 1 - if dims[-1] != leaf: - leaf_name = ( - self.full_names[leaf] if self.full_names is not None else str(leaf) - ) - dim_names = tuple( - self.full_names[d] if self.full_names is not None else str(d) - for d in dims - ) + if not dims or dims[-1] != len(self.lengths) - 1: raise ValueError( - "The last dimension of data_dims must be the last variable " - f"dimension {leaf_name!r} (ie. {leaf}); got data_dims={dim_names} " - f"(ie. {tuple(dims)}" + "The last dimension of data_dims must be the last variable dimension" ) if dims == self.data_dims: @@ -721,6 +646,8 @@ def reduce_foldedtensor(self: FoldedTensor): ( self.data.as_tensor(), self.lengths, + self.data_dims, + self.full_names, self.indexer.clone() if self.indexer.is_shared() and self.indexer.storage().is_cuda else self.indexer, diff --git a/foldedtensor/functions.cpp b/foldedtensor/functions.cpp index d8730b9..6af2d40 100644 --- a/foldedtensor/functions.cpp +++ b/foldedtensor/functions.cpp @@ -16,8 +16,18 @@ std::tuple< std::vector// new shape > make_refolding_indexer( - std::vector> &lengths, - std::vector &new_data_dims) { + const std::vector> &lengths, + const std::vector &new_data_dims) { + if (new_data_dims.empty() || lengths.empty() || new_data_dims.back() != lengths.size() - 1) { + throw py::value_error("data_dims must end with the last variable dimension"); + } + int previous = -1; + for (int dim : new_data_dims) { + if (dim <= previous || dim >= lengths.size()) { + throw py::value_error("data_dims must be strictly increasing dimension indices within lengths"); + } + previous = dim; + } const size_t n_lengths = lengths.size(); const size_t n_new_dims = new_data_dims.size(); @@ -26,59 +36,39 @@ make_refolding_indexer( new_dim_map[new_data_dims[i]] = i; } - std::vector offsets(n_lengths - 1, 0); + std::vector offsets(n_lengths, 0); std::vector new_idx(n_new_dims, 0); - // Operations are tuples of (indexer, length) that we will use as follows: - // new_indexer[offset:offset + length] = flat_index(indexer) + range(length) + // Store the padded start coordinates and length of each innermost sequence std::vector, int64_t>> operations; operations.reserve(lengths.back().size()); long long n_elements = 0; std::vector new_shape(n_new_dims, 0); - for (int length: lengths.back()) { - operations.emplace_back(new_idx, length); - - n_elements += length; - - new_idx.back() += length; - new_shape[n_new_dims - 1] = std::max(new_shape[n_new_dims - 1], new_idx.back()); - - int dim = n_lengths - 2; - int8_t new_mapped_dim = new_dim_map[dim]; - if (new_mapped_dim >= 0) { - new_idx[new_mapped_dim] += 1; - new_shape[new_mapped_dim] = std::max(new_shape[new_mapped_dim], new_idx[new_mapped_dim]); - for (size_t i = new_mapped_dim + 1; i < n_new_dims; i++) { - new_idx[i] = 0; - } + // Visit every parent so empty sequences keep their position in padded layouts + auto visit = [&](auto &&self, size_t dim) -> void { + const auto length = lengths[dim][offsets[dim]++]; + const auto mapped_dim = new_dim_map[dim]; + if (dim == n_lengths - 1) { + operations.emplace_back(new_idx, length); + n_elements += length; + new_idx.back() += length; + new_shape.back() = std::max(new_shape.back(), new_idx.back()); + return; } - - for (dim = n_lengths - 2; dim >= 0; dim--) { - lengths[dim][offsets[dim]] -= 1; - if (lengths[dim][offsets[dim]] > 0) { - break; - } - - offsets[dim] += 1; - - if (dim == 0) { - break; + for (int64_t i = 0; i < length; ++i) { + if (mapped_dim >= 0) { + new_shape[mapped_dim] = std::max(new_shape[mapped_dim], new_idx[mapped_dim] + 1); } - - int next_dim = dim - 1; - int8_t next_new_data_mapped_dim = new_dim_map[next_dim]; - if (next_new_data_mapped_dim >= 0) { - new_idx[next_new_data_mapped_dim] += 1; - new_shape[next_new_data_mapped_dim] = std::max(new_shape[next_new_data_mapped_dim], new_idx[next_new_data_mapped_dim]); - - for (int8_t i = next_new_data_mapped_dim + 1; i < n_new_dims; i++) { - new_idx[i] = 0; - } + self(self, dim + 1); + if (mapped_dim >= 0) { + ++new_idx[mapped_dim]; + std::fill(new_idx.begin() + mapped_dim + 1, new_idx.end(), 0); } } - } - // Init new strides (full of 1, size = n_old_dims) and compute them in reverse - // for data size and new data sizes from n_new_dims and n_old_dims offsets + }; + visit(visit, 0); + + // Row major strides map sequence starts to flat storage positions std::vector new_strides(n_new_dims, 1); for (int i = n_new_dims - 2; i >= 0; i--) { new_strides[i] = new_strides[i + 1] * new_shape[i + 1]; @@ -86,8 +76,8 @@ make_refolding_indexer( auto new_indexer = py::array_t(n_elements); size_t offset = 0; - for (auto operation: operations) { - std::vector &idx = std::get<0>(operation); + for (const auto &operation: operations) { + const auto &idx = std::get<0>(operation); auto length = std::get<1>(operation); int64_t begin_idx = 0; @@ -107,6 +97,98 @@ make_refolding_indexer( } +// Expand range boundaries into storage indices for pooling, excluding padding +std::tuple make_indices_ranges( + py::array_t begins, + py::array_t ends, + const std::vector &indice_dims, + const std::vector> &lengths, + const std::vector &data_dims) { + if (indice_dims.empty() || begins.ndim() != 2 || ends.ndim() != 2 || + begins.shape(0) != indice_dims.size() || ends.shape(0) != indice_dims.size() || + begins.shape(1) != ends.shape(1)) { + throw py::value_error("Coordinate arrays must match indice_dims and the number of ranges"); + } + int previous = -1; + for (int dim : indice_dims) { + if (dim <= previous || dim >= lengths.size()) { + throw py::value_error("indice_dims must be strictly increasing dimension indices within lengths"); + } + previous = dim; + } + if (data_dims.empty() || data_dims.back() != lengths.size() - 1) { + throw py::value_error("data_dims must end with the last variable dimension"); + } + + // Prefix offsets map nested coordinates to boundaries in the flattened tensor + std::vector> prefixes(lengths.size()); + for (size_t dim = 0; dim < lengths.size(); ++dim) { + prefixes[dim].push_back(0); + for (int64_t length : lengths[dim]) { + prefixes[dim].push_back(prefixes[dim].back() + length); + } + } + auto flat_boundary = [&](const py::array_t &coords, size_t span) { + auto values = coords.unchecked<2>(); + int dim = indice_dims.front(); + int64_t pos = values(0, span); + if (pos < 0 || pos > prefixes[dim].back() || + (pos == prefixes[dim].back() && indice_dims.size() > 1)) { + throw py::index_error("Index out of bounds at first dimension"); + } + for (size_t i = 1; i < indice_dims.size(); ++i) { + int64_t start = pos, stop = pos + 1; + for (++dim; dim <= indice_dims[i]; ++dim) { + start = prefixes[dim].at(start); + stop = prefixes[dim].at(stop); + } + dim = indice_dims[i]; + int64_t coord = values(i, span); + if (coord < 0 || coord > stop - start || + (coord == stop - start && i + 1 < indice_dims.size())) { + throw py::index_error("Index out of bounds in addressed dimension"); + } + pos = start + coord; + } + for (++dim; dim < lengths.size(); ++dim) { + pos = prefixes[dim].at(pos); + } + return pos; + }; + + const auto count = begins.shape(1); + std::vector starts(count), stops(count); + py::array_t offsets(count); + auto offset = offsets.mutable_unchecked<1>(); + int64_t total = 0; + for (py::ssize_t i = 0; i < count; ++i) { + starts[i] = flat_boundary(begins, i); + stops[i] = flat_boundary(ends, i); + if (stops[i] < starts[i]) { + throw py::value_error("Range end before begin"); + } + offset(i) = total; + total += stops[i] - starts[i]; + } + + // Map flattened positions to storage rows for padded layouts + py::array_t indexer; + if (data_dims.size() > 1) { + indexer = std::get<0>(make_refolding_indexer(lengths, data_dims)); + } + py::array_t indices(total), spans(total); + auto indices_data = indices.mutable_data(); + auto spans_data = spans.mutable_data(); + for (py::ssize_t i = 0; i < count; ++i) { + for (int64_t pos = starts[i], j = offset(i); pos < stops[i]; ++pos, ++j) { + indices_data[j] = data_dims.size() == 1 ? pos : indexer.data()[pos]; + spans_data[j] = i; + } + } + return {indices, offsets, spans}; +} + + #pragma clang diagnostic push #pragma ide diagnostic ignored "misc-no-recursion" @@ -306,547 +388,13 @@ static bool init_numpy() { return true; } -static std::vector cumsum(const std::vector &v) { - std::vector out; - out.reserve(v.size() + 1); - out.push_back(0); - int64_t total = 0; - for (auto x : v) { - total += x; - out.push_back(total); - } - return out; -} - -/** - * Compute per-dimension child start offsets (exclusive prefix sums). - * - * For every variable dimension j>0, `lengths[j]` contains, for each parent - * entity at dimension j-1, the number of children at dimension j. The - * exclusive prefix-sum of this array maps a parent global id to the global id - * of its first child at the next dimension. - * - * - starts[j].size() == lengths[j].size() + 1 - * - For a parent global id g at dimension j-1, the first child global id at - * dimension j is `starts[j][g]`, and the number of children is - * `lengths[j][g]`. - * - starts[0] is left empty (unused), since there is no dimension -1. - * - * @param lengths Variable lengths per dimension. For j>0, lengths[j][g] is the - * number of children in dim j for parent g in dim j-1. - * @return For each j, starts[j] = cumsum(lengths[j]) (exclusive prefix sum). - */ -static std::vector> child_start_offsets( - const std::vector> &lengths) { - const size_t D = lengths.size(); - std::vector> starts(D); - for (size_t j = 1; j < D; ++j) { - starts[j] = cumsum(lengths[j]); - } - return starts; -} - -static std::vector> leaf_offsets_per_dim( - const std::vector> &lengths -) { - const size_t D = lengths.size(); - if (D < 2) { - return std::vector>(); - } - // Start with one leaf per word - size_t n_words = 0; - for (auto x : lengths[D - 2]) n_words += x; - std::vector counts(n_words, 1); - - std::vector> offsets(D); - // For words (D-2): [0,1,2,...,n_words] - offsets[D - 2].resize(n_words + 1); - for (size_t i = 0; i < n_words + 1; ++i) offsets[D - 2][i] = (int64_t)i; - - for (int d = (int)D - 3; d >= 0; --d) { - std::vector new_counts; - new_counts.reserve(lengths[d + 1].size()); - auto it = counts.begin(); - for (auto n_children : lengths[d + 1]) { - int64_t s = 0; - for (int64_t k = 0; k < n_children; ++k) { - if (it == counts.end()) break; - s += *it; - ++it; - } - new_counts.push_back(s); - } - counts.swap(new_counts); - offsets[d] = cumsum(counts); - } - return offsets; -} - -/** - * Compute the flat begin index for every leaf under a refolded layout. - * - * Given the nested `lengths` description and the list of data dimensions - * `data_dims` (which must end at the leaf dimension D-1), this function - * simulates iterating leaves (tokens) while incrementing the multi-dimensional - * index over the data layout. It returns, for each leaf (global leaf id), the - * flat index at which that leaf begins in the contiguous, refolded array. - * - * The resulting flat indices are computed using strides derived from the - * maximum extents observed during the simulated iteration of `data_dims`. - * - * @param lengths Variable lengths per dimension - * @param data_dims Contiguous data dimensions in order, must end at D-1. - * @return Vector `begins[leaf_gid]` giving the flat begin offset of each leaf. - * @throws std::invalid_argument if `data_dims` does not end with D-1. - */ -static std::vector begin_idx_per_leaf( - std::vector> lengths, - const std::vector &data_dims) { - const size_t D = lengths.size(); - const size_t n_new = data_dims.size(); - if (n_new == 0) return {}; - if ((size_t)data_dims.back() != D - 1) { - throw std::invalid_argument("data_dims must end with last variable dimension"); - } - - std::vector new_dim_map(D, -1); - for (size_t i = 0; i < n_new; ++i) new_dim_map[data_dims[i]] = (int8_t)i; - - std::vector new_idx(n_new, 0); - std::vector new_shape(n_new, 0); - std::vector offsets(D - 1, 0); - - std::vector, int64_t>> ops; // (idx snapshot, leaf length) - ops.reserve(lengths.back().size()); - - for (auto leaf_len : lengths.back()) { - ops.emplace_back(new_idx, leaf_len); - - new_idx.back() += leaf_len; - if (new_idx.back() > new_shape.back()) new_shape.back() = new_idx.back(); - - int dim = (int)D - 2; - int8_t mapped = new_dim_map[dim]; - if (mapped >= 0) { - new_idx[mapped] += 1; - if (new_idx[mapped] > new_shape[mapped]) new_shape[mapped] = new_idx[mapped]; - for (size_t i = mapped + 1; i < n_new; ++i) new_idx[i] = 0; - } - - for (dim = (int)D - 2; dim >= 0; --dim) { - lengths[dim][offsets[dim]] -= 1; - if (lengths[dim][offsets[dim]] > 0) { - break; - } - offsets[dim] += 1; - if (dim == 0) break; - int next_dim = dim - 1; - int8_t next_mapped = new_dim_map[next_dim]; - if (next_mapped >= 0) { - new_idx[next_mapped] += 1; - if (new_idx[next_mapped] > new_shape[next_mapped]) new_shape[next_mapped] = new_idx[next_mapped]; - for (int8_t i = next_mapped + 1; i < (int8_t)n_new; ++i) new_idx[i] = 0; - } - } - } - - // strides - std::vector strides(n_new, 1); - for (int i = (int)n_new - 2; i >= 0; --i) { - int64_t s = new_shape[i + 1]; - if (s <= 0) s = 1; - strides[i] = strides[i + 1] * s; - } - - std::vector begins; - begins.reserve(ops.size()); - for (auto &op : ops) { - auto &idx = op.first; - int64_t base = 0; - for (size_t i = 0; i + 1 < n_new; ++i) base += idx[i] * strides[i]; - base += idx.back(); - begins.push_back(base); - } - return begins; -} - -/** - * Resolve a (possibly multi-dimensional) coordinate into a flat token index. - * - * The coordinate spans the contiguous variable dimensions given by - * `indice_dims`. Depending on the last addressed dimension, the function - * supports boundary indices (equal to the size) and maps them to the logical - * end position after the last token of the addressed entity/leaf. - * - * Single-dimension addressing rules: - * - If d == D-1 (token dimension): idx in [0, total_tokens] -> begin_of_leaf + offset. - * - If d == D-2 (leaf/word id): idx in [0, total_words] -> begin_of_leaf. - * - Else (higher level): idx in [0, leaf_offs[d].size()-1] -> first token of entity. - * In all cases, idx == size selects the end position after the last token. - * - * Multi-dimension addressing (contiguous): interpret `coord` as offsets within - * the subtree rooted at `indice_dims[0]`, descend using `starts` to compute the - * parent global id, and resolve the last coordinate either to a token offset or - * to the first token of the targeted child and boundary at the last dimension is - * supported analogously. - * - * @param lengths Variable lengths per dimension. - * @param data_dims Data dimensions (must end at D-1). - * @param indice_dims Contiguous addressed variable dimensions. - * @param starts Per-dimension child start offsets: starts[j] = cumsum(lengths[j]). - * @param leaf_offs For each dimension, offsets into the leaf (token) axis. - * @param begins_per_leaf Flat begin index per leaf (from begin_idx_per_leaf). - * @param token_starts Global token cumsum across leaves. - * @param coord Coordinate values aligned with `indice_dims`. - * @return Flat token index (or end position) in the refolded layout. - * @throws std::out_of_range on invalid coordinates beyond the allowed boundary. - */ -// Helper: memoized count of descendants at a target dimension under an entity. -// cache[target_dim][from_dim] is a vector of size = number of entities at from_dim, -// storing the count of target_dim entities under each entity at from_dim. -static int64_t count_descendants_memo( - const std::vector> &lengths, - const std::vector> &starts, - int from_dim, - int64_t gid, - int target_dim, - std::vector>> &cache) { - if (from_dim == target_dim) return 1; // the entity itself counts as 1 at its own dimension - auto &level_cache = cache[target_dim][from_dim]; - if (gid < 0 || gid >= (int64_t)level_cache.size()) return 0; - int64_t val = level_cache[gid]; - if (val >= 0) return val; - // Sum descendant counts over immediate children - int next_dim = from_dim + 1; - int64_t n_children = lengths[next_dim][gid]; - int64_t start = starts[next_dim][gid]; - int64_t total = 0; - for (int64_t i = 0; i < n_children; ++i) { - total += count_descendants_memo(lengths, starts, next_dim, start + i, target_dim, cache); - } - level_cache[gid] = total; - return total; -} - -// Map a flattened offset within descendants at target_dim to a concrete child gid at target_dim. -static int64_t descendant_gid_by_flat_offset( - const std::vector> &lengths, - const std::vector> &starts, - int from_dim, - int64_t gid, - int target_dim, - int64_t offset, - std::vector>> &cache) { - if (from_dim == target_dim) return gid; - int next_dim = from_dim + 1; - int64_t n_children = lengths[next_dim][gid]; - int64_t start = starts[next_dim][gid]; - for (int64_t i = 0; i < n_children; ++i) { - int64_t child_gid = start + i; - int64_t cnt = count_descendants_memo(lengths, starts, next_dim, child_gid, target_dim, cache); - if (offset < cnt) { - return descendant_gid_by_flat_offset(lengths, starts, next_dim, child_gid, target_dim, offset, cache); - } - offset -= cnt; - } - // Should not reach here if offset < total descendants - throw std::out_of_range("Offset beyond descendant count"); -} - -static int64_t compute_flat_index( - const std::vector> &lengths, - const std::vector &data_dims, - const std::vector &indice_dims, - const std::vector> &starts, - const std::vector> &leaf_offs, - const std::vector &begins_per_leaf, - const std::vector &token_starts, - const std::vector &coord) { - const size_t D = lengths.size(); - const int dend = indice_dims.back(); - - // Single-dimension convenience addressing - if (indice_dims.size() == 1) { - const int d = indice_dims[0]; - int64_t idx = coord[0]; - if (d == (int)D - 1) { - // Global token index (pooled across leaves), boundary allowed - int64_t total_tokens = token_starts.back(); - if (idx < 0 || idx > total_tokens) throw std::out_of_range("Token index out of bounds"); - if (idx == total_tokens) { - int64_t last_leaf = (int64_t)lengths.back().size() - 1; - return begins_per_leaf[last_leaf] + lengths.back()[last_leaf]; - } - auto it = std::upper_bound(token_starts.begin(), token_starts.end(), idx); - int64_t leaf = (int64_t)(it - token_starts.begin()) - 1; - int64_t offset = idx - token_starts[leaf]; - return begins_per_leaf[leaf] + offset; - } else if (d == (int)D - 2) { - // Global word index (leaf id), boundary allowed - int64_t total_words = 0; for (auto x : lengths[D - 2]) total_words += x; - if (idx < 0 || idx > total_words) throw std::out_of_range("Word index out of bounds"); - if (idx == total_words) { - int64_t last_leaf = total_words - 1; - return begins_per_leaf[last_leaf] + lengths.back()[last_leaf]; - } - return begins_per_leaf[idx]; - } else { - // Higher-level entity: map to first token of its first leaf, boundary allowed - // Bounds are based on the total number of entities at this level across parents, - // which corresponds to leaf_offs[d].size() - 1, not lengths[d].size(). - int64_t total_entities = (int64_t)leaf_offs[d].size() - 1; - if (idx < 0 || idx > total_entities) throw std::out_of_range("Index out of bounds"); - if (idx == total_entities) { - int64_t leaf_end = leaf_offs[d].back(); - if (leaf_end == 0) return 0; // empty - int64_t last_leaf = leaf_end - 1; - return begins_per_leaf[last_leaf] + lengths.back()[last_leaf]; - } - int64_t leaf_idx = leaf_offs[d][idx]; - return begins_per_leaf[leaf_idx]; - } - } - - // General non-contiguous multi-dimension addressing (flatten intermediate dims) - // Validate strictly increasing dims - for (size_t i = 1; i < indice_dims.size(); ++i) { - if (indice_dims[i] <= indice_dims[i - 1]) { - throw std::invalid_argument("indice_dims must be strictly increasing"); - } - } - - // Build a cache for descendant counts for all target dims that may be addressed - // Prepare cache[target_dim][from_dim][gid] = count, initialized to -1 - std::vector>> cache; - cache.resize(D); - for (size_t t = 0; t < D; ++t) { - cache[t].resize(D); - for (size_t fd = 0; fd < D; ++fd) { - size_t n_entities = 0; - if (fd == D - 1) { - n_entities = lengths.back().size(); - } else if (fd + 1 < D) { - n_entities = lengths[fd + 1].size(); - } - cache[t][fd] = std::vector(n_entities, -1); - } - } - - // Resolve the first coordinate to a global entity id at its dim - int d0 = indice_dims[0]; - int64_t gid = coord[0]; - // Bounds for first coordinate (no boundary allowed except if last dim only, already handled) - if (d0 == (int)D - 1) { - throw std::invalid_argument("First indice_dim cannot be the leaf/token dimension when multiple dims are provided"); - } else if (d0 == (int)D - 2) { - int64_t total_words = 0; for (auto x : lengths[D - 2]) total_words += x; - if (gid < 0 || gid >= total_words) throw std::out_of_range("Index out of bounds at first dimension"); - } else { - int64_t total_entities = (int64_t)leaf_offs[d0].size() - 1; - if (gid < 0 || gid >= total_entities) throw std::out_of_range("Index out of bounds at first dimension"); - } - - if (indice_dims.size() == 2 && dend == (int)D - 1) { - // Common case: [d_parent, token] with flattening across intermediates - int parent_dim = d0; - int64_t last_idx = coord[1]; - // Number of tokens under this parent entity - int64_t token_count = count_descendants_memo(lengths, starts, parent_dim, gid, (int)D - 1, cache); - if (last_idx < 0 || last_idx > token_count) throw std::out_of_range("Token index out of bounds"); - int64_t leaf_begin = leaf_offs[parent_dim][gid]; - int64_t leaf_end = leaf_offs[parent_dim][gid + 1]; - if (last_idx == token_count) { - if (leaf_end == leaf_begin) return begins_per_leaf[leaf_begin]; - int64_t last_leaf = leaf_end - 1; - return begins_per_leaf[last_leaf] + lengths.back()[last_leaf]; - } - // Map token offset to absolute token index across global token_starts - int64_t base_tokens = token_starts[leaf_begin]; - int64_t abs_token = base_tokens + last_idx; - auto it = std::upper_bound(token_starts.begin(), token_starts.end(), abs_token); - int64_t leaf = (int64_t)(it - token_starts.begin()) - 1; - int64_t offset = abs_token - token_starts[leaf]; - return begins_per_leaf[leaf] + offset; - } - - // Traverse successive addressed dims, skipping/flattening intermediates - for (size_t i = 1; i + 1 < indice_dims.size(); ++i) { - int target_dim = indice_dims[i]; - int64_t off = coord[i]; - if (off < 0) throw std::out_of_range("Negative index not allowed"); - int64_t cnt = count_descendants_memo(lengths, starts, indice_dims[i - 1], gid, target_dim, cache); - if (off >= cnt) throw std::out_of_range("Index out of bounds at intermediate dimension"); - gid = descendant_gid_by_flat_offset(lengths, starts, indice_dims[i - 1], gid, target_dim, off, cache); - } - - // Handle last dimension - int prev_dim = indice_dims[indice_dims.size() - 2]; - int64_t last_idx = coord.back(); - if (dend == (int)D - 1) { - // last is token, gid is entity at prev_dim - int64_t token_count = count_descendants_memo(lengths, starts, prev_dim, gid, (int)D - 1, cache); - if (last_idx < 0 || last_idx > token_count) throw std::out_of_range("Token index out of bounds"); - int64_t leaf_begin = leaf_offs[prev_dim][gid]; - int64_t leaf_end = leaf_offs[prev_dim][gid + 1]; - if (last_idx == token_count) { - if (leaf_end == leaf_begin) return begins_per_leaf[leaf_begin]; - int64_t last_leaf = leaf_end - 1; - return begins_per_leaf[last_leaf] + lengths.back()[last_leaf]; - } - int64_t base_tokens = token_starts[leaf_begin]; - int64_t abs_token = base_tokens + last_idx; - auto it = std::upper_bound(token_starts.begin(), token_starts.end(), abs_token); - int64_t leaf = (int64_t)(it - token_starts.begin()) - 1; - int64_t offset = abs_token - token_starts[leaf]; - return begins_per_leaf[leaf] + offset; - } else { - // last is an addressed non-leaf level, select descendant at dend with boundary allowed - int64_t cnt = count_descendants_memo(lengths, starts, prev_dim, gid, dend, cache); - if (last_idx < 0 || last_idx > cnt) throw std::out_of_range("Index out of bounds at last dimension"); - if (last_idx == cnt) { - int64_t leaf_begin = leaf_offs[prev_dim][gid]; - int64_t leaf_end = leaf_offs[prev_dim][gid + 1]; - if (leaf_end == leaf_begin) return begins_per_leaf[leaf_begin]; - int64_t last_leaf = leaf_end - 1; - return begins_per_leaf[last_leaf] + lengths.back()[last_leaf]; - } - int64_t child_gid = descendant_gid_by_flat_offset(lengths, starts, prev_dim, gid, dend, last_idx, cache); - int64_t leaf_idx = (dend == (int)D - 2) ? child_gid : leaf_offs[dend][child_gid]; - return begins_per_leaf[leaf_idx]; - } -} - -static py::array_t map_indices_cpp( - const std::vector> &lengths, - const std::vector &data_dims, - const std::vector &indice_dims, - const std::vector> &indices) { - const size_t D = lengths.size(); - if (D == 0) return py::array_t(0); - - if (data_dims.empty() || (size_t)data_dims.back() != D - 1) { - throw std::invalid_argument("data_dims must end with the last variable dimension"); - } - if (indices.size() != indice_dims.size()) { - throw std::invalid_argument("indices and indice_dims must have the same length"); - } - size_t n = indices.empty() ? 0 : indices[0].size(); - for (auto &v : indices) if (v.size() != n) throw std::invalid_argument("indices must be same length"); - - // Precompute helpers - std::vector begins = begin_idx_per_leaf(lengths, data_dims); - std::vector> leaf_offs = leaf_offsets_per_dim(lengths); - std::vector> starts = child_start_offsets(lengths); - std::vector token_starts = cumsum(lengths.back()); - - // Validate monotonic increasing dims and bounds - if (!indice_dims.empty()) { - for (size_t i = 1; i < indice_dims.size(); ++i) { - if (indice_dims[i] <= indice_dims[i - 1]) { - throw std::invalid_argument("indice_dims must be strictly increasing"); - } - } - if (indice_dims.back() > (int)D - 1) { - throw std::invalid_argument("Final indice_dim must be <= leaf dimension"); - } - } - - py::array_t out(n); - auto *out_ptr = (int64_t *) out.mutable_data(); - std::vector coord(indice_dims.size()); - for (size_t i = 0; i < n; ++i) { - for (size_t j = 0; j < indice_dims.size(); ++j) coord[j] = indices[j][i]; - out_ptr[i] = compute_flat_index(lengths, data_dims, indice_dims, starts, leaf_offs, begins, token_starts, coord); - } - return out; -} - -// Extracted from inline binding: build flat indices for spans between begins and ends -static py::tuple make_indices_ranges_cpp( - const std::vector> &lengths, - const std::vector &data_dims, - const std::vector &indice_dims, - const std::vector> &begins, - const std::vector> &ends) { - const size_t D = lengths.size(); - if (data_dims.empty() || (size_t)data_dims.back() != D - 1) { - throw std::invalid_argument("data_dims must end with the last variable dimension"); - } - if (begins.size() != ends.size() || begins.size() != indice_dims.size()) { - throw std::invalid_argument("begins/ends must match indice_dims length"); - } - size_t n = begins.empty() ? 0 : begins[0].size(); - for (auto &v : begins) if (v.size() != n) throw std::invalid_argument("begins arrays must be same length"); - for (auto &v : ends) if (v.size() != n) throw std::invalid_argument("ends arrays must be same length as begins"); - - // Precompute helpers - std::vector begins_per_leaf = begin_idx_per_leaf(lengths, data_dims); - std::vector> leaf_offs = leaf_offsets_per_dim(lengths); - std::vector> starts = child_start_offsets(lengths); - std::vector token_starts = cumsum(lengths.back()); - - // Validate monotonic increasing dims - if (!indice_dims.empty()) { - for (size_t i = 1; i < indice_dims.size(); ++i) { - if (indice_dims[i] <= indice_dims[i - 1]) { - throw std::invalid_argument("indice_dims must be strictly increasing"); - } - } - if (indice_dims.back() > (int)D - 1) { - throw std::invalid_argument("Final indice_dim must be <= leaf dimension"); - } - } - - // First pass: compute starts and total length - std::vector starts_vec; - starts_vec.reserve(n); - std::vector> be_pairs; - be_pairs.reserve(n); - int64_t total = 0; - for (size_t i = 0; i < n; ++i) { - std::vector bcoord(indice_dims.size()); - std::vector ecoord(indice_dims.size()); - for (size_t j = 0; j < indice_dims.size(); ++j) { - bcoord[j] = begins[j][i]; - ecoord[j] = ends[j][i]; - } - int64_t b = compute_flat_index(lengths, data_dims, indice_dims, starts, leaf_offs, begins_per_leaf, token_starts, bcoord); - int64_t e = compute_flat_index(lengths, data_dims, indice_dims, starts, leaf_offs, begins_per_leaf, token_starts, ecoord); - if (e < b) throw std::invalid_argument("Range end before begin"); - starts_vec.push_back(total); - be_pairs.emplace_back(b, e); - total += (e - b); - } - - // Build outputs - py::array_t indices(total); - auto *ind_ptr = (int64_t *) indices.mutable_data(); - // Also build span indices: the span number for each expanded position - py::array_t span_indices(total); - auto *span_ptr = (int64_t *) span_indices.mutable_data(); - for (size_t i = 0; i < be_pairs.size(); ++i) { - auto &p = be_pairs[i]; - for (int64_t x = p.first; x < p.second; ++x) { - *ind_ptr++ = x; - *span_ptr++ = (int64_t) i; - } - } - py::array_t offsets(starts_vec.size()); - auto *off_ptr = (int64_t *) offsets.mutable_data(); - for (auto s : starts_vec) *off_ptr++ = s; - return py::make_tuple(indices, offsets, span_indices); -} - PYBIND11_MODULE(_C, m) { // Initialize the NumPy API. init_numpy(); + m.def("make_indices_ranges", &make_indices_ranges, "Expand ranges into storage indices excluding padding"); m.def("make_refolding_indexer", &make_refolding_indexer, "Build an indexer to refold data into a different shape"); m.def("nested_py_list_to_padded_array", &nested_py_list_to_padded_np_array, "Converts a nested Python list to a padded array"); - m.def("map_indices", &map_indices_cpp, "Maps indices to flat leaf starts with boundary support"); - m.def("make_indices_ranges", &make_indices_ranges_cpp, "Expand ranges between begins and ends into flat indices, start offsets, and span indices"); } -// PARTS TO SIMPLIFY -- END - #pragma clang diagnostic pop diff --git a/scripts/benchmark.py b/scripts/benchmark.py index d160dbe..c03564f 100644 --- a/scripts/benchmark.py +++ b/scripts/benchmark.py @@ -1,5 +1,4 @@ # ruff: noqa: F401, E501 -import argparse import contextlib import random import subprocess @@ -133,16 +132,7 @@ def format_time(dt): if __name__ == "__main__": # fmt: off - parser = argparse.ArgumentParser(description="Run foldedtensor benchmarks.") - parser.add_argument( - "-c", - "--cases", - type=int, - nargs="*", - help="Space-separated case IDs to run (1-8). Default: all.", - ) - args = parser.parse_args() - cases = args.cases or list(range(1, 9)) + cases = [1, 2, 3, 4, 5, 6] if 1 in cases: print("\n## Case 1 (pad variable lengths nested list)\n") @@ -238,50 +228,44 @@ def format_time(dt): print(f"Speedup against best alternative: **{min(alt) / ft_time:.2f}x** :rocket:") if 7 in cases: - print("\n## Case 7 (summing vectors inside each differently-sized sequence, all concatenated)\n") + # Test case not working yet + + def sum_all_words_per_sample(ft): + lengths = ft.lengths + ids = torch.arange(lengths[0][0]) + for i in range(1, len(lengths)): + ids = torch.repeat_interleave( + ids, + lengths[i], + output_size=len(lengths[i + 1]) + if i < len(lengths) - 1 + else ft.size(len(ft.data_dims) - 1), + ) + + out = torch.zeros(lengths[0][0], ft.shape[-1]) + out.index_add_(source=ft.as_tensor(), dim=0, index=ids) + + return out + + + print("\n## Case 7 (flat sums)\n") with block_code(): exec_and_print( - 'def sum_all_words_per_sample(t):\n' - ' begins = torch.arange(len(t.lengths[1]))\n' - ' ends = begins + 1\n' - ' indices, offsets, spans = t.lengths.make_indices_ranges(\n' - ' begins=(begins,), ends=(ends,), indice_dims=(0,)\n' - ' )\n' - ' return torch.nn.functional.embedding_bag(\n' - ' input=indices,\n' - ' weight=t.view(-1, t.size(-1)),\n' - ' offsets=offsets,\n' - ' mode="sum",\n' - ' )\n\n' "embedder = torch.nn.Embedding(500, 128)\n" "nested_list = make_nested_list(320, (150, 250), value=1)\n" - "ft = foldedtensor.as_folded_tensor(nested_list).refold(1)\n" - #"nt = torch.nested.nested_tensor([torch.LongTensor(sub) for sub in nested_list])\n" + "ft = foldedtensor.as_folded_tensor(nested_list).refold(2)\n" + "nt = torch.nested.nested_tensor([torch.LongTensor(sub) for sub in nested_list])\n" "ft = embedder(ft)\n" - #"nt = embedder(nt)\n" + "nt = embedder(nt)\n" ) - pd_time = timeit("ft.refold(0, 1).sum(-2)") - ft_time = timeit("sum_all_words_per_sample(ft)") - print(f"Speedup against pad-then-sum: **{pd_time / ft_time:.2f}x** :rocket:") - - if 8 in cases: - print("\n## Case 8 (CamemBERT tokenization with padding)\n") + nt_time = timeit("nt.sum(dim=1)") + ft_time = timeit("sum_all_words_per_sample(ft)") - with block_code(): - exec_and_print( - "from transformers import AutoTokenizer\n" - "tokenizer = AutoTokenizer.from_pretrained('camembert-base')\n" - "texts = [\n" - " ('Le chat est sur le tapis. ' * random.randint(50, 150)).strip()\n" - " for _ in range(64)\n" - "]" - ) - hf_time = timeit("tokenizer(texts, return_tensors='pt', padding=True)") - ft_time = timeit("foldedtensor.as_folded_tensor(tokenizer(texts)['input_ids'])") + print(f"Speedup against best alternative: **{nt_time / ft_time:.2f}x** :rocket:") - print( - f"Speedup against baseline: **{hf_time / ft_time:.2f}x** :rocket:" - ) + # timeit("embedder(ft)") + # timeit("embedder(ft).refold(0, 1)") + # timeit("embedder(nt)") # fmt: on diff --git a/tests/test_folded_tensor.py b/tests/test_folded_tensor.py index 67139a4..1e93b45 100644 --- a/tests/test_folded_tensor.py +++ b/tests/test_folded_tensor.py @@ -1,7 +1,14 @@ +import pickle + import pytest import torch -from foldedtensor import FoldedTensor, as_folded_tensor +from foldedtensor import ( + FoldedTensor, + FoldedTensorLengths, + as_folded_tensor, + reduce_foldedtensor, +) def test_as_folded_tensor_from_nested_list(): @@ -349,7 +356,7 @@ def test_no_data_dims(): def test_as_tensor(ft): tensor = ft.as_tensor() - assert type(tensor) == torch.Tensor + assert type(tensor) is torch.Tensor assert tensor.shape == (2, 5, 2) assert tensor.storage().data_ptr() == ft.storage().data_ptr() @@ -448,46 +455,61 @@ def test_missing_dims(): assert "line" in str(e.value) -def test_get_lengths(): - tensor = as_folded_tensor( - [ - [0, 1, 2], - [3, 4], - ], - full_names=("sample", "token"), - dtype=torch.long, - ) +def test_layout_metadata(): + tensor = as_folded_tensor([[0, 1, 2], [3, 4]], full_names=("sample", "token")) assert tensor.lengths == [[2], [3, 2]] assert tensor.lengths["token"] == [3, 2] - - -def test_recreate_folded_tensor_manually(): - tensor = as_folded_tensor( - [ - [0, 1, 2], - [3, 4], - ], - full_names=("sample", "token"), - dtype=torch.long, - ) - as_folded_tensor( - data=tensor.data, + recreated = as_folded_tensor(tensor.as_tensor(), lengths=tensor.lengths) + renamed = as_folded_tensor( + tensor.as_tensor(), lengths=tensor.lengths, - data_dims=tensor.data_dims, full_names=("sample_bis", "token_bis"), ) - - -def test_fail_on_refold_missing_last_dim(): - tensor = as_folded_tensor( - [ - [[0], [1, 2]], - [[3, 4], [8, 9, 10, 11]], - ], - full_names=("sample", "sent", "word"), - dtype=torch.long, + assert recreated.full_names == tensor.full_names + assert renamed.full_names == ("sample_bis", "token_bis") + assert renamed.tolist() == tensor.tolist() + with pytest.raises(ValueError, match="The last dimension"): + tensor.refold("sample") + + +@pytest.mark.parametrize("positional", [False, True]) +def test_constructor_dimensions(positional): + source = as_folded_tensor([[1.0, 2.0], [3.0]], full_names=("sample", "word")) + args = dict( + data=source.as_tensor(), + lengths=list(source.lengths), + data_dims=(0, 1), + full_names=("sample", "word"), + indexer=source.indexer, + mask=source.mask, ) - with pytest.raises(ValueError) as e: - tensor.refold("sent") - - assert "The last dimension" in str(e.value) + tensor = FoldedTensor(*args.values()) if positional else FoldedTensor(**args) + constructor, state = reduce_foldedtensor(tensor) + for restored in (tensor, constructor(*state)): + assert restored.refold("word").tolist() == [1.0, 2.0, 3.0] + assert restored.lengths.full_names == ("sample", "word") + assert restored.mask.tolist() == [[True, True], [True, False]] + + +def test_pickle_tensor_attributes(): + # Tensor attributes supply dimensions when serialized lengths contain only counts + source = as_folded_tensor([[1.0, 2.0], [3.0]], full_names=("sample", "word")) + lengths = FoldedTensorLengths.__new__(FoldedTensorLengths) + lengths.__dict__["data"] = list(source.lengths) + tensor = source.as_tensor().as_subclass(FoldedTensor) + tensor.__dict__.update( + lengths=lengths, + data_dims=(0, 1), + full_names=("sample", "word"), + indexer=source.indexer, + _mask=source.mask, + ) + for original in (source, tensor): + restored = pickle.loads(pickle.dumps(original)) + assert restored.refold("word").tolist() == [1.0, 2.0, 3.0] + assert restored.lengths.full_names == ("sample", "word") + assert restored.mask.tolist() == [[True, True], [True, False]] + renamed = source.with_data(source.as_tensor()) + renamed.full_names = ("doc", "token") + assert renamed.lengths.full_names == ("doc", "token") + assert source.full_names == ("sample", "word") diff --git a/tests/test_indices.py b/tests/test_indices.py index 94749f8..9affbad 100644 --- a/tests/test_indices.py +++ b/tests/test_indices.py @@ -1,13 +1,11 @@ import numpy as np +import pytest import torch import foldedtensor as ft def build_tensor(): - # Two samples total. First sample mirrors the example in the prompt - # and totals 14 tokens (contexts: 5 and 9). Second sample has one - # small context to ensure strides are unchanged for the (context, token) view. data = [ [ [ @@ -32,218 +30,125 @@ def build_tensor(): return ft.as_folded_tensor(data, full_names=("sample", "context", "word", "token")) -def build_tensor_single_sample(): - # Single sample as in the prompt example - data = [ - [ - [ - [0, 2, 3], - [10], - [4], - ], - [ - [0, 1, 2], - [2, 3], - [10, 11], - [100, 101], - ], - ], +@pytest.mark.parametrize("data_dims", [(3,), (1, 3), (0, 1, 2, 3)]) +@pytest.mark.parametrize( + "dims,begins,ends,expected,offsets", + [ + (("token",), ([0, 3, 14],), ([3, 14, 17],), list(range(17)), [0, 3, 14]), + (("word",), ([0, 1, 3],), ([1, 3, 9],), list(range(17)), [0, 3, 5]), + (("context",), ([0, 1, 2],), ([1, 2, 3],), list(range(17)), [0, 5, 14]), + (("sample",), ([0, 1],), ([1, 2],), list(range(17)), [0, 14]), + ( + ("context", "word"), + ([0, 0, 1], [0, 1, 2]), + ([0, 1, 1], [1, 3, 4]), + list(range(12)) + list(range(10, 14)), + [0, 3, 12], + ), + ( + ("sample", "word"), + ([0, 1], [2, 1]), + ([0, 1], [4, 2]), + [4, 5, 6, 7, 15, 16], + [0, 4], + ), + ], +) +def test_ranges(data_dims, dims, begins, ends, expected, offsets): + tensor = build_tensor().refold(*data_dims) + indices, actual_offsets, owners = tensor.lengths.make_indices_ranges( + begins=begins, + ends=ends, + indice_dims=dims, + ) + assert indices == tensor.indexer[expected].tolist() + assert actual_offsets == offsets + assert owners == [ + i + for i, (a, b) in enumerate(zip(offsets, offsets[1:] + [len(expected)])) + for _ in range(a, b) ] - return ft.as_folded_tensor(data, full_names=("sample", "context", "word", "token")) -def test_map_indices_flat_unpadded_tokens_by_token(): - t = build_tensor() - - assert t.refold("token").lengths.map_indices( - indices=([0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13],), - indice_dims=("token",), - ) == [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13] - - -def test_map_indices_flat_unpadded_tokens_by_word(): - t = build_tensor() - - assert t.refold("token").lengths.map_indices( - indices=([0, 1, 2, 3, 4, 5, 6, 7],), - indice_dims=("word",), - ) == [0, 3, 4, 5, 8, 10, 12, 14] - - -def test_map_indices_flat_unpadded_tokens_by_context(): - t = build_tensor() - - assert t.refold("token").lengths.map_indices( - indices=([0, 1, 2],), - indice_dims=("context",), - ) == [0, 5, 14] - - -def test_map_indices_flat_unpadded_tokens_by_sample(): - t = build_tensor() - - assert t.refold("token").lengths.map_indices( - indices=([0, 1],), - indice_dims=("sample",), - ) == [0, 14] - - -def test_map_indices_subset_words(): - t = build_tensor() - - assert t.refold("token").lengths.map_indices( - indices=([0, 1, 2, 4, 6],), - indice_dims=("word",), - ) == [0, 3, 4, 8, 12] - - -def test_map_indices_context_word_to_padded_context_token(): - t = build_tensor() - - assert t.refold("context", "token").lengths.map_indices( - indices=([0, 0, 0, 0, 1, 1, 1, 1], [0, 1, 2, 3, 0, 1, 2, 3]), +@pytest.mark.parametrize("array", [torch.as_tensor, np.asarray]) +def test_ranges_broadcast_and_output_type(array): + tensor = build_tensor() + indices, offsets, owners = tensor.lengths.make_indices_ranges( + begins=(array([[0], [1]]), array([0, 1])), + ends=(array([[0], [1]]), array([1, 3])), indice_dims=("context", "word"), - ) == [0, 3, 4, 5, 9, 12, 14, 16] - - -def test_make_indices_ranges_flat_tokens(): - t = build_tensor_single_sample() - - indices, offsets, spans = t.refold("token").lengths.make_indices_ranges( - begins=(torch.as_tensor([0, 0, 1]), torch.as_tensor([0, 1, 2])), - ends=(torch.as_tensor([0, 1, 1]), torch.as_tensor([1, 3, 4])), - indice_dims=("context", "word"), - return_tensors=False, - ) - - assert indices == [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 10, 11, 12, 13] - assert offsets == [0, 3, 12] - assert spans == [0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2] - - -def test_make_indices_ranges_one_dim_token(): - t = build_tensor_single_sample() - indices, offsets, spans = t.refold("token").lengths.make_indices_ranges( - begins=(torch.tensor([0, 3, 12]),), - ends=(torch.tensor([3, 12, 14]),), - indice_dims=("token",), - return_tensors=False, - ) - assert indices == list(range(0, 3)) + list(range(3, 12)) + list(range(12, 14)) - assert offsets == [0, 3, 12] - assert spans == [0] * 3 + [1] * 9 + [2] * 2 - - -def test_make_indices_ranges_one_dim_word(): - t = build_tensor_single_sample() - indices, offsets, spans = t.refold("token").lengths.make_indices_ranges( - begins=(torch.tensor([0, 1, 3]),), - ends=(torch.tensor([1, 3, 7]),), - indice_dims=("word",), - return_tensors=False, - ) - assert indices == list(range(0, 3)) + list(range(3, 5)) + list(range(5, 14)) - assert offsets == [0, 3, 5] - assert spans == [0] * 3 + [1] * 2 + [2] * 9 - - -def test_word_span_mean_pooler_with_embedding_bag_flat_indices(): - t = build_tensor().refold("context", "token") - # 0 -> 2: [[0, 2, 3], [10]] - # 5 -> 7: [[10, 11], [100, 101]] - indices, offsets, spans = t.lengths.make_indices_ranges( - begins=(torch.tensor([0, 5]),), - ends=(torch.tensor([2, 7]),), - indice_dims=("word",), - ) - embeds = t.unsqueeze(-1).expand(-1, -1, 2).float() - res = torch.nn.functional.embedding_bag( - input=indices, - weight=embeds.view(-1, 2), - offsets=offsets, - mode="mean", - ) - assert res.tolist() == [[3.75, 3.75], [55.5, 55.5]] - - -def test_word_span_mean_pooler_with_embedding_bag_multidim_indices(): - t = build_tensor().refold("context", "token") - # 0 -> 2: [[0, 2, 3], [10]] - # 5 -> 7: [[10, 11], [100, 101]] - indices, offsets, spans = t.lengths.make_indices_ranges( - begins=( - torch.tensor([0, 1]), - torch.tensor([0, 2]), - ), - ends=( - torch.tensor([0, 1]), - torch.tensor([2, 4]), - ), - indice_dims=( - "context", - "word", - ), ) - embeds = t.unsqueeze(-1).expand(-1, -1, 2).float() - res = torch.nn.functional.embedding_bag( - input=indices, - weight=embeds.view(-1, 2), - offsets=offsets, - mode="mean", + assert all( + type(x) is type(array([])) for x in (indices, offsets, owners) # noqa: E721 ) - assert res.tolist() == [[3.75, 3.75], [55.5, 55.5]] - - -def test_map_indices_format_torch_multidimensional(): - t = build_tensor() - - assert torch.allclose( - t.lengths.map_indices( - indices=( - torch.as_tensor([[0, 1, 2, 3, 4, 5, 6], [7, 8, 9, 10, 11, 12, 13]]), - ), - indice_dims=("token",), - ), - torch.tensor([[0, 1, 2, 3, 6, 12, 13], [14, 15, 16, 18, 19, 21, 22]]), + assert offsets.tolist() == [[0, 3], [5, 8]] + assert ( + indices.tolist() + == tensor.indexer[[0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11]].tolist() ) -def test_map_indices_format_numpy_multidimensional(): - t = build_tensor() - - assert np.allclose( - t.lengths.map_indices( - indices=(np.asarray([[0, 1, 2, 3, 4, 5, 6], [7, 8, 9, 10, 11, 12, 13]]),), - indice_dims=("token",), - ), - np.asarray([[0, 1, 2, 3, 6, 12, 13], [14, 15, 16, 18, 19, 21, 22]]), +@pytest.mark.parametrize("dims", [(2,), (0, 2), (0, 1, 2)]) +def test_ranges_and_refolding_with_empty_contexts(dims): + # Empty parents keep their rows when the pooler gathers padded wordpieces + layout = ft.FoldedTensorLayout( + [[3], [0, 4, 2], [2, 0, 1, 3, 1, 2]], + data_dims=(2,), + full_names=("context", "word", "piece"), ) - - -def test_make_indices_ranges_format_torch_multidimensional(): - t = build_tensor_single_sample() - - indices, offsets, spans = t.lengths.make_indices_ranges( - begins=(torch.as_tensor([[0, 0, 1]]), torch.as_tensor([[0, 1, 2]])), - ends=(torch.as_tensor([[0, 1, 1]]), torch.as_tensor([[1, 3, 4]])), - indice_dims=("context", "word"), + values = torch.arange(9, dtype=torch.float32, requires_grad=True) + tensor = ft.as_folded_tensor(values, lengths=layout).refold(*dims) + nested = ft.as_folded_tensor( + [[], [[0.0, 1.0], [], [2.0], [3.0, 4.0, 5.0]], [[6.0], [7.0, 8.0]]], + data_dims=dims, ) - - assert isinstance(indices, torch.Tensor) - assert isinstance(offsets, torch.Tensor) - assert isinstance(spans, torch.Tensor) - assert offsets.shape == (1, 3) - - -def test_make_indices_ranges_format_numpy_multidimensional(): - t = build_tensor_single_sample() - - indices, offsets, spans = t.lengths.make_indices_ranges( - begins=(np.asarray([[0, 0, 1]]), np.asarray([[0, 1, 2]])), - ends=(np.asarray([[0, 1, 1]]), np.asarray([[1, 3, 4]])), + assert torch.equal(tensor.as_tensor(), nested.as_tensor()) + indices, offsets, owners = tensor.lengths.make_indices_ranges( + begins=([0, 1, 1, 1, 1, 2], [0, 0, 1, 1, 2, 0]), + ends=([0, 1, 1, 1, 1, 2], [0, 4, 1, 4, 3, 2]), indice_dims=("context", "word"), ) - assert isinstance(indices, np.ndarray) - assert isinstance(offsets, np.ndarray) - assert isinstance(spans, np.ndarray) - assert offsets.shape == (1, 3) + expected = list(range(6)) + list(range(2, 6)) + [2] + list(range(6, 9)) + actual = tensor.as_tensor().reshape(-1)[indices] + assert actual.tolist() == expected + assert offsets == [0, 0, 6, 6, 10, 11] + assert owners == [1] * 6 + [3] * 4 + [4] + [5] * 3 + actual.sum().backward() + assert values.grad.tolist() == np.bincount(expected, minlength=9).tolist() + + +def test_ranges_empty_and_scalar(): + tensor = ft.as_folded_tensor([[], []], full_names=("sample", "word")) + assert tensor.refold("word").refold("sample", "word").shape == (2, 0) + assert tensor.lengths.make_indices_ranges( + begins=([0, 1], [0]), + ends=([0, 1], [0]), + indice_dims=("sample", "word"), + ) == ([], [0, 0], []) + assert tensor.lengths.make_indices_ranges( + begins=([],), + ends=([],), + indice_dims=("word",), + ) == ([], [], []) + assert build_tensor().lengths.make_indices_ranges( + begins=(0,), + ends=(1,), + indice_dims=("token",), + ) == ([0], 0, [0]) + + +@pytest.mark.parametrize( + "begins,ends,dims,error", + [ + (([-1],), ([1],), ("word",), IndexError), + (([0], [4]), ([0], [4]), ("context", "word"), IndexError), + (([0], [0]), ([3], [0]), ("context", "word"), IndexError), + (([2],), ([1],), ("word",), ValueError), + (([0], [0]), ([1], [1]), ("word", "word"), ValueError), + ], +) +def test_invalid_ranges(begins, ends, dims, error): + with pytest.raises(error): + build_tensor().lengths.make_indices_ranges( + begins=begins, ends=ends, indice_dims=dims + )