This release introduces significant performance optimizations, memory efficiency improvements, and enhanced system robustness: 🚀 Performance Optimizations: - Add three new fused CUDA kernels (rope_index, rot_pos_emb, get_window_index) for accelerated multimodal preprocessing - Implement FSDP2 support for distributed training with improved memory efficiency - Add Torch.compile integration for additional performance gains - Optimize memory usage: reduce peak allocation from 48GB to 24GB on 8-GPU setup 🔧 System Robustness: - Fix missing token position inputs in prediction pipeline - Add type-robust negation operations in RoPE CUDA kernels (half/bfloat16 support) - Fix dataset root parameter initialization in LeRobot data loader - Enhanced error handling and input validation across fusion operators 📚 Documentation & Usability: - Add comprehensive memory usage benchmarks and hardware recommendations - Update citation format with proper arXiv reference - Improve training configuration documentation with quick start guide - Add detailed API documentation for new fusion operators 🛠️ Technical Details: - Version bump to 1.0.1 - New CUDA kernels: rope_index.cu, rot_pos.cu, window_index.cu - FSDP2 state dict loading with distribute_tensor support - Enhanced multimodal RoPE with 3D position encoding - Window attention optimization for Vision Transformers Breaking Changes: None - all changes are backward compatible
29 lines
1016 B
Markdown
29 lines
1016 B
Markdown
# Fusion Operators (CSRC)
|
|
|
|
High-performance CUDA kernels for accelerating model training, with specialized support for multimodal and MoE architectures.
|
|
|
|
## Operators
|
|
|
|
### Asymmetric Dual Expert GEMM
|
|
- `asym_dual_gmm`: Simultaneous matrix multiplication for two experts
|
|
- Supports all transpose combinations (NN, TN, NT, TT)
|
|
|
|
### Token Permutation
|
|
- `permute`: Token permutation for MoE routing
|
|
- `unpermute`: Token recovery after expert computation
|
|
- `unpermute_bwd`: Backward pass for token recovery
|
|
|
|
### Multimodal RoPE
|
|
- `rope`: Rotary Position Embedding forward pass
|
|
- `rope_bwd`: RoPE backward pass
|
|
- `rope_index`: Generates position indices for multimodal RoPE
|
|
- `rot_pos_emb`: Fused rotary position embedding computation
|
|
|
|
### Vision Transformer Optimization
|
|
- `get_window_index`: Window attention index generation
|
|
|
|
|
|
## Acknowledgments
|
|
|
|
The `permute` and `unpermute` operators are adapted from [fanshiqing/grouped_gemm](https://github.com/fanshiqing/grouped_gemm). Thanks for their open-source contributions.
|