Skip to content

0.8.1: Enhanced PyTree Error Diagnostics and Refactoring

Choose a tag to compare

@mctigger mctigger released this 10 Sep 11:08
· 45 commits to main since this release
f93849f

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 TensorDict or fields for TensorDataClass.
  • Type Mismatches: Incompatible container types (e.g., TensorDict vs. list).
  • Device Mismatches: Containers located on different devices (e.g., cpu vs. 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 ContextWithAnalysis abstract base class standardizes structural analysis, and TensorContainerPytreeContext centralizes common logic to reduce code duplication.
  • Dataclass Conversion: PyTree context classes were converted from NamedTuple to dataclass to 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