be630106b22b7060ade1e842b704a83772704a24
Full translation of cub/block/block_scan.cuh BLOCK_SCAN_RAKING_MEMOIZE strategy to _torch_chunk_gated_delta_rule cross-chunk scan loop: CCCL RAKING_MEMOIZE: 'preserve upsweep segment values in registers while performing warp-synchronous scan, allowing downsweep not to re-read from shared memory.' Translation: precompute g_exp_full, g_last_exp, g_diff_exp tensors outside the sequential cross-chunk loop. Loop body now uses indexed lookups into precomputed tensors instead of calling exp() 3 times per chunk iteration. For seq_len=100K with chunk_size=64: 1562 chunks × 3 exp() = 4686 exp() calls eliminated from the hot loop. Replaced with 3 bulk exp() + tensor indexing. Memory tradeoff (same as RAKING_MEMOIZE's register pressure): +3 tensors of shape (batch, heads, num_chunks, chunk_size) float32 = 3 × 1 × 6 × 1562 × 64 × 4B ≈ 7MB (negligible vs 16GB model weights) Also inherits overflow_cast protection: g is clamped to [-20,20] before exp(), so precomputed values stay in safe float32 range.
project_6
Description
Languages
C++
41.8%
Cuda
31.6%
Python
22.2%
C
2.1%
CMake
1.1%
Other
1.1%