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).
-----------------------------------------------