diff --git a/README.md b/README.md index d6bf614..185ea92 100644 --- a/README.md +++ b/README.md @@ -106,6 +106,36 @@ print(refolded_embedding.shape) # torch.Size([2, 5, 16]) # 2 samples, 5 words max, 16 dims ``` +### Pooling spans + +`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 + +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",), +) +pooled = torch.nn.functional.embedding_bag( + indices, + tensor.as_tensor().reshape(-1, 1), + offsets, + mode="mean", +) +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 5450eaa..69642b5 100644 --- a/changelog.md +++ b/changelog.md @@ -1,5 +1,12 @@ # Changelog +## Unreleased + +- 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 - Fix `storage` torch warning diff --git a/foldedtensor/__init__.py b/foldedtensor/__init__.py index 5499ab8..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 @@ -8,7 +7,99 @@ import torch from torch.autograd import Function -from . import _C +from . import _C # type: ignore[import] + + +def make_indices_ranges( + *, begins, ends, indice_dims, lengths, data_dims, return_tensors=None +): + """ + Expand half open ranges into storage indices for span pooling, excluding padding + + Parameters + ---------- + 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 + ------- + 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,), + ) + # indices: [0, 1, 2, 2, 3] + # ------- ---- + # offsets: [0, 3, 5] + # span_indices: [0, 0, 0, 1, 1] + ``` + """ + 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, + data_dims, + ) + 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 + ) + 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 = { torch.bool: bool, @@ -49,13 +140,100 @@ __version__ = "0.4.0" -class FoldedTensorLengths(UserList): +class FoldedTensorLayout(UserList): + """ + 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=(), *, 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))) + ) + + 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) + 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=None, return_tensors=None + ): + """ + 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=self.resolve_dims(indice_dims), + lengths=self, + data_dims=self.data_dims + if data_dims is None + else self.resolve_dims(data_dims), + return_tensors=return_tensors, + ) -if typing.TYPE_CHECKING: - FoldedTensorLengths = List[List[int]] # noqa: F811 + +FoldedTensorLengths = FoldedTensorLayout # noinspection PyMethodOverriding @@ -88,11 +266,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 +325,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 +348,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 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( @@ -189,11 +371,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 +398,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 +426,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, ) @@ -266,33 +444,48 @@ 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: FoldedTensorLengths, - data_dims: Sequence[int], - full_names: Sequence[str], - 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, ): - data_dims = data_dims - full_names = full_names + 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.data_dims = data_dims - instance.full_names = full_names instance.indexer = indexer instance._mask = mask return instance @@ -301,12 +494,30 @@ 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 + + @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: @@ -329,8 +540,6 @@ def to(self, *args, **kwargs): return FoldedTensor( data=result, 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 ), @@ -359,9 +568,10 @@ 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 structures" + 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: @@ -414,6 +624,11 @@ def refold(self, *dims: Union[Sequence[Union[int, str]], int, str]): f"could not be refolded with dimensions {list(dims)}" ) + if not dims or dims[-1] != len(self.lengths) - 1: + raise ValueError( + "The last dimension of data_dims must be the last variable dimension" + ) + if dims == self.data_dims: return self diff --git a/foldedtensor/functions.cpp b/foldedtensor/functions.cpp index 718c69b..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" @@ -310,6 +392,7 @@ 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"); } diff --git a/tests/test_folded_tensor.py b/tests/test_folded_tensor.py index a4d2e24..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() @@ -446,3 +453,63 @@ def test_missing_dims(): tensor.refold("line", "token") assert "line" in str(e.value) + + +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] + recreated = as_folded_tensor(tensor.as_tensor(), lengths=tensor.lengths) + renamed = as_folded_tensor( + tensor.as_tensor(), + lengths=tensor.lengths, + full_names=("sample_bis", "token_bis"), + ) + 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, + ) + 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 new file mode 100644 index 0000000..9affbad --- /dev/null +++ b/tests/test_indices.py @@ -0,0 +1,154 @@ +import numpy as np +import pytest +import torch + +import foldedtensor as ft + + +def build_tensor(): + 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")) + + +@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) + ] + + +@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"), + ) + assert all( + type(x) is type(array([])) for x in (indices, offsets, owners) # noqa: E721 + ) + 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() + ) + + +@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"), + ) + 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 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"), + ) + 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 + )