Unified TensorDistribution
Overview
TensorContainer 0.6.3 introduces significant improvements to the tensor distribution module, featuring unified implementations, enhanced performance, and better developer experience. This release focuses on simplifying distribution classes while maintaining full compatibility with existing code.
What's New
Major Improvements
Unified TensorDistribution Implementation
- Comprehensive Refactoring: All tensor distribution classes have been unified under a consistent implementation pattern
- Code Reduction: Removed over 1000 lines of redundant code while adding new functionality
- Better Inheritance: Enhanced class hierarchy to properly leverage parent class functionality
- Consistent APIs: All distributions now follow the same patterns for parameter handling and validation
Enhanced Parameter Handling
- Unified Broadcasting: Implemented consistent parameter broadcasting across all distributions using
broadcast_all - Improved Validation: Better error messages and validation logic for distribution parameters
- Scalar Parameter Support: Enhanced handling of scalar parameters in distribution initialization
New Features
unflattenDistribution Method: New functionality in the base distribution class for handling flattened parameters- Enhanced
validate_argsSupport: Improved argument validation across all distribution classes - TensorOneHotCategoricalStraightThrough: New distribution implementation for straight-through gradient estimation
Performance & Quality Improvements
Code Quality
- Modern Type Hints: Updated to use Python 3.10+ union syntax (
|instead ofUnion) - Better Error Handling: Replaced generic RuntimeError with more specific ValueError for parameter validation
- Consistent Code Style: Unified formatting and style across all distribution implementations
Testing Enhancements
- Expanded Test Coverage: Added comprehensive tests for pytree integration and parameter validation
- Better Error Testing: Improved test cases for error handling and edge cases
- Parameter Factory Fixtures: Added reusable test fixtures for distribution parameter testing
Technical Details
Changed Components
- 55 files modified across the tensor distribution module
- All 30+ distribution classes refactored for consistency:
- Bernoulli, Beta, Binomial, Categorical, Cauchy, Chi2
- ContinuousBernoulli, Dirichlet, Exponential, FisherSnedecor
- Gamma, Geometric, Gumbel, HalfCauchy, HalfNormal
- InverseGamma, Kumaraswamy, Laplace, LogisticNormal
- Multinomial, MultivariateNormal, NegativeBinomial, Normal
- OneHotCategorical, Pareto, Poisson, RelaxedBernoulli
- RelaxedOneHotCategorical, StudentT, TanhNormal
- TruncatedNormal, Uniform, VonMises, Weibull, Wishart
Breaking Changes
While most user code should continue to work without changes, there are some minor breaking changes:
- Some
RuntimeErrorinstances replaced withValueErrorfor better parameter validation - Simplified initialization signatures in some distribution classes
Migration Guide
Most existing code will continue to work without modification. For cases where changes are needed:
- Update error handling code to catch
ValueErrorinstead ofRuntimeErrorfor parameter validation - Review distribution initialization if using advanced parameter configurations
Compatibility
- Python: 3.9+ (dropped Python 3.8 support)
- PyTorch: 2.0+