NewYour coding agent can read the release notes before it upgrades.Set up the MCP server →
PyPI · #3332 most downloaded on PyPI
TensorDict is a pytorch dedicated tensor container.
Last release 25 days ago
09 Sep 2026
Release timing varies
gaps range from 8 days to 5 months
Nearly every release is documented
notes for 34 of 34 stable releases
Nothing withdrawn
no release was ever pulled
4 years old
40 releases · first in 2022
One column per quarter.
…and CUDA graph replay. No new public APIs, breaking changes, or deprecations are introduced.
TensorDict 0.14.2 is a patch release correcting compiled shallow copies, non-tensor concatenation, and CUDA graph replay. No new public APIs, breaking changes, or deprecations are introduced.
Install after publication: pip install tensordict==0.14.2
Thanks to @vmoens.
…tensorclass deserialization. It contains no breaking changes, new features, or deprecations.
TensorDict 0.14.1 is a focused patch release with correctness fixes for module metadata, lazy and non-tensor indexing, and frozen tensorclass deserialization. It contains no breaking changes, new features, or deprecations.
TensorDictModuleBase subclasses, including their child modules and input and output keys (#1767).LazyStackedTensorDict (#1772).TensorDictModule output-key metadata synchronized when out_keys is reassigned (#1773).NonTensorStack return independent entries, preventing indexed results from mutating their source or one another (#1778).pip install tensordict==0.14.1Thanks to @priba, @gtnv, @simonschlaepfer, and @vmoens for their contributions to this release.
Full Changelog: v0.14.0...v0.14.1
TensorDict 0.14.0 introduces Zarr-backed persistent storage, portable single-file memory-mapped archives, and direct autograd through TensorDict.backw
TensorDict 0.14.0 introduces Zarr-backed persistent storage, portable single-file memory-mapped archives, and direct autograd through TensorDict.backward(). It also improves serialization performance, distributed point-to-point operations, UnbatchedTensor interoperability, and memmap correctness and security.
TensorDict.to_module() now preserves existing module parameter and buffer registrations by default. Pass preserve_module_state=False to retain the previous replacement behavior (#1720).TensorDict.copy_at_() now uses its optimized tensor-only path by default. Pass fast=False to retain the previous update_at_() fallback behavior (#1691).allow_pickle across load and refresh APIs. In 0.14, omitting it still loads pickled non-tensor data with a warning; the default will become allow_pickle=False in 0.15. Pass True only for trusted artifacts or False to reject pickle now (#1763).PersistentTensorDict storage through to_zarr(), from_zarr(), and compatible store backends (#1754)..tdz archives for memory-mapped TensorDicts, including lazy loading, subpath access, compression, writable mappings, and pack/unpack utilities (#1735).TensorDict.backward() for differentiating all compatible tensor leaves in one autograd operation (#1733).group_dst and group_src for distributed send and receive operations (#1739).vmap support and Python scalar conversions for UnbatchedTensor (#1729, #1728).NonTensorStack entries in contiguous() (#1746).TensorDictParams.load_state_dict(strict=False) to skip missing submodules (#1734)._SubTensorDict memmap round trips (#1736).UnbatchedTensor stack metadata (#1731, #1730).Thanks to @peterdsharpe, and @javierdejesusda, both first-time contributors!
Full Changelog: https://github.com/pytorch/tensordict/compare/v0.13.0...v0.14.0
TensorDict 0.13.0 adds new data import/export capabilities, broader in-place shape-changing operations, and several torch.compile reliability improvem
TensorDict 0.13.0 adds new data import/export capabilities, broader in-place shape-changing operations, and several torch.compile reliability improvements. This release also improves memmap/file-backed workflows, TensorClass behavior under compile/export, to_module() state preservation migration, and release infrastructure for CPU-only TensorDict wheels.
robust_key now defaults to robust filesystem-safe key encoding for memmap save/load APIs (#1717). Passing robust_key=False keeps the legacy filename behavior for existing legacy memmap directories.robust_key FutureWarning targeting v0.12 has been enacted by making robust encoding the default (#1717).TensorDict.to_module() now warns when writing a plain tensor over an existing nn.Parameter would deregister it from module.state_dict(); pass preserve_module_state=False to keep the v0.13 behavior or preserve_module_state=True to opt in to the v0.14 default (#1720).copy_at_(..., fast=None) fallback warning remains active; the default is scheduled to change to fast=True in v0.14.**kwargs support to update() and update_() (#1677).ProbabilisticTensorDictModule (#1689).canonical=True support to contiguous() (#1698).inplace=True support to pad() to reduce peak memory (#1706).inplace=True support to gather, repeat, repeat_interleave, roll, reshape, flatten, unflatten, and contiguous (#1707).preserve_module_state to TensorDict.to_module() so existing nn.Parameter and buffer registrations can be preserved when loading TensorDict leaves into modules (#1720).unravel_keys inconsistency that prevented torch.compile in some cases (#1674).ragged_idx handling for consolidated tensors (#1675).MetaData (#1697).SymInt values in _parse_batch_size and storing batch dimensions in pytree context (#1704).__post_init__ and field default application under torch.compile (#1709).TensorDict.copy_at_ (#1691).frozenset and dict_keys (#1703).clone(recurse=False) via cached schema.Thanks to all contributors to this release, including first-time contributors @mathieuorhan, @arpitjain099, and @felixmaximilian.
Full Changelog: https://github.com/pytorch/tensordict/compare/v0.12.0...v0.13.0
TensorDict 0.12.4 is a bugfix-only release. It contains no new features.
TensorDict 0.12.4 is a bugfix-only release. It contains no new features.
SymInt batch sizes and preserved batch-dimension metadata in pytree context (#1704).__post_init__, field defaults, default factories, and omitted None defaults are applied correctly under torch.compile (#1709).non_blocking handling (#1711).frozenset - dict_keys now that it is traceable in newer PyTorch versions.[BugFix] Fix auto batch size with MetaData by @vmoens in https://github.com/pytorch/tensordict/pull/1697
TensorDictBase.auto_batch_size_() no longer collapses the batch size to an
empty shape when a TensorDict contains scalar MetaData / NonTensorData
(pass-through) entries. Previously a single scalar metadata leaf forced the
inferred batch size down to torch.Size([]), ignoring real tensor leaves.
Full Changelog: https://github.com/pytorch/tensordict/compare/v0.12.2...v0.12.3
TensorDictBase.auto_batch_size_() no longer collapses the batch size to an
empty shape when a TensorDict contains scalar MetaData / NonTensorData
(pass-through) entries. Previously a single scalar metadata leaf forced the
inferred batch size down to torch.Size([]), ignoring real tensor leaves.
Full Changelog: v0.12.2...v0.12.3
Patch release with a bug fix for consolidated nested tensors.
Patch release with a bug fix for consolidated nested tensors.
_ragged_idx loss during consolidation of nested tensors, which caused numerical incorrectness when the nested tensor had more than 2 dimensions and ragged_idx != 1 (#1675)pip install tensordict==0.12.2
Full Changelog: https://github.com/pytorch/tensordict/compare/v0.12.1...v0.12.2
Patch release with a torch.compile bug fix.
Patch release with a torch.compile bug fix.
unravel_keys inconsistency that prevented torch.compile from working correctly when called with a single key (#1674)pip install tensordict==0.12.1
TensorDict v0.12.0 introduces TypedTensorDict for schema-enforced tensor dictionaries, a full distributed collectives suite (broadcast, all_reduce, al
TensorDict v0.12.0 introduces TypedTensorDict for schema-enforced tensor dictionaries, a full distributed collectives suite (broadcast, all_reduce, all_gather, scatter), TensorDictStore with Redis/Dragonfly/KeyDB backends, and major torch.compile and performance improvements. The UnbatchedTensor has been rewritten as a proper tensor subclass, and state_dict handling has been overhauled for consistency.
UnbatchedTensor is now a __torch_dispatch__-based tensor subclass (was previously a wrapper) (#1638, #1648)state_dict is now flat by default, with auto-detection in load_state_dict for backwards compatibilitystate_dict now uses logical keystorch.compile support (#1657, #1659, #1660, #1662, #1663)broadcast, all_reduce, all_gather, scatter, consolidated send/recv and init_remote/from_remote_init with UCXX transport support (#1611)set_printoptions: Configurable TensorDict repr with verbose mode (#1654, #1655, #1665)torch.func support: jacrev, jacfwd, and hessian now work with TensorDict (#1613)vmap with unbatched data: TensorDicts containing unbatched tensors can now be vmapped (#1625)TensorClass.select(as_tensordict=...) parameter (#1544)TensorDictBase.is_non_tensor(key) for consistent non-tensor key detectionHigherOrderOperator support in __torch_function__ (#1668)td[key] = [] handling (#1666)UnbatchedTensor.tolist() (#1664)UnbatchedTensor CUDA pickling for multiprocessing (#1656)UnbatchedTensor indexing without batch dim, GPU failures, getitem/stack (#1607, #1626, #1633)replace() recompiles under torch.compile (#1605)auto_batch_size regression with NonTensorStack (#1609)NonTensorData positional args causing graph breaks under torch.compile (#1630)state_dict error messages, params forwarding, detachpybind11>=2.13 for Python 3.13 compatibilityTensorDict.__init__ by bypassing set() dispatch chain (#1551)TensorClass.__init__ by bypassing __setattr__ and dataclass __init__ (#1552)torch.cuda.synchronize() in .to() (#1545)torch.compile (#1636)Removed deprecated features scheduled for v0.11 (#1526):
This release brings Python 3.14 support (including free-threaded mode), drops Python 3.9, adds extensive torch.compile compatibility improvements, introduces new tensor manipulation methods, and expands NPU device support.
pad_sequence device parameter removed (deprecated since v0.5)MemoryMappedTensor._tensor property now raises RuntimeError - use the tensor directly since it's a tensor subclassset_lazy_legacy now raises RuntimeErrorlock, unlock, rename_key (without trailing underscore) now raise RuntimeError - use lock_, unlock_, rename_key_ insteadPython 3.14 and free-threading support (#1481, #1438) - Fixed race conditions in value caching that caused segfaults under free-threaded Python (PYTHON_GIL=0). The library now declares it does not require the GIL.
TensorClassModuleBase (#1473) - New base class for type-safe PyTorch modules that operate on TensorClass instances. Provides compile-time type checking, automatic TensorDict conversion via as_td_module(), nested TensorClass support, and ONNX export compatibility.
New tensor manipulation methods - Added several torch-compatible operations:
td.movedim() (#1504)td.swapaxes()td.flip(), td.fliplr()td.roll()td.rot90()td.narrow(), td.tile(), td.broadcast_to()td.atleast_Nd()td.mod()Quantiles support (#1450) - Added quantile computation for TensorDicts
Context manager for td.to() (#1448) - td.to() can now be used as a context manager for temporary device transfers
key_transform for reduce ops (#1451) - Added key transformation support for reduction operations
selected_out_keys for ProbabilisticTensorDictSequential (#1497) - Added ability to specify which output keys to keep
Extended NPU support (#1471, #1460) - Expanded NPU device support across more use cases previously limited to CUDA
torch.compile compatibility fixes:
Robust key setting with memmap tensordicts (#1443) - Major fix for handling key operations with memory-mapped tensordicts
Saving and loading metadata (#1444) - Fixed memmap + serialization issues with MetaData typed objects
tensordict.cat preserves TensorClass (#1456)
Names propagation for nested TensorDicts (#1520) - Fixed names propagation with extra batch dimensions
Device mismatch in CUDA/NPU scenarios (#1439) - Fixed "expect all tensors on the same device" errors
Nested lazy stacks handling (#1461) - Fixed as_list and other operations with nested lazy stacks
PyTorch version compatibility (#1498, #1499, #1437) - Various fixes for compatibility with older PyTorch versions
[Deprecation] Upgrade set_list_to_stack behavior (#1382): Updated behavior of set_list_to_stack with proper deprecation warnings for the old API
We are excited to announce the release of TensorDict 0.10.0! This release includes significant improvements to type annotations, new features for metadata handling, enhanced tensor operations, and numerous bug fixes that improve the overall stability and usability of the library.
to_struct_array now supports string data types (#1410)to_mds method creates an MDS dataset on your favourite location -- no more painful columns definition etc (#1426).grad function (#1417)to_struct_array functionality to handle string data typestensor_split operation to match PyTorch tensor API__all__ declarations across all modules for better IDE support and import control@overload for methods that have a reduce arg (#1427): Added proper type overloads for methods with reduce parameters__getitem__ (#1402): Improved type annotations for indexing operationsTensorDictModuleWrapper forward (#1415): Fixed forward pass in TensorDictModuleWrapperrepeat_interleave to properly support tensor argumentsset_list_to_stack with proper deprecation warnings for the old APISpecial thanks to all the contributors who made this release possible:
For comprehensive documentation, tutorials, and examples, visit:
For a complete list of changes, see the full changelog.
If you encounter any issues, please report them on our GitHub Issues page.
This minor releases brings the following improvements and bug fixes:
This minor releases brings the following improvements and bug fixes:
Full Changelog: https://github.com/pytorch/tensordict/compare/v0.9.0...v0.9.1
`NormalParamWrapper`: Deprecated in favor of tensordict.nn.NormalParamExtractor
TensorDict 0.9.0 introduces significant improvements in performance, new features for lazy operations, enhanced CUDA graph support, and various bug fixes. This release focuses on stability improvements and new functionality for distributed and lazy tensor operations.
to_lazystack(): New method to convert TensorDict instances to lazy stacks (#1351) (a5aab970)tensordict.stack now preserves names when stacking TensorDict instances (#1348) (2053031b)update_batch_size in where(): Enhanced where() operation now supports update_batch_size parameter (#1365) (847a86c4)tolist_first(): New method for converting TensorDict to list with first-level flattening (#1334) (73fe89bc)torch.maximum support: Added support for torch.maximum operation in TensorDict (#1362) (85f26e41)torch.sum, torch.mean, torch.var and loss functions (l1, smooth_l1, mse) (#1361) (17ca2ffa)CudaGraphModule.state_dict(): New method to access state dictionary of CUDA graph modules (#1346) (909907b1)NonTensorDataBase and MetaData: New base classes for handling non-tensor data in TensorDict (#1324) (8d0241d4)TensorDict.__copy__(): New method for creating shallow copies of TensorDict instances (#1321) (b6feadd1)broadcast tensordicts: New functionality for broadcasting TensorDict instances across different shapes (#1307) (29598633)remote_init with subclasses: Enhanced remote initialization support for TensorDict subclasses (#1308) (5859a2cf)return_early for isend: New parameter for early return in send operations (#1306) (40127677)split_size validation in TensorDict.split() (#1370) (0bb94c08)update_batch_size when source is TD and destination is LTD (#1371) (c8bfda2c)new_ operations on NonTensorStack (#1366) (5e67c32)tensor_only construction (#1364) (75e2c26d)update_batch_size in lazy stack updates (#1359) (c9f0e406)_non_tensordict (#1353) (c11a95b8)tensorclass __enter__ and __exit__ methods (#1352) (08abb06c)tensordict.stack forcing all names to None when no match (#1350) (c90df00a)maybe_dense_stacks (#1340) (9c8dd2db)new_* operations for Lazy stacks (#1317) (1c8be198)__setitem__ (#1313) (1d642b04)CudaGraphModule on devices that are not 0 (#1315) (89d05a1c)isend early return (#1316) (8d3c4706)_validate_value (#1310) (a36f7f92)return_composite defaults to True only when >1 distribution (#1328) (2c739240)TDParams compatibility with export (#1285) (ecdde0b1)_is_list_tensor_compatible missing return value (#1277) (a9cc6321).item() warning on tensors that require grad (#1283) (910c9533)_get_item: Optimized item retrieval operations (#1288) (1e33a18f)tensor_only for tensorclass: Enhanced performance for tensor-only operations (#1280) (d4bc34cd)_C extension now statically linked against Python library (#1304) (af17524b)NormalParamWrapper: Deprecated in favor of tensordict.nn.NormalParamExtractoris_functional, make_functional, and get_functional have been removed from tensordictFutureWarning will be raised if lists are assigned to TensorDict without setting the appropriate context manager.is_functional, make_functional, or get_functional, these have been removedtensordict.nn.NormalParamExtractor instead of NormalParamWrapperset_list_to_stack context manager for consistent list behaviorSpecial thanks to all contributors who made this release possible, including:
For a complete list of all changes, please refer to the git log from version 0.8.0 to 0.9.0.
This minor release provides some fixes to CudaGraphModule, allowing the module to run on different devices than the default.
This minor release provides some fixes to CudaGraphModule, allowing the module to run on different devices than the default.
It also adds __copy__ to the TensorDict ops, such that copy(td) triggers td.copy(), resulting in a copy of the TD stucture without new memory allocation.
Full Changelog: https://github.com/pytorch/tensordict/compare/v0.8.2...v0.8.3
This release fixes an apparent memory leak due to the value validation in tensordict. The leak is apparent, as in it disappears in gc.collect() is inv
This release fixes an apparent memory leak due to the value validation in tensordict.
The leak is apparent, as in it disappears in gc.collect() is invoked.
See #1309 for context.
Full Changelog: https://github.com/pytorch/tensordict/compare/v0.8.1...v0.8.2
This new minor fixes the _C build pipeline, which was failing on some machines as the extension was build with dynamic linkage against libpython
This new minor fixes the _C build pipeline, which was failing on some machines as the extension was build with dynamic linkage against libpython
[Deprecation] Deprecate KJT support 63730afb @vmoens
We're excited to announce a new tensordict release, packed with new features, packaging perks as well as bug fixes.
The interaction with non-tensor data is now much easier to get by ():
set_list_to_stack(True).set() # Ask for new behaviour
td = TensorDict(batch_size=(3, 2))
td["numbers"] = [["0", "1"], ["2", "3"], ["4", "5"]]
print(td)
# TensorDict(
# fields={
# numbers: NonTensorStack(
# [['0', '1'], ['2', '3'], ['4', '5']],
# batch_size=torch.Size([3, 2]),
# device=None)},
# batch_size=torch.Size([3, 2]),
# device=None,
# is_shared=False)
Stacks of non-tensor data can also be reshaped 58ccbf52. Using the previous example:
td = td.view(-1)
td["numbers"]
# ['0', '1', '2', '3', '4', '5']
We also made it easier to get values of lazy stacks (f7bc8394):
tds = [TensorDict(a=torch.zeros(3)), TensorDict(a=torch.ones(2))]
td = lazy_stack(tds)
print(td.get("a", as_list=True))
# [tensor([0., 0., 0.]), tensor([1., 1.])]
print(td.get("a", as_nested_tensor=True))
# NestedTensor(size=(2, j1), offsets=tensor([0, 3, 5]), contiguous=True)
print(td.get("a", as_padded_tensor=True, padding_value=-1))
# tensor([[ 0., 0., 0.],
# [ 1., 1., -1.]])
You can now install tensordict with any PyTorch version. We only provide test coverage for the latest pytorch (currently 2.7.0), so for any other version you will be on your own in terms of compatibility but there should be no limitations in term of installing the library with older version of pytorch.
strict kwarg in TDModule 06215b60 @vmoenszip(..., strict=True) in TDModules 0a8638d3 @vmoensnone value for _CAPTURE_NONTENSOR_STACK 2f60f8d9 @vmoens_get_item 1e33a18f @vmoensAddStateIndependentNormalScale (#1216) 309c0a3c Faury Louis*Full Changelog: https://github.com/pytorch/tensordict/compare/v0.7.0...v0.8.0
We are pleased to announce the release of tensordict v0.7.2, which includes several bug fixes and backend improvements.
We are pleased to announce the release of tensordict v0.7.2, which includes several bug fixes and backend improvements.
Improved errors for TensorDictSequential (#1227)
Improved documentation for TensorDictModuleBase (#1226)
Full Changelog: https://github.com/pytorch/tensordict/compare/v0.7.1...v0.7.2
We are pleased to announce the release of tensordict v0.7.1, which includes several bug fixes and a deprecation notice.
We are pleased to announce the release of tensordict v0.7.1, which includes several bug fixes and a deprecation notice.
Softly deprecated extra-tensors with respect to out_keys (#1215). We make sure a warning is raised when the number of output tensors and output keys do not match.
Full Changelog: https://github.com/pytorch/tensordict/compare/v0.7.0...v0.7.1
…[Tests] Skip deprecation warning tests on FB fbcode (#1128) (22da6791) by @vmoens ghstack-source-id: fb0cc381a670377667194324f5b019076b8e762d [Version…
v0.7.0 brings a lot of new features and bug fixes. Thanks to the vibrant community to help us keeping this project alive!
A special thanks to our new contributors (+ people interacting with us on the PyTorch forum, discord or via issues and social platforms)!
min and max operations as we do with torch.Tensor.min. Previously, tensorclasses
were used, but that lead to some undefined behaviors when indexing (in PyTorch, min returns a namedtuple that can be
indexed to get the values or argmax, whereas indexing a tensorclass indexes it along the batch dimension).== sign are exactly equivalent:td = TensorDict(..., batch_size=[3, 5])
t = torch.randn(5)
td + t == td + t.expand(td.shape)
TL;DR: We're changing the way log-probs and entropies are collected and written in ProbabilisticTensorDictModule and
in CompositeDistribution. The "sample_log_prob" default key will soon be "<value>_log_prob (or
("path", "to", "<value>_log_prob") for nested keys). For CompositeDistribution, a different log-prob will be
written for each leaf tensor in the distribution. This new behavior is controlled by the
tensordict.nn.set_composite_lp_aggregate(mode: bool) function or by the COMPOSITE_LP_AGGREGATE environment variable.
We strongly encourage users to adopt the new behavior by setting tensordict.nn.set_composite_lp_aggregate(False).set()
at the beginning of their training script.
We've had multiple rounds of refactoring for CompositeDistribution which relied on some very specific assumptions and
resulted in a brittle and painful API. We now settled on the following API that will be enforced in v0.9, unless the
tensordict.nn.set_composite_lp_aggregate(mode) value is explicitly set to True (the current default).
The bulk of the problem was that log-probs were aggregated in a single tensor and registered in td["sample_log_prob"].
This had the following problems:
"sample_log_prob" is a generic but inappropriate name (the data may not be a random sample but anything else.)TensorClass class to do simple inheritance-style coding, which is accompanied by a stub file
that encodes all the TensorDict op signatures (we ensure this in the CI). See #1067@tensorclass(shadow=True) or class T(TensorClass["shadow"]): ...
and you will be able to use dedicated names like get or values as attribute names. This is slightly unsafe when you
nest the tensorclass, as we can't guarantee that the container won't be calling these methods directly on the
tensorclass.@tensorclass(nocast=True) and TensorClass["nocast"] will deactivate the auto-casting in tensorclasses.
The behavior is now:
int or np.arrays will be cast to torch.Tensor
instances for example).autocast=True will cause @tensorclass to go one step further and attempt to cast values to the type indicated
in the dataclass definition.nocast=True keeps values as they are. All non-tensor (or non-tensordict/tensorclass) values will be wrapped in
a NonTensorData.NonTensorStack(*values).See the full list of features here:
[Feature] Add __abs__ docstrings, __neg__, __rxor__, __ror__, __invert__, __and__, __rand__, __radd__, __rtruediv__, __rmul__, __rsub__, __rpow__, bitwise_and, logical_and (#1154) (d1363eb9) by @vmoens ghstack-source-id: 97ce710b5a4b552d9477182e1836cf3777c2d756
[Feature] Add expln map to NormalParamExtractor (#1204) (e900b24c) by @vmoens ghstack-source-id: 9003ceafbe8ecb73c701ea1ce96c0a342d0679b0
[Feature] Add missing __torch_function__ (#1169) (bc6390c3) by @vmoens ghstack-source-id: 3dbefb4f5322a944664bbc2d29af7f862cb92342
[Feature] Better list casting in TensorDict.from_any (#1108) (1ffc4637) by @vmoens ghstack-source-id: 427d19d5ef7c0d2779e064e64522fc0094a885af
[Feature] Better logs of key errors in assert_close (#1082) (747c593b) by @vmoens ghstack-source-id: 46cb41d0da34b17ccc248119c43ddba586d29d80
[Feature] COMPOSITE_LP_AGGREGATE env variable (#1190) (9733d6e6) by @vmoens ghstack-source-id: 16b07d0eac582cfd419612f87e38e1a7acffcfc0
[Feature] CompositeDistribution.from_distributions (#1113) (a45c7e39) by @vmoens ghstack-source-id: 04a62439b0fe60422fbc901172df46306e161cc5
[Feature] Ensure all dists work with DETERMINSTIC type without warning (#1182) (8e63112f) by @vmoens ghstack-source-id: 63117f9b3ac4125a2be4e3e55719cc718051fc10
[Feature] Expose WrapModule (#1118) (d849756f) by @vmoens ghstack-source-id: 55caa5d7c39e0f98c1e0558af2a076fee15f7984
[Feature] Fix type assertion in Seq build (#1143) (eaafc18c) by @vmoens ghstack-source-id: 83d3dcafe45568c366207395a22b22fb35f61de1
[Feature] Force log_prob to return a tensordict when kwargs are passed to ProbabilisticTensorDictSequential.log_prob (#1146) (98c57eef) by @vmoens ghstack-source-id: 326d0763c9bbb13b51daac91edca4f0e821adf62
[Feature] Make ProbabilisticTensorDictSequential account for more than one distribution (#1114) (c7bd20c7) by @vmoens ghstack-source-id: b62b81b5cfd49168b5875f7ba9b4f35b51cd2423
[Feature] NonTensorData(*sequence_of_any) (#1160) (70d4ed17) by @vmoens ghstack-source-id: 537f3d87b0677a1ae4992ca581a585420a10a284
[Feature] NonTensorStack.data (#1132) (4404abe0) by @vmoens ghstack-source-id: 86065377cc1cd7c7283ed0a468f5d5602d60526d
[Feature] NonTensorStack.from_list (#1107) (f924afcd) by @vmoens ghstack-source-id: e8f349cb06a72dcb69a639420b14406c9c08aa99
[Feature] Optional in_keys for WrapModule (#1145) (2d37d92b) by @vmoens ghstack-source-id: a18dd5dff39937b027243fcebc6ef449b547e0b0
[Feature] OrderedDict for TensorDictSequential (#1142) (7df20626) by @vmoens ghstack-source-id: a8aed1eaefe066dafaa974f5b96190860de2f8f1
[Feature] ProbabilisticTensorDictModule.num_samples (#1117) (978d96cc) by @vmoens ghstack-source-id: dc6b1c98cee5fefc891f0d65b66f0d17d10174ba
[Feature] ProbabilisticTensorDictSequential.default_interaction_type (#1123) (68ce9c3b) by @vmoens ghstack-source-id: 37d38df36263e8accd84d6cb895269d50354e537
[Feature] Subclass conservation in td ops (#1186) (070ca618) by @vmoens ghstack-source-id: 83e79abda6a4bb6839d99240052323380981855c
[Feature] TensorClass (#1067) (a6a0dd61) by @vmoens ghstack-source-id: c3d4e17599a3204d4ad06bceb45e4fdcd0fd1be5
[Feature] TensorClass shadow attributes (#1159) (c744bcfa) by @vmoens ghstack-source-id: b5cc7c7fea2d48394e63d289ee2d6f215c2333bc
[Feature] TensorDict.<reduction>(dim='feature') (#1121) (ba431598) by @vmoens ghstack-source-id: 68f21aca722895e8a240dbca66e97310c20a6b5d
[Feature] TensorDict.clamp (#1165) (646683c6) by @vmoens ghstack-source-id: 44f0937c195d969055de10709402af7c4473df32
[Feature] TensorDict.logsumexp (#1162) (e564b3a6) by @vmoens ghstack-source-id: 84148ad9c701029db6d02dfb84ddb0a9b26c9ab7
[Feature] TensorDict.separates (#1120) (674f356d) by @vmoens ghstack-source-id: be142a150bf4378a0806347257c3cf64c78e4eda
[Feature] TensorDict.softmax (#1163) (c0c6c14f) by @vmoens ghstack-source-id: a88bebc23e6aaa02ec297db72dbda68ec9628ce7
[Feature] TensorDictModule in_keys allowed as Dict[str, tuple | list] to enable multi use of a sample feature (#1101) (e871b7df) by @bachdj-px
[Feature] UnbatchedTensor (#1170) (74cae098) by @vmoens ghstack-source-id: fa25726d61e913a725a71f1579eb06b09455e7c8
[Feature] intersection for assert_close (#1078) (84d31dba) by @vmoens ghstack-source-id: 3ae83c4ef90a9377405aebbf1761ace1a39417b1
[Feature] allow tensorclass to be customized (#1080) (31c73301) by @vmoens ghstack-source-id: 0b65b0a2dfb0cd7b5113e245c9444d3a0b55d085
[Feature] broadcast pointwise ops for tensor/tensordict mixed inputs (#1166) (aeff8376) by @vmoens ghstack-source-id: bbefbb1a2e9841847c618bb9cf49160ff1a5c36a
[Feature] compatibility of consolidate with compile (quick version) (#1061) (3cf52a0f) by @vmoens ghstack-source-id: 1bf3ca550dfe5499b58f878f72c4f1687b0f247e
[Feature] dist_params_keys and dist_sample_keys (#1179) (a728a4f7) by @vmoens ghstack-source-id: d1e53e780132d04ddf37d613358b24467520230f
[Feature] flexible return type when indexing prob sequences (#1189) (790bef60) by @vmoens ghstack-source-id: 74d28ee84d965c11c527c60b20d9123ef30007f6
[Feature] from_any with UserDict (#1106) (3485c2c8) by @vmoens ghstack-source-id: 420464209cff29c3a1c58ec521fbf4ed69d1355f
[Feature] inplace to method (#1066) (fbb71a4d) by @vmoens ghstack-source-id: 21cbc9f21287d3041c95d31fe2e5259b4ed36a42
[Feature] min, amin, max, amax, cummin, cummax (#1057) (e8c0087f) by @vmoens ghstack-source-id: 9873c08f98e84b372c6f701a3326e900454dc1d0
[Feature] repeat and repeat_interleave (#1115) (004f9796) by @vmoens ghstack-source-id: d90a1a7bd87115c5f7af1a413788a30cbc2096ee
[Feature] super() calls within TensorClass subclasses (#1133) (5349f2af) by @vmoens ghstack-source-id: 060a89982413869c54e1fb4aa74f90e2b9cdaac4
[Feature] tensorclass nocast (#1079) (5125217d) by @vmoens ghstack-source-id: edaba79a8a3b42cb3dac19b9fc145c1ceca4c70f
[Doc,BugFix,Feature] Better doc for reductions (#1122) (26110533) by @vmoens ghstack-source-id: 575544e823f29031739dc49561bc2a71125071ca
[Feature,Refactor] More args in constructors, refactor free functions (#1116) (b539bebe) by @vmoens ghstack-source-id: 35e2444bb5d4bf92b78437063e2f5aec83651713
[Feature,Refactor] Refactor from_dict, add from_any, from_dataclass (#1102) (c95a7039) by @vmoens ghstack-source-id: eb25fe4b201fd7f27d60b140278820c0d5d51eb8
[BugFix] Add checks and to_module to tensorclass accepted methods (#1124) (dd88950e) by @vmoens ghstack-source-id: 5a1cd0d2ac9a0880111f503fc9cb12519d85ef42
[BugFix] Better comparison of tensorclasses (#1137) (eb4a56ed) by @vmoens ghstack-source-id: 8def6f01f2b6d09714319a56f96b166ac1fd49d5
[BugFix] Better deterministic sample for composite (#1205) (5630fc82) by @vmoens ghstack-source-id: 9dd872e39ed9b697412ac8618b870b4d94670293
[BugFix] Better handling of struct arrays (#1193) (9b663389) by @vmoens ghstack-source-id: 37aecd40226db2362515fd00826bb336607e8749
[BugFix] Better repr of lazy stacks (#1076) (eaba7110) by @vmoens ghstack-source-id: 7256b4c95b239bf9e6467c0ea687abe2c9179922
[BugFix] Better return_log_prob=True for tensordict outputs (#1155) (aff4f4c1) by @vmoens ghstack-source-id: 977af3880f39cb341c1c715f1b8c9d59b7c580a0
[BugFix] Cloning empty tensordicts (#1119) (c38e256f) by @vmoens ghstack-source-id: f3db930052a3ff8d7e75e0d238a578c79acd6bd7
[BugFix] Computation of log prob in composite distribution for TensorDict without batch size (#1065) (e64a4c36) by @priba Co-authored-by: @priba pau.riba@helsing.ai
[BugFix] Consistent behavior for pad_sequence with one and many non-tensors (#1172) (be44018b) by @vmoens ghstack-source-id: c74edd95ed9846c14ffe26cb176d93c6e5e0dfbf
[BugFix] Do not unlock td if it's not locked in TDParams (for compile compat) (#1125) (4949efa8) by @vmoens ghstack-source-id: 9b6923f9c219e12af5560c97c1c6c58ed7870a8a
[BugFix] Ensure grads and noned when needed (#1069) (082d5429) by @vmoens ghstack-source-id: 5e9c5a974e5a5c73b033e5b85c3eb70c2f433512
[BugFix] Fix from_any tests (#1110) (e2444ed9) by @vmoens ghstack-source-id: 8c3b3d825555c727c7c18c7e8a87311f718a94b6
[BugFix] Fix grad tests (#1070) (752e6dc4) by @vmoens ghstack-source-id: 32c9ee0690db01b54de2e9cc0d83b595b87ac527
[BugFix] Fix mem leak when locking (#1188) (259c9411) by @vmoens ghstack-source-id: d6e44e1d9b9afc9903a0f45945c10a94dcf5a0ca
[BugFix] Fix none ref in during reduction (#1090) (c11024e0) by @vmoens
[BugFix] Fix typo in OneHotCategorical (#1130) (d7506c49) by @jgrhab
[BugFix] Fix unitary ops for tensorclass (#1164) (148c8231) by @vmoens ghstack-source-id: 2d117645769890b72f5856f68acbe1b48015cfbb
[BugFix] Fix update_ KeyError when a key from source is missing in dest (#1150) (efb89a6f) by @vmoens ghstack-source-id: 63013752ba61f05079cb6a60bb06312968b79ae9
[BugFix] Make min/max tensorclasses be interchangeable with PT equivalent (#1180) (44734854) by @vmoens
[BugFix] Proper masks for padding with custom pad value (#1185) (bbf773be) by @vmoens ghstack-source-id: 0580f89ce9bbaf5a13bab33f9c9b8f5a9e9df96f
[BugFix] Reference cycle in TensorDict._cast_reduction (#1056) (3316bdcc) by @egaznep
[BugFix] Same behaviour in compiled and non-compiled versions of new_unsafe (#1197) (64cfa585) by @vmoens ghstack-source-id: 075f953ca7fbd7fc54d797e12437db59b44bde03
[BugFix] Use 'spawn' mp context in all tests (#1111) (c8427308) by @vmoens ghstack-source-id: a7d786fe77c2c12d5c8c85579123a64ef5c87cf2
[BugFix] Use correct default cuda device (#1161) (ecb692e7) by @vmoens ghstack-source-id: 9afb5b03ddf75afec357e9e54caadfc92ebf4ded
[BugFix] auto-batch-size in dipatch (#1109) (2728dbf9) by @vmoens ghstack-source-id: ca5b36195c28da65a20d42699346fbc06083181c
[BugFix] calling pad with immutable sequence (#1075) (760c5374) by @SamAdamDay
[BugFix] clear_refs_for_compile() to clear weakrefs when compiling (#1196) (ba135417) by @vmoens ghstack-source-id: ecbad083704930313dcfdd5bf3f6bb6b984030e8
[BugFix] densify non tensor stack is a no-op (#1194) (9a25b88f) by @vmoens ghstack-source-id: edbc22ce562cd918ce5dd5c0441e47cdadf7d88a
[BugFix] fix inline TDParams kwargs for nontensordata (#1094) (a5656cb0) by @vmoens ghstack-source-id: afd50385b6b1e8bd8ccfaabfa387ca5611ca07e2
[BugFix] inline TDParams kwargs in prob modules (#1093) (0b7ce938) by @vmoens ghstack-source-id: 9fda35811b4656bd9939c9fb31cb253d7751b55c
[BugFix] select_out_keys for Prob sequential (#1103) (df61d64a) by @vmoens ghstack-source-id: a566ae225c54f07a680b4bf380b16d8e797f62ea
[BugFix] smarter check in set_interaction_type (#1088) (db2b5e65) by @vmoens ghstack-source-id: 1821309ad24827c22c40c41f3544e7a768325f72
[BE] Better warning for composite_lp_aggregate value (#1201) (0fdaba7b) by @vmoens ghstack-source-id: c0a801e4922878d65b0f81357afd7ecddc6943fb [BE] Check ordering and exclusivity of tensorclass registers (#1176) (b493178d) by @vmoens ghstack-source-id: 3dc907f4dd3047238adb0bb309d9ae75d24c5085 [BE] No warning if user sets the log_prob_key explicitly and only one variable is sampled from the ProbTDMod (#1209) (60e71796) by @vmoens ghstack-source-id: ccc966d5e698a4fb394081a92bafc31649951ab7 [BE] TensorClass stub method check (#1174) (22add568) by @vmoens ghstack-source-id: 5d85310221eb18ca6d0b1c4a4a88557f3fe8819d [BE] single dim check helper (#1192) (52dfeb87) by @vmoens ghstack-source-id: 6606e4b96061f73b98787b25129c29671a78dc1e [BE] tensorclass method registration check (#1175) (dedec043) by @vmoens ghstack-source-id: fa5dff657d58a035a05a39dbca84e3f9795c7fee [Quality] Avoid torch.distributed imports at root (#1134) (0e9a8544) by @vmoens ghstack-source-id: f773ed94d4b7ca13c603f32251f0735c751ebf94 [Quality] Better error message for incongruent lists of keys (#1077) (78b78022) by @vmoens ghstack-source-id: 34940a47d84bcf171bf4511187fcc82df88f801f [Quality] Better use of StrEnum in set_interaction_type (#1087) (79a33458) by @vmoens ghstack-source-id: c91a7a6be513fb46be6914df0b3bde779fa5528f
[Performance] Faster _is_tensor_collection in eager mode (#1060) (3963e51a) by @vmoens ghstack-source-id: b81c1d243a7c72a9d5fd68bf8e65e97a934ae61c [Performance] Faster copy of TDParams (#1096) (bbfe8c7f) by @vmoens [Performance] Faster to (#1073) (6272510f) by @vmoens ghstack-source-id: 3dfb0b66fae82dc8cf5ef2a14eccb1bec5237ebb
[Refactor] Add missing functions in tensorclass register (#1153) (05881f39) by @vmoens ghstack-source-id: 48311d7a98a9895b10e5552e5b4a4f13764607e0
[Refactor] Avoid TDParams parameters and buffers construction when obvious (#1100) (df870ef8) by @vmoens ghstack-source-id: bd701ecfaf68605801a215d3cd9d49268b888bb3
[Refactor] Avoid lambda functions in core functionality (#1136) (1d93434a) by @vmoens ghstack-source-id: bdd43bbcd353f1148b8b0da79f670e82c3b55c47
[Refactor] Better compile checks (#1139) (2aea3dd5) by @vmoens ghstack-source-id: c6a8d4587df45e374f0d6cb59fe1c982c7818276
[Refactor] Better handling of params and buffers in bytes (#1059) (3d3ea24a) by @vmoens ghstack-source-id: 87945c47b376d223bb3dc33bd6ec7cb9bb047455
[Refactor] Fix property handling in TC (#1178) (c7505980) by @vmoens ghstack-source-id: 8be779a7a85fdf45000181a9ea0f830822f19e37
[Refactor] Make CompositeDistribution a tensordict-exclusive class (#1112) (b4b8b31f) by @vmoens ghstack-source-id: 56c1dd2ad856a18613ec1a4c0ca70aedd28a52e3
[Refactor] Make _set_dispatch_td_nn_modules compatible with compile (#1084) (853b7d98) by @vmoens ghstack-source-id: 85a78cd6086233b414fcfe221dd8129e2e38f71c
[Refactor] Refactor context managers (#1098) (270d7bab) by @vmoens ghstack-source-id: c16baa83f6e41c4afd6637f3b3739d4e5cf25f1e
[Refactor] Refactor keys, items and values (#1058) (9232c467) by @vmoens ghstack-source-id: 9d5436c6bbc743e3c754d5fe5f6d87b005dde014
[Refactor] __eq__ to identity check in non-tensor stacking (#1083) (b39f0dbf) by @vmoens ghstack-source-id: ccbe882e12370b4145d7d834012cc3cfa6376f6c
[Refactor] composite_lp_aggregate to handle log-probs aggregates globally (#1181) (0013e38b) by @vmoens
[Doc] Add doc on export with nested keys (#1085) (3cb5855a) by @vmoens ghstack-source-id: 9c95e2dba6751d93c20c66d0dba0d4219dc61c0b
[Doc] Add missing classes to doc (#1203) (b4fd3804) by @vmoens ghstack-source-id: 2e174577aa33cc8d69c0f423c90ea2e5ee0fdef6
[Doc] Better doc for non-tensor data handling (#1173) (d3fba082) by @vmoens ghstack-source-id: c25571282d7bb63a14cc7b4ba9fb217785060cc7
[Doc] Better docstring for to_module (#1081) (9607cf09) by @vmoens ghstack-source-id: 16cedee8c0d38da6f377a262d5d7478a66fce07f
[Doc] Update distributed_replay_buffer.py reference (#1105) (d5fcace5) by @emmanuel-ferdman Signed-off-by: @emmanuel-ferdman emmanuelferdman@gmail.com
[CI] Add missing packages for smoke test (#1149) (42ad0660) by @vmoens ghstack-source-id: 1ae2795ee734baac4419b722d4e11d522051b112 [CI] Fix GLIBCXX_3.4.29 error (#1207) (bae04ced) by @vmoens [CI] Fix failing CI runs for release (#1206) (76e1810e) by @vmoens [CI] Fix linux_job_v2 after https://github.com/pytorch/test-infra/pull/6104 (#1191) (f6426cd9) by @amdfaa [CI] Fix nightly build (#1148) (b19fc054) by @vmoens ghstack-source-id: 46010c7ef465c2fdfe5422e094b5c227b67dbd4f [CI] Fix smoke tests (#1147) (24d8474a) by @vmoens ghstack-source-id: d529689fe2b08bed90f55cb2edd0592571619d85 [CI] Update the download-artifacts version (#1177) (a63e7f35) by @vmoens [CI] Upgrade CUDA to 12.4 (#1198) (1f35ab47) by @vmoens ghstack-source-id: 2d06089d4364d906581b60e9d5da1d65ecf01203 [CI] Upgrade to linux_job_v2 (#1097) (f606f7b1) by @vmoens ghstack-source-id: cf280238d4a241c02edb5d3e8ffb52c24fd228c9
[FBCode] Deactivate vmap monkey-patching in FBCode (#1135) (d394750c) by @vmoens ghstack-source-id: 4f1ebfd0fd4ff5b7378c8692a064406e72fc68c0 [Minor] Fix wrong fbcode manual links (#1064) (b06de95a) by @vmoens ghstack-source-id: 8ff9fb4a98ba3cbfc7c90d12a963a8a1e950e95b [Minor] print_directory_tree returns a string (#1086) (2b19ef11) by @vmoens ghstack-source-id: d57f19dd8efcef06676fca40a4d6f95367ff1d55 [Test] Ignore annoying jit_utils warning with device context manager (#1151) (c671d8ce) by @vmoens ghstack-source-id: 23fffb80e79bb839b34178cf5e20faea7a8115c5 [Test] fix inline TDParams kwargs for nontensordata (#1095) (978eb6c0) by @vmoens ghstack-source-id: da8b7f40d05715170a3e9f0b47763efe356afe5e [Tests] Skip deprecation warning tests on FB fbcode (#1128) (22da6791) by @vmoens ghstack-source-id: fb0cc381a670377667194324f5b019076b8e762d [Versioning] 0.6.2 (#1089) (73b0fd71) by @vmoens ghstack-source-id: 929ad7fb89e0bd6d25a70ed64b340ad7245fd693 [Versioning] python 3.8 compatibility fix (#1127) (d7529ab5) by @vmoens ghstack-source-id: ba7e9325c4125892522ee63253c148ed34adac7c [Versioning] v0.6.1 (#1072) (f12d31d1) by @vmoens ghstack-source-id: a899c95c12a3b1b986ed429b6507711c4126189e [Versioning] v0.7 version and deprecations (#1202) (02339fd7) by @vmoens ghstack-source-id: 6eea85522d8e55d3c1bbfd93ab8c551142de45c1
[BugFix] Fix none ref in during reduction (#1090) 88c86f83 [Versioning] v0.6.2 fix cdd5cb39 [BugFix] smarter check in set_interaction_type 0b3c7783 [V
Minor bug fixes
[BugFix] Fix none ref in during reduction (#1090) 88c86f83
[Versioning] v0.6.2 fix cdd5cb39
[BugFix] smarter check in set_interaction_type 0b3c7783
[Versioning] 0.6.2 b26bbe35
[Quality] Better use of StrEnum in set_interaction_type 477d85bc
[Minor] print_directory_tree returns a string 05c0fe7e
[Doc] Add doc on export with nested keys d64c33d8
[Feature] Better logs of key errors in assert_close 1ef1188d
[Refactor] Make _set_dispatch_td_nn_modules compatible with compile f24e3d84
[Doc] Better docstring for to_module 178dfd9d
[Feature] intersection for assert_close 85833929
[Quality] Better error message for incongruent lists of keys 866943c3
[BugFix] Better repr of lazy stacks e00965c8
[BugFix] calling pad with immutable sequence (#1075) d3bcb6ef
[Performance] Faster to f031bf26
Full Changelog: https://github.com/pytorch/tensordict/compare/v0.6.1...v0.6.2
This release offers some minor bug fixes in bytes (#1059), a better handling of edge cases in keys, values and items (#1058), better caching of grads
This release offers some minor bug fixes in bytes (#1059), a better handling of edge cases in keys, values and items (#1058), better caching of grads (#1069) and a broken reference cycle in reductions during calls to export (#1056).
[Deprecate] Make calls to make_functional error #1034 by @vmoens
TensorDict 0.6.0 makes the @dispatch decorator compatible with torch.export and related APIs,
allowing you to get rid of tensordict altogether when exporting your models:
from torch.export import export
model = Seq(
# 1. A small network for embedding
Mod(nn.Linear(3, 4), in_keys=["x"], out_keys=["hidden"]),
Mod(nn.ReLU(), in_keys=["hidden"], out_keys=["hidden"]),
Mod(nn.Linear(4, 4), in_keys=["hidden"], out_keys=["latent"]),
# 2. Extracting params
Mod(NormalParamExtractor(), in_keys=["latent"], out_keys=["loc", "scale"]),
# 3. Probabilistic module
Prob(
in_keys=["loc", "scale"],
out_keys=["sample"],
distribution_class=dists.Normal,
),
)
model_export = export(model, args=(), kwargs={"x": x})
See our new tutorial to learn more about this feature.
The library integration with the PT2 stack is also further improved by the introduction of CudaGraphModule,
which can be used to speed-up model execution under a certain set of assumptions; mainly that the inputs and outputs
are non-differentiable, that they are all tensors or constant and that the whole graph can be executed on cuda with
buffers of constant shape (ie, dynamic shape is not allowed).
We also introduce a new tutorial on streaming tensordicts.
Note: The aarch64 binaries are attached to these release notes and not available in PyPI at the moment.
existsok in memmap* methods #990 by @vmoensdata_ptr() method #1024 by @vmoensinplace arg in TDM constructor #992 by @vmoensselected_out_keys arg in TDS constructor #993 by @vmoens__name__ to TDModules #1045 by @vmoens__init__ (#1014) by @vmoens__setitem__ (#985) by @vmoensFull Changelog: https://github.com/pytorch/tensordict/compare/v0.5.0...v0.6.0
Co-authored-by: Vincent Moens vmoens@meta.com by @albertbou92
This release is packed with new features and performance improvements.
This release is packed with new features and performance improvements.
TensorDict.consolidateThere is now a TensorDict.consolidate method that will put all the tensors in a single storage. This will greatly speed-up serialization in multiprocessed and distributed settings.
TensorDict common ops (get, set, index, arithmetic ops etc) now work within torch.compile.
The list of supported operations can be found in test/test_compile.py. We encourage users to report any graph break caused by tensordict to us, as we are willing to improve the coverage as much as can be.
#807 enables python 3.12 support, a long awaited feature!
mean, std and other reduction methodsIt is now possible to get the grand average of a tensordict content using tensordict.mean(reduce=True).
This applies to mean, nanmean, prod, std, sum, nansum and var.
from_pytree and to_pytreeWe made it easy to convert a tensordict to a given pytree structure and build it from any pytree using to_pytree and from_pytree. #832
Similarly, conversion to namedtuple is now made easy thanks to #788.
map_iterOne can now iterate through a TensorDIct batch-dimension and apply a function on a separate process thanks to map_iter.
This should enable the construction of datasets using TensorDict, where the preproc step is executed on a separate process. #847
It is not possible to use flatten_keys and flatten as context managers (#908, #779):
with tensordict.flatten_keys() as flat_td:
flat_td["flat.key"] = 0
assert td["flat", "key"] == 0
We made it easy to build tensordicts with simple keyword arguments, like a dict is built in python:
td = TensorDict(a=0, b=1)
assert td["a"] == torch.tensor(0)
assert td["b"] == torch.tensor(1)
The batch_size is now optional for both tensordict and tensorclasses. #905
Thanks to #769, it is now possible to load a tensordict directly on a destination device (including "meta" device):
td = TensorDict.load(path, device=device)
to(device, pin_memory, num_threads) by @vmoens in https://github.com/pytorch/tensordict/pull/846out argument in apply by @vmoens in https://github.com/pytorch/tensordict/pull/794to for consolidated TDs by @vmoens in https://github.com/pytorch/tensordict/pull/851zero_grad and requires_grad_ by @vmoens in https://github.com/pytorch/tensordict/pull/901_make_dtype_promotion backward compat by @vmoens in https://github.com/pytorch/tensordict/pull/842pad_sequence behavior for non-tensor attributes of tensorclass by @kurtamohler in https://github.com/pytorch/tensordict/pull/884_tensordict, except for NonTensorData by @vmoens in https://github.com/pytorch/tensordict/pull/841_run_checks from __init__ by @vmoens in https://github.com/pytorch/tensordict/pull/843Full Changelog: https://github.com/pytorch/tensordict/compare/v0.4.0...v0.5.0
[Refactor] Cleanup deprecation of empty td filtering in apply by @vmoens in https://github.com/pytorch/tensordict/pull/665
This new version of tensordict comes with a great deal of new features:
You can now operate pointwise arithmetic operations on tensordict. For locked tensordicts and inplace operations such as += or data.mul_, fused cuda kernels will be used which will drastically improve the runtime.
See
Casting tensordicts to device is now much faster out-of-the box as data will be cast asynchronously (and it's safe too!) [BugFix,Feature] Optional non_blocking in set, to_module and update by @vmoens in https://github.com/pytorch/tensordict/pull/718 [BugFix] consistent use of non_blocking in tensordict and torch.Tensor by @vmoens in https://github.com/pytorch/tensordict/pull/734 [Feature] non_blocking=None by default by @vmoens in https://github.com/pytorch/tensordict/pull/748
The non-tensor data API has also been improved, see [BugFix] Allow inplace modification of non-tensor data in locked tds by @vmoens in https://github.com/pytorch/tensordict/pull/694 [BugFix] Fix inheritance from non-tensor by @vmoens in https://github.com/pytorch/tensordict/pull/709 [Feature] Allow non-tensordata to be shared across processes + memmap by @vmoens in https://github.com/pytorch/tensordict/pull/699 [Feature] Better detection of non-tensor data by @vmoens in https://github.com/pytorch/tensordict/pull/685
@tensorclass now supports automatic type casting: annotating a value as a tensor or an int can ensure that the value will be cast to that type if the tensorclass decorator takes the autocast=True argument
[Feature] Type casting for tensorclass by @vmoens in https://github.com/pytorch/tensordict/pull/735
TensorDict.map now supports the "fork" start method. Preallocated outputs are also a possibility. [Feature] mp_start_method in tensordict map by @vmoens in https://github.com/pytorch/tensordict/pull/695 [Feature] map with preallocated output by @vmoens in https://github.com/pytorch/tensordict/pull/667
Miscellaneous performance improvements [Performance] Faster flatten_keys by @vmoens in https://github.com/pytorch/tensordict/pull/727 [Performance] Faster update_ by @vmoens in https://github.com/pytorch/tensordict/pull/705 [Performance] Minor efficiency improvements by @vmoens in https://github.com/pytorch/tensordict/pull/703 [Performance] Random speedups by @albanD in https://github.com/pytorch/tensordict/pull/728 [Feature] Faster to(device) by @vmoens in https://github.com/pytorch/tensordict/pull/740
Finally, we have opened a discord channel for tensordict! [Badge] Discord shield by @vmoens in https://github.com/pytorch/tensordict/pull/736
We cleaned up the API a bit, creating a save and a load methods, or adding some utils such as fromkeys. One can also check if a key belongs to a tensordict as it is done with a regular dictionary with key in tensordict!
[Feature] contains, clear and fromkeys by @vmoens in https://github.com/pytorch/tensordict/pull/721
Thanks for all our contributors and community for the support!
Other PRs
pad_sequence refactoring by @dtsaras in https://github.com/pytorch/tensordict/pull/652__exit__ update when td is locked by @vmoens in https://github.com/pytorch/tensordict/pull/671reset_parameters_recursive is a no-op by @matteobettini in https://github.com/pytorch/tensordict/pull/693from_modules method for MOE / ensemble learning by @vmoens in https://github.com/pytorch/tensordict/pull/677Full Changelog: https://github.com/pytorch/tensordict/compare/v0.3.0...v0.4.0
[BugFix,Feature] Optional non_blocking in set, to_module and update (#718) [Refactor] Refactor contiguous (#716) [Test] Add proper tests for torch.sta
[BugFix,Feature] Optional non_blocking in set, to_module and update (#718) [Refactor] Refactor contiguous (#716) [Test] Add proper tests for torch.stack with lazy stacks (#715) [BugFix] Fix dense stack usage in torch.stack (#714) [BugFix] Dense stack lazy tds defaults to dense_stack_tds (#713) [Feature] Store non tensor stacks in a single json (#711) [Feature] TensorDict logger (#710) [BugFix, Feature] tensorclass.to_dict and from_dict (#707) [BugFix] Fix inheritance from non-tensor (#709) [Performance] Faster update_ (#705) [Benchmark] Benchmark update_ (#704) [Performance] Minor efficiency improvements (#703) [Feature] Allow non-tensordata to be shared across processes + memmap (#699) [CI] Unpin mpmath (#702) [CI] Remove snapshot from CI (#701) [BugFix] Support empty tuple in lazy stack indexing (#696) [CI] Pinning mpmath (#697) [BugFix] Allow inplace modification of non-tensor data in locked tds (#694) [Feature] Better detection of non-tensor data (#685) [Feature] Warn when reset_parameters_recursive is a no-op (#693) [BugFix,Feature] filter_empty in apply (#661)
See the release on PyPI: https://pypi.org/project/tensordict/0.3.2/
Solves several bugs and performance issues.
Solves several bugs and performance issues.
List of changes:
We deprecate MemmapTensor in favour of MemoryMappedTensor, which is fully backed by torch.Tensor and not numpy anymore. The new API is faster and way…
In this release we introduce a bunch of exciting features to TensorDict:
We deprecate MemmapTensor in favour of MemoryMappedTensor, which is fully backed by torch.Tensor and not numpy anymore. The new API is faster and way more bug-free than it used too. See #541
Saving tensordicts on disk can now be done via memmap, memmap_ and memmap_like which all support multithreading. If possible, serialization is pickle free (memmap + json) and torch.save is only used for classes that fail to be serialized with json. Serializing models using tensordict is now 3-10x faster than using torch.save, even for SOTA LLMs such as LLAMA.
TensorDict can now carry non tensor data through the NonTensorData class. Assigning non-tensor data can also be done via __setitem__ and they can be retrieved via __getitem__. #601
A bunch of new operations have appeared too such as named_apply (apply with key names) or tensordict.auto_batch_size_(), and operations like update can now be achieved for only a subset of keys.
Almost all operations in the library are now faster!
We are slowing deprecating lazy classes except for LazyStackedTensorDict. Whereas torch.stack used to systematically return a lazy stack, it now returns a dense stack if the set_lazy_legacy(mode) decorator is set to False (which will be the default in the next release). The old behaviour can be set with set_lazy_legacy(True). Lazy stacks can still be obtained using LazyStackedTensorDict.lazy_stack. Appropriate warnings are raised unless you have patched your code accordingly.
str-typed value to _device attribute in MemmapTensor by @kurt-stolle in https://github.com/pytorch/tensordict/pull/552__init__ by @vmoens in https://github.com/pytorch/tensordict/pull/576TensorDictModule.__getattr__ for private attributes by @vmoens in https://github.com/pytorch/tensordict/pull/579empty method by @vmoens in https://github.com/pytorch/tensordict/pull/622index in list error by @vmoens in https://github.com/pytorch/tensordict/pull/627auto_batch_size_ by @vmoens in https://github.com/pytorch/tensordict/pull/630Full Changelog: https://github.com/pytorch/tensordict/compare/v0.2.1...v0.3.0
[CI] Migration CCI->GHA by @vmoens in https://github.com/pytorch/tensordict/pull/536
to by @vmoens in https://github.com/pytorch/tensordict/pull/524pad argument to TensorDict.where by @vmoens in https://github.com/pytorch/tensordict/pull/539torch.where: Expand mask on the right in lazy stack tds by @vmoens in https://github.com/pytorch/tensordict/pull/542Full Changelog: https://github.com/pytorch/tensordict/compare/v0.2.0...v0.2.1
[Feature] First class dim compatibility by @vmoens in https://github.com/pytorch/tensordict/pull/525
dense_stack_tds by @matteobettini in https://github.com/pytorch/tensordict/pull/506[Benchmark] Benchmark flatten / unflatten by @vmoens in https://github.com/pytorch/tensordict/pull/439
[Benchmark] Benchmark items, values and keys by @vmoens in https://github.com/pytorch/tensordict/pull/428
[Benchmark] Benchmark leaf membership by @vmoens in https://github.com/pytorch/tensordict/pull/462
[Benchmark] Benchmark set nested by @vmoens in https://github.com/pytorch/tensordict/pull/459
[Benchmark] Benchmark to(device) by @vmoens in https://github.com/pytorch/tensordict/pull/385
[Benchmark] Benchmark unbind by @vmoens in https://github.com/pytorch/tensordict/pull/470
[Benchmark] Benchmark vmap and stacked tds by @vmoens in https://github.com/pytorch/tensordict/pull/430
[Benchmark] Fix GPU benchmark by @vmoens in https://github.com/pytorch/tensordict/pull/386
[Benchmark] Memmap tensordict benchmarks by @vmoens in https://github.com/pytorch/tensordict/pull/432
[Benchmark] More benchmarks for get by @vmoens in https://github.com/pytorch/tensordict/pull/477
[Benchmark] More benchmarks for locked tds by @vmoens in https://github.com/pytorch/tensordict/pull/437
[Benchmark] Total time per benchmark for efficiency by @vmoens in https://github.com/pytorch/tensordict/pull/446
[BugFix] Fix EnsembleModule parameters exposure by @vmoens in https://github.com/pytorch/tensordict/pull/508
[BugFix] Fix None device when loading by @vmoens in https://github.com/pytorch/tensordict/pull/425
[BugFix] Fix benchmarks in PR comparison by @vmoens in https://github.com/pytorch/tensordict/pull/442
[BugFix] Fix benchmarks in PR comparison by @vmoens in https://github.com/pytorch/tensordict/pull/447
[BugFix] Fix error for incongruent devices in load_state_dict by @vmoens in https://github.com/pytorch/tensordict/pull/529
[BugFix] Fix error in state_dict tests by @vmoens in https://github.com/pytorch/tensordict/pull/531
[BugFix] Fix het lazy stack ops by @vmoens in https://github.com/pytorch/tensordict/pull/416
[BugFix] Fix hook call by @vmoens in https://github.com/pytorch/tensordict/pull/457
[BugFix] Fix indexing of lazy stacks by @vmoens in https://github.com/pytorch/tensordict/pull/488
[BugFix] Fix lazy stack / stack_onto + masking lazy stacks by @vmoens in https://github.com/pytorch/tensordict/pull/497
[BugFix] Fix lazy stack lock by @vmoens in https://github.com/pytorch/tensordict/pull/452
[BugFix] Fix lazy-stack to call to to_tensordict() by @vmoens in https://github.com/pytorch/tensordict/pull/519
[BugFix] Fix logical operations in MemmapTensor by @vmoens in https://github.com/pytorch/tensordict/pull/484
[BugFix] Fix max length by @vmoens in https://github.com/pytorch/tensordict/pull/408
[BugFix] Fix opt deps by @vmoens in https://github.com/pytorch/tensordict/pull/464
[BugFix] Fix params.clone by @vmoens in https://github.com/pytorch/tensordict/pull/509
[BugFix] Fix pickle PR by @vmoens in https://github.com/pytorch/tensordict/pull/397
[BugFix] Fix regex in tensorclass error checks by @vmoens in https://github.com/pytorch/tensordict/pull/398
[BugFix] Fix state-dict by @vmoens in https://github.com/pytorch/tensordict/pull/528
[BugFix] Fix test rename_key by @vmoens in https://github.com/pytorch/tensordict/pull/441
[BugFix] Fix unbind by @vmoens in https://github.com/pytorch/tensordict/pull/471
[BugFix] Fix wheels and nightly by @vmoens in https://github.com/pytorch/tensordict/pull/417
[BugFix] Fix writing on stacked tds in vmap by @vmoens in https://github.com/pytorch/tensordict/pull/456
[BugFix] Heterogeneous lazy stack apply(_) by @vmoens in https://github.com/pytorch/tensordict/pull/400
[BugFix] Keyword arguments with dispatch + select keys by @vmoens in https://github.com/pytorch/tensordict/pull/409
[BugFix] Nested keys to probabilistic modules by @matteobettini in https://github.com/pytorch/tensordict/pull/479
[BugFix] Overload expand by @vmoens in https://github.com/pytorch/tensordict/pull/530
[BugFix] Repr of het stacked tds by @vmoens in https://github.com/pytorch/tensordict/pull/399
[BugFix] Retrieve all keys for LazyStackedTensorDict by @matteobettini in https://github.com/pytorch/tensordict/pull/512
[BugFix] Torch 1.13 compat by @vmoens in https://github.com/pytorch/tensordict/pull/421
[BugFix] Torch 1.13 compat by @vmoens in https://github.com/pytorch/tensordict/pull/424
[BugFix] Torch 1.13 compat by @vmoens in https://github.com/pytorch/tensordict/pull/429
[BugFix] Unflatten names by @vmoens in https://github.com/pytorch/tensordict/pull/472
[BugFix] fix infinite loop by @xmaples in https://github.com/pytorch/tensordict/pull/392
[BugFix] from_module without detach by @vmoens in https://github.com/pytorch/tensordict/pull/466
[BugFix] stateless lstms by @vmoens in https://github.com/pytorch/tensordict/pull/475
[Bugfix] Indexing memmaps with torch.Tensor by @tcbegley in https://github.com/pytorch/tensordict/pull/383
[Bugfix] fix getitem by @apbard in https://github.com/pytorch/tensordict/pull/384
[CI] Concurrency on gha by @vmoens in https://github.com/pytorch/tensordict/pull/387
[CI] Fix Doc CI bis by @matteobettini in https://github.com/pytorch/tensordict/pull/490
[CI] Fix doc CI by @matteobettini in https://github.com/pytorch/tensordict/pull/489
[CI] Python 3.11 compatibility by @vmoens in https://github.com/pytorch/tensordict/pull/516
[CI] Removing myself from notified users in benchmarks by @tcbegley in https://github.com/pytorch/tensordict/pull/483
[CI] nightly run by @vmoens in https://github.com/pytorch/tensordict/pull/422
[CI] nightly test by @vmoens in https://github.com/pytorch/tensordict/pull/420
[Doc] Fix missing packaging package by @vmoens in https://github.com/pytorch/tensordict/pull/449
[Doc] Update citation by @vmoens in https://github.com/pytorch/tensordict/pull/404
[Feature] Add EnsembleModule by @smorad in https://github.com/pytorch/tensordict/pull/485
[Feature] Add reset_parameters support to TensorDictModuleBase by @smorad in https://github.com/pytorch/tensordict/pull/476
[Feature] Better error message in tdmodule by @vmoens in https://github.com/pytorch/tensordict/pull/390
[Feature] Cat LazyStackedTensorDicts by @matteobettini in https://github.com/pytorch/tensordict/pull/499
[Feature] Create memmaps folders with makesdir by @vmoens in https://github.com/pytorch/tensordict/pull/431
[Feature] Functional reset_parameters by @smorad in https://github.com/pytorch/tensordict/pull/478
[Feature] Kwargs for tensordict modules by @vmoens in https://github.com/pytorch/tensordict/pull/405
[Feature] Module serialization by @vmoens in https://github.com/pytorch/tensordict/pull/396
[Feature] Module to add state independent normal scale by @albertbou92 in https://github.com/pytorch/tensordict/pull/515
[Feature] Print lazy fields by @matteobettini in https://github.com/pytorch/tensordict/pull/493
[Feature] Registering a tensordict as a module buffer by @vmoens in https://github.com/pytorch/tensordict/pull/395
[Feature] Support classmethods in tensorclass by @vmoens in https://github.com/pytorch/tensordict/pull/448
[Feature] TensorDict.map by @vmoens in https://github.com/pytorch/tensordict/pull/518
[Feature] Torch 1.13 compatibility by @vmoens in https://github.com/pytorch/tensordict/pull/414
[Feature] Unravel nested keys by @vmoens in https://github.com/pytorch/tensordict/pull/403
[Feature] as_decorator and enter exit methods by @vmoens in https://github.com/pytorch/tensordict/pull/455
[Feature] batch_dims for from_dict by @vmoens in https://github.com/pytorch/tensordict/pull/460
[Feature] unravel_key_list by @vmoens in https://github.com/pytorch/tensordict/pull/474
[Minor] Add dist/ to gitignore. by @skandermoalla in https://github.com/pytorch/tensordict/pull/492
[Minor] Better errors for kwargs and key selection by @vmoens in https://github.com/pytorch/tensordict/pull/415
[Minor] Fix typo in _propagate_unlock by @vmoens in https://github.com/pytorch/tensordict/pull/443
[Minor] Key filtering error fixes during deletion by @vmoens in https://github.com/pytorch/tensordict/pull/481
[Misc] Move to PyTorch compat by @vmoens in https://github.com/pytorch/tensordict/pull/533
[Perf] Perf improvements by @vmoens in https://github.com/pytorch/tensordict/pull/412
[Performance] Better names handling in LazyStackTD by @vmoens in https://github.com/pytorch/tensordict/pull/482
[Performance] Faster functional by @vmoens in https://github.com/pytorch/tensordict/pull/418
[Performance] Faster keys by @vmoens in https://github.com/pytorch/tensordict/pull/427
[Performance] Faster vmap by @vmoens in https://github.com/pytorch/tensordict/pull/426
[Performance] Refactor get by @vmoens in https://github.com/pytorch/tensordict/pull/445
[Performance] Speed improvements by @vmoens in https://github.com/pytorch/tensordict/pull/410
[Performance] use dict for modules in swap states by @vmoens in https://github.com/pytorch/tensordict/pull/419
[Refactor] Better key membership check by @vmoens in https://github.com/pytorch/tensordict/pull/461
[Refactor] Better key membership check by @vmoens in https://github.com/pytorch/tensordict/pull/465
[Refactor] Better keys for cache to reduce mem footprint by @vmoens in https://github.com/pytorch/tensordict/pull/454
[Refactor] Bottom-up get_tuple by @vmoens in https://github.com/pytorch/tensordict/pull/451
[Refactor] Faster benchmarks by @vmoens in https://github.com/pytorch/tensordict/pull/453
[Refactor] Faster creation on device by @vmoens in https://github.com/pytorch/tensordict/pull/389
[Refactor] Faster indexing of memmap tensors by @vmoens in https://github.com/pytorch/tensordict/pull/433
[Refactor] Fix warnings by @vmoens in https://github.com/pytorch/tensordict/pull/402
[Refactor] Minor refactorings to tensorclass by @vmoens in https://github.com/pytorch/tensordict/pull/435
[Refactor] Refactor to create empty tds where necessary by @vmoens in https://github.com/pytorch/tensordict/pull/522
[Refactor] Refactor to() by @vmoens in https://github.com/pytorch/tensordict/pull/381
[Refactor] Simplify exclude by @vmoens in https://github.com/pytorch/tensordict/pull/391
[Refactor] TensorDict.to() signature by @vmoens in https://github.com/pytorch/tensordict/pull/520
[Refactor] _set_str and _set_tuple by @vmoens in https://github.com/pytorch/tensordict/pull/480
[Refactor] unravel_key fixes by @vmoens in https://github.com/pytorch/tensordict/pull/473
[Setup] Update setup.py python versions by @vmoens in https://github.com/pytorch/tensordict/pull/521
[Test,CI,Feature] Total time per test by @vmoens in https://github.com/pytorch/tensordict/pull/407
@xmaples made their first contribution in https://github.com/pytorch/tensordict/pull/392
@smorad made their first contribution in https://github.com/pytorch/tensordict/pull/476
@matteobettini made their first contribution in https://github.com/pytorch/tensordict/pull/479
@skandermoalla made their first contribution in https://github.com/pytorch/tensordict/pull/492
@albertbou92 made their first contribution in https://github.com/pytorch/tensordict/pull/515
Full Changelog: https://github.com/pytorch/tensordict/compare/v0.1.2...v0.2.0
Nothing published for this version
[Benchmark] More indexing benchmarks by @vmoens in https://github.com/pytorch-labs/tensordict/pull/375
Full Changelog: https://github.com/pytorch-labs/tensordict/compare/v0.1.1...v0.1.2
[Refactor] Deprecate set_default by @tcbegley in https://github.com/pytorch-labs/tensordict/pull/236
TensorDictBase available at root by @vmoens in https://github.com/pytorch-labs/tensordict/pull/295__setitem__ with broadcasting of tensordicts by @vmoens in https://github.com/pytorch-labs/tensordict/pull/316__new__ attribute and property creation by @vmoens in https://github.com/pytorch-labs/tensordict/pull/353Full Changelog: https://github.com/pytorch-labs/tensordict/compare/0.1.0...v0.1.1
First official release of tensordict!
First official release of tensordict!
__getitems__ for tensorclass by @vmoens in https://github.com/pytorch-labs/tensordict/pull/277Full Changelog: https://github.com/pytorch-labs/tensordict/compare/0.0.3...v0.1.0
[Refactor] Deprecating MetaTensor by @vmoens in https://github.com/pytorch-labs/tensordict/pull/168
TensorDict.split method by @wonnor-pro in https://github.com/pytorch-labs/tensordict/pull/56range in indexing operations by @tcbegley in https://github.com/pytorch-labs/tensordict/pull/139__post_init__ support by @tcbegley in https://github.com/pytorch-labs/tensordict/pull/172__setitem__ must allow non-type stringent writing by @vmoens in https://github.com/pytorch-labs/tensordict/pull/203_getitem_batch_size in various edge cases. by @tcbegley in https://github.com/pytorch-labs/tensordict/pull/211__eq__ refactoring by @vmoens in https://github.com/pytorch-labs/tensordict/pull/210clear_device by @tcbegley in https://github.com/pytorch-labs/tensordict/pull/242Full Changelog: https://github.com/pytorch-labs/tensordict/compare/0.0.1...0.0.3
[Refactor] Deprecating MetaTensor by @vmoens in https://github.com/pytorch-labs/tensordict/pull/168
TensorDict.split method by @wonnor-pro in https://github.com/pytorch-labs/tensordict/pull/56range in indexing operations by @tcbegley in https://github.com/pytorch-labs/tensordict/pull/139__post_init__ support by @tcbegley in https://github.com/pytorch-labs/tensordict/pull/172__setitem__ must allow non-type stringent writing by @vmoens in https://github.com/pytorch-labs/tensordict/pull/203_getitem_batch_size in various edge cases. by @tcbegley in https://github.com/pytorch-labs/tensordict/pull/211__eq__ refactoring by @vmoens in https://github.com/pytorch-labs/tensordict/pull/210Full Changelog: https://github.com/pytorch-labs/tensordict/compare/0.0.1...v0.0.2b
[Refactor] Deprecating MetaTensor by @vmoens in https://github.com/pytorch-labs/tensordict/pull/168
TensorDict.split method by @wonnor-pro in https://github.com/pytorch-labs/tensordict/pull/56range in indexing operations by @tcbegley in https://github.com/pytorch-labs/tensordict/pull/139__post_init__ support by @tcbegley in https://github.com/pytorch-labs/tensordict/pull/172Full Changelog: https://github.com/pytorch-labs/tensordict/compare/0.0.1...v0.0.2a
Introduces a new @tensorclass prototype. Check it out here
TensorDict.split method by @wonnor-pro in https://github.com/pytorch-labs/tensordict/pull/56range in indexing operations by @tcbegley in https://github.com/pytorch-labs/tensordict/pull/139Full Changelog: https://github.com/pytorch-labs/tensordict/compare/0.0.1...0.0.1c
TensorDict makes it easy to organise data and write reusable, generic PyTorch code. Originally developed for TorchRL, we’ve spun it out into a separat
TensorDict makes it easy to organise data and write reusable, generic PyTorch code. Originally developed for TorchRL, we’ve spun it out into a separate library.
TensorDict is primarily a dictionary but also a tensor-like class: it supports multiple tensor operations that are mostly shape and storage-related. It is designed to be efficiently serialised or transmitted from node to node or process to process. Finally, it is shipped with its own tensordict.nn module which is compatible with functorch and aims at making model ensembling and parameter manipulation easier.
On this page we will motivate TensorDict and give some examples of what it can do.
TensorDict allows you to write generic code modules that are re-usable across paradigms. For instance, the following loop can be re-used across most SL, SSL, UL and RL tasks.
>>> for i, tensordict in enumerate(dataset):
... # the model reads and writes tensordicts
... tensordict = model(tensordict)
... loss = loss_module(tensordict)
... loss.backward()
... optimizer.step()
... optimizer.zero_grad()
With its tensordict.nn module, the package provides many tools to use TensorDict in a code base with little or no effort.
In multiprocessing or distributed settings, tensordict allows you to seamlessly dispatch data to each worker:
>>> # creates batches of 10 datapoints
>>> splits = torch.arange(tensordict.shape[0]).split(10)
>>> for worker in range(workers):
... idx = splits[worker]
... pipe[worker].send(tensordict[idx])
Some operations offered by TensorDict can be done via tree_map too, but with a greater degree of complexity:
>>> td = TensorDict(
... {"a": torch.randn(3, 11), "b": torch.randn(3, 3)}, batch_size=3
... )
>>> regular_dict = {"a": td["a"], "b": td["b"]}
>>> td0, td1, td2 = td.unbind(0)
>>> # similar structure with pytree
>>> regular_dicts = tree_map(lambda x: x.unbind(0))
>>> regular_dict1, regular_dict2 regular_dict3 = [
... {"a": regular_dicts["a"][i], "b": regular_dicts["b"][i]}
... for i in range(3)]
The nested case is even more compelling:
>>> td = TensorDict(
... {"a": {"c": torch.randn(3, 11)}, "b": torch.randn(3, 3)}, batch_size=3
... )
>>> regular_dict = {"a": {"c": td["a", "c"]}, "b": td["b"]}
>>> td0, td1, td2 = td.unbind(0)
>>> # similar structure with pytree
>>> regular_dicts = tree_map(lambda x: x.unbind(0))
>>> regular_dict1, regular_dict2 regular_dict3 = [
... {"a": {"c": regular_dicts["a"]["c"][i]}, "b": regular_dicts["b"][i]}
... for i in range(3)
Decomposing the output dictionary in three similarly structured dictionaries after applying the unbind operation quickly becomes significantly cumbersome when working naively with pytree. With tensordict, we provide a simple API for users that want to unbind or split nested structures, rather than computing a nested split / unbound nested structure.
A TensorDict is a dict-like container for tensors. To instantiate a TensorDict, you must specify key-value pairs as well as the batch size. The leading dimensions of any values in the TensorDict must be compatible with the batch size.
>>> import torch
>>> from tensordict import TensorDict
>>> tensordict = TensorDict(
... {"zeros": torch.zeros(2, 3, 4), "ones": torch.ones(2, 3, 4, 5)},
... batch_size=[2, 3],
... )
The syntax for setting or retrieving values is much like that for a regular dictionary.
>>> zeros = tensordict["zeros"]
>>> tensordict["twos"] = 2 * torch.ones(2, 3)
One can also index a tensordict along its batch_size which makes it possible to obtain congruent slices of data in just a few characters (notice that indexing the nth leading dimensions with tree_map using an ellipsis would require a bit more coding):
>>> sub_tensordict = tensordict[..., :2]
One can also use the set method with inplace=True or the set_ method to do inplace updates of the contents. The former is a fault-tolerant version of the latter: if no matching key is found, it will write a new one.
The contents of the TensorDict can now be manipulated collectively. For example, to place all of the contents onto a particular device one can simply do
>>> tensordict = tensordict.to("cuda:0")
To reshape the batch dimensions one can do
>>> tensordict = tensordict.reshape(6)
The class supports many other operations, including squeeze, unsqueeze, view, permute, unbind, stack, cat and many more. If an operation is not present, the TensorDict.apply method will usually provide the solution that was needed.
The values in a TensorDict can themselves be TensorDicts (the nested dictionaries in the example below will be converted to nested TensorDicts).
>>> tensordict = TensorDict(
... {
... "inputs": {
... "image": torch.rand(100, 28, 28),
... "mask": torch.randint(2, (100, 28, 28), dtype=torch.uint8)
... },
... "outputs": {"logits": torch.randn(100, 10)},
... },
... batch_size=[100],
... )
Accessing or setting nested keys can be done with tuples of strings
>>> image = tensordict["inputs", "image"]
>>> logits = tensordict.get(("outputs", "logits")) # alternative way to access
>>> tensordict["outputs", "probabilities"] = torch.sigmoid(logits)
Some operations on TensorDict defer execution until items are accessed. For example stacking, squeezing, unsqueezing, permuting batch dimensions and creating a view are not executed immediately on all the contents of the TensorDict. Instead they are performed lazily when values in the TensorDict are accessed. This can save a lot of unnecessary calculation should the TensorDict contain many values.
>>> tensordicts = [TensorDict({
... "a": torch.rand(10),
... "b": torch.rand(10, 1000, 1000)}, [10])
... for _ in range(3)]
>>> stacked = torch.stack(tensordicts, 0) # no stacking happens here
>>> stacked_a = stacked["a"] # we stack the a values, b values are not stacked
It also has the advantage that we can manipulate the original tensordicts in a stack:
>>> stacked["a"] = torch.zeros_like(stacked["a"])
>>> assert (tensordicts[0]["a"] == 0).all()
The caveat is that the get method has now become an expensive operation and, if repeated many times, may cause some overhead. One can avoid this by simply calling tensordict.contiguous() after the execution of stack. To further mitigate this, TensorDict comes with its own meta-data class (MetaTensor) that keeps track of the type, shape, dtype and device of each entry of the dict, without performing the expensive operation.
Suppose we have some function foo() -> TensorDict and that we do something like the following:
>>> tensordict = TensorDict({}, batch_size=[N])
>>> for i in range(N):
... tensordict[i] = foo()
When i == 0 the empty TensorDict will automatically be populated with empty tensors with batch size N. In subsequent iterations of the loop the updates will all be written in-place.
To make it easy to integrate TensorDict in one’s code base, we provide a tensordict.nn package that allows users to pass TensorDict instances to nn.Module objects.
TensorDictModule wraps nn.Module and accepts a single TensorDict as an input. You can specify where the underlying module should take its input from, and where it should write its output. This is a key reason we can write reusable, generic high-level code such as the training loop in the motivation section.
>>> from tensordict.nn import TensorDictModule
>>> class Net(nn.Module):
... def __init__(self):
... super().__init__()
... self.linear = nn.LazyLinear(1)
...
... def forward(self, x):
... logits = self.linear(x)
... return logits, torch.sigmoid(logits)
>>> module = TensorDictModule(
... Net(),
... in_keys=["input"],
... out_keys=[("outputs", "logits"), ("outputs", "probabilities")],
... )
>>> tensordict = TensorDict({"input": torch.randn(32, 100)}, [32])
>>> tensordict = module(tensordict)
>>> # outputs can now be retrieved from the tensordict
>>> logits = tensordict["outputs", "logits"]
>>> probabilities = tensordict.get(("outputs", "probabilities"))
To facilitate the adoption of this class, one can also pass the tensors as kwargs:
>>> tensordict = module(input=torch.randn(32, 100))
which will return a TensorDict identical to the one in the previous code box.
A key pain-point of multiple PyTorch users is the inability of nn.Sequential to handle modules with multiple inputs. Working with key-based graphs can easily solve that problem as each node in the sequence knows what data needs to be read and where to write it.
For this purpose, we provide the TensorDictSequential class which passes data through a sequence of TensorDictModules. Each module in the sequence takes its input from, and writes its output to the original TensorDict, meaning it’s possible for modules in the sequence to ignore output from their predecessors, or take additional input from the tensordict as necessary. Here’s an example.
>>> class Net(nn.Module):
... def __init__(self, input_size=100, hidden_size=50, output_size=10):
... super().__init__()
... self.fc1 = nn.Linear(input_size, hidden_size)
... self.fc2 = nn.Linear(hidden_size, output_size)
...
... def forward(self, x):
... x = torch.relu(self.fc1(x))
... return self.fc2(x)
...
... class Masker(nn.Module):
... def forward(self, x, mask):
... return torch.softmax(x * mask, dim=1)
>>> net = TensorDictModule(
... Net(), in_keys=[("input", "x")], out_keys=[("intermediate", "x")]
... )
>>> masker = TensorDictModule(
... Masker(),
... in_keys=[("intermediate", "x"), ("input", "mask")],
... out_keys=[("output", "probabilities")],
... )
>>> module = TensorDictSequential(net, masker)
>>> tensordict = TensorDict(
... {
... "input": TensorDict(
... {"x": torch.rand(32, 100), "mask": torch.randint(2, size=(32, 10))},
... batch_size=[32],
... )
... },
... batch_size=[32],
... )
>>> tensordict = module(tensordict)
>>> intermediate_x = tensordict["intermediate", "x"]
>>> probabilities = tensordict["outputs", "probabilities"]
In this example, the second module combines the output of the first with the mask stored under (“inputs”, “mask”) in the TensorDict.
TensorDictSequential offers a bunch of other features: one can access the list of input and output keys by querying the in_keys and out_keys attributes. It is also possible to ask for a sub-graph by querying select_subsequence() with the desired sets of input and output keys that are desired. This will return another TensorDictSequential with only the modules that are indispensable to satisfy those requirements. The TensorDictModule is also compatible with vmap and other functorch capabilities.
Functional Programming We provide and API to use TensorDict in conjunction with functorch. For instance, TensorDict makes it easy to concatenate model weights to do model ensembling:
>>> from torch import nn
>>> from tensordict import TensorDict
>>> from tensordict.nn import make_functional
>>> import torch
>>> from functorch import vmap
>>> layer1 = nn.Linear(3, 4)
>>> layer2 = nn.Linear(4, 4)
>>> model = nn.Sequential(layer1, layer2)
>>> # we represent the weights hierarchically
>>> weights1 = TensorDict(layer1.state_dict(), []).unflatten_keys(separator=".")
>>> weights2 = TensorDict(layer2.state_dict(), []).unflatten_keys(separator=".")
>>> params = make_functional(model)
>>> # params provided by make_functional match state_dict:
>>> assert (params == TensorDict({"0": weights1, "1": weights2}, [])).all()
>>> # Let's use our functional module
>>> x = torch.randn(10, 3)
>>> out = model(x, params=params) # params is the last arg (or kwarg)
>>> # an esemble of models: we stack params along the first dimension...
>>> params_stack = torch.stack([params, params], 0)
>>> # ... and use it as an input we'd like to pass through the model
>>> y = vmap(model, (None, 0))(x, params_stack)
>>> print(y.shape)
torch.Size([2, 10, 4])
The functional API is comparable if not faster than the current FunctionalModule implemented in functorch.
TensorDict.split method by @wonnor-pro in https://github.com/pytorch-labs/tensordict/pull/56Full Changelog: https://github.com/pytorch-labs/tensordict/compare/0.0.1...0.0.1b
Nothing published for this version
Your coding agent can read these notes before it upgrades. Set up the MCP server →