Skip to content

Unified TensorDistribution

Choose a tag to compare

@mctigger mctigger released this 05 Aug 06:38
· 153 commits to main since this release
1d7d2d2

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

  • unflatten Distribution Method: New functionality in the base distribution class for handling flattened parameters
  • Enhanced validate_args Support: 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 of Union)
  • 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 RuntimeError instances replaced with ValueError for 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:

  1. Update error handling code to catch ValueError instead of RuntimeError for parameter validation
  2. Review distribution initialization if using advanced parameter configurations

Compatibility

  • Python: 3.9+ (dropped Python 3.8 support)
  • PyTorch: 2.0+