0.8.1: Enhanced PyTree Error Diagnostics and Refactoring
Enhanced PyTree Error Diagnostics and Refactoring
This update introduces significant enhancements to error reporting for PyTree operations (e.g., torch.stack, torch.cat) on TensorContainer subclasses such as TensorDict and TensorDataClass. The primary objective is to provide developers with clear, actionable diagnostics when structural mismatches occur, thereby improving the debugging workflow. This is accompanied by a refactoring of the underlying PyTree context implementation for improved modularity and extensibility.
Key Changes
1. Detailed Mismatch Diagnosis
A new utility, diagnose_pytree_structure_mismatch, has been implemented to provide detailed diagnostics for structural differences. This replaces generic RuntimeError or ValueError exceptions with specific, informative messages that identify the root cause of an issue. The diagnostics cover:
- Key/Field Mismatches: Differences in keys for
TensorDictor fields forTensorDataClass. - Type Mismatches: Incompatible container types (e.g.,
TensorDictvs.list). - Device Mismatches: Containers located on different devices (e.g.,
cpuvs.cuda). - Nesting Mismatches: Discrepancies in the nested structure of containers.
2. Example Error Message
When attempting to stack TensorContainer instances with incompatible keys, the new error message provides a precise description of the problem:
Structure mismatch at container: containers have incompatible layouts.
Container 0 at container: TensorDict(keys=['a', 'b'], device=cpu)
Container 1 at container: TensorDict(keys=['a', 'c'], device=cpu)
Fix: Key mismatch detected. Missing keys in container 1: ['b']. Extra keys in container 1: ['c'].
3. Internal Refactoring
To support these improvements, the following architectural changes were made:
- Structured Exceptions: A set of specific exception classes (
TypeMismatch,ContextMismatch,KeyPathMismatch) were introduced to represent distinct error conditions. - PyTree Context Abstraction: A new
ContextWithAnalysisabstract base class standardizes structural analysis, andTensorContainerPytreeContextcentralizes common logic to reduce code duplication. - Dataclass Conversion: PyTree context classes were converted from
NamedTupletodataclassto improve readability and extensibility.
What's Changed
- Improved PyTree Error Messages and Refactoring by @mctigger in #21
- fix(build): bump project version to 0.8.1 by @mctigger in #22
Full Changelog: 0.8.0...0.8.1