jax
06fd7b81 - Fixed perf regression introduced by checkin "fix correctness, concurrency, multi-GPU and FFI safety in aiter mha"

Commit
159 days ago
Fixed perf regression introduced by checkin "fix correctness, concurrency, multi-GPU and FFI safety in aiter mha" Root cause ---------- regressed checkin replaced a thread_local WorkspacePool with per-call AsyncWorkspace instances using hipMallocAsync / hipFreeAsync. On ROCm the default hipMemPool release threshold is 0, and stream- ordered frees cannot be reused by the very next hipMallocAsync on the same stream until the stream catches up. The bwd handler allocates five workspaces per call (the largest, dq_acc, can reach ~2 GB), so every backward call paid full real allocations instead of pool reuse. Fix --- Replace AsyncWorkspace with a process-wide workspace cache in jaxlib/gpu/aiter_mha_bwd.cc, keyed by (device, stream, slot): - WsSlot enum tags the five buffers (dq_acc, dk_exp, dv_exp, dbias, dummy_rng). - std::unordered_map<WsKey, WsEntry> guarded by std::mutex holds {ptr, capacity} per (dev, stream, slot). - get_workspace() reuses the cached buffer when bytes <= cap and grows it via hipMalloc otherwise; the optional zero-init runs stream-ordered (hipMemsetAsync) outside the mutex. - All five backward call sites switched from AsyncWorkspace ws; ws.allocate(bytes, stream) to get_workspace(WsSlot::..., dev_idx, stream, bytes, zero, &ptr). -----------------------------------------------
Author
Parents
Loading