You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
fix(rocm): lock the device map in rocm::device() #2197
PR #2196 (#2183) locked two process-globals on the ROCm launch path: the JIT module cache (LOCAL_FIXES item 35) and HipEventPool in event.hip (item 36, whose unlocked version failed 20 of 20 eight-thread runs). Its review found the same pattern in rocm::device(): an unlocked find and try_emplace on a process-global std::unordered_map<int, Device>. The review also found that the event pool's map is a destructible function-static, so a HipEvent released during static teardown pushes into a destroyed map; the JIT cache in item 35 is leaked to avoid exactly this.
Current Behavior
All in src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/ on main (84d3a7b):
device.cpp:1483-1486: get_devices() returns a function-static std::unordered_map<int, Device>, never locked.
device.cpp:1557-1571: device(mlx::core::Device) does devices.find(index) and, on a miss, hipSetDevice, ensure_device_flags, devices.try_emplace(index, index).
device.cpp:1579-1584: get_command_encoder(Stream) calls device() first. It runs for every primitive in gpu::eval (eval.cpp:60), in finalize, synchronize and new_stream (eval.cpp:33-44,122-128), and rocm::device(s.device) is also called directly inside eval_gpu of qmm, SDPA, flash attention, conv, compiled and fp8 convert. So the map is read on every launch, on whichever thread evaluates: the server's scheduler, embedding, rerank and audio workers each evaluate on their own thread-local stream.
device.cpp:591-603: record_stream_error calls get_devices().find() from the event error path (event.hip:434), unlocked.
device.cpp:1586-1591: clear_all_encoders() (from gpu::clear_streams, eval.cpp:130-132) iterates the map unlocked, and Device::clear_encoders() (device.cpp:454-456) clears encoders_ without taking encoders_mtx_, which get_command_encoder and find_encoder do take.
event.hip: HipEventPool::cache_for holds static std::map<std::pair<int, int>, std::vector<HipEventHandle>> cache. After fix(rocm): lock the JIT module cache and the HIP event pool #2196 it is locked by a static std::mutex mu; both are still destroyed at exit, and the map's destructor runs hipEventDestroy on every pooled handle, which fails on a faulted device (item 7).
Reachability
Latent; not reachable on a single-GPU host or in any shipped code path today.
Every insert comes from a stream's creation: an eval needs a stream, and mlx::core::new_stream (patches-rocm/mlx/stream.cpp:67-78) holds the all_streams()std::unique_lock while it calls gpu::new_stream -> get_command_encoder -> device(). So two first inserts never run concurrently, on any host.
On a single-GPU host (this one: one gfx1151) the only key is 0, inserted by the process's first new_stream before any thread can hold a stream. Every later call is a find on a map that no longer changes.
The reachable race needs two GPUs: thread A creates the first stream on GPU 1 (the insert) while thread B evaluates on GPU 0 (an unlocked find walking the same buckets). That is a data race and undefined behavior, even though libstdc++ does not rehash at the second insert.
Make the device map safe in the same way upstream's is, keeping the fork's lazy, per-device construction.
Upstream at 81ba1c6a (mlx/backend/cuda/device.cpp:581-597) leaks a std::vector<Device> and fills it in a magic-static initializer that constructs every device, so lookups are read-only after one thread-safe init. Eager construction is rejected here: Device::Device calls make_current() (hipSetDevice), and the fork must never touch a GPU it was not asked for, because a context or queue on the other GPU of a multi-GPU host wedges the discrete GPU's queue over a TB5 link (eval.cpp:18-23, device.cppensure_device_flags).
In device.cpp:
Replace get_devices() with a file-local leaked map and lock, matching item 35: static auto* devices = new std::unordered_map<int, Device>; and static auto* mtx = new std::shared_mutex; in an anonymous-namespace accessor. Remove the forward declaration at device.cpp:591, so nothing reaches the map without the lock.
device(): take a std::shared_lock and find; on a hit, return it->second. On a miss, take a std::unique_lock, find again, and only if still missing run hipSetDevice(index), ensure_device_flags(index) and try_emplace(index, index). Return the reference after releasing the lock: unordered_map nodes are never erased, so references stay valid. If the Device constructor throws, nothing is inserted and the unwind releases the lock.
record_stream_error: under the shared lock, look up the Device* and release the lock before calling find_encoder. Keep the leaked fallback path unchanged.
clear_all_encoders: copy the Device* values under the shared lock, release it, then call clear_encoders() on each. Never call into a Device with the map lock held (except construction in device()), so an encoder destructor that reaches record_stream_error or device() cannot deadlock on a non-recursive lock.
Device::clear_encoders(): under encoders_mtx_, swap encoders_ into a local map, then let the local map destroy the encoders after the lock is released, so ~CommandEncoder never runs while encoders_mtx_ is held.
In event.hip, if #2196 has merged without it, leak the pool: static auto* cache = new std::map<...>; and static auto* mu = new std::mutex;, so release during static teardown is safe and no hipEventDestroy runs at exit. If #2196 has not merged when this is implemented, rebase onto it first; do not re-implement its lock.
Scope
In scope:patches-rocm/mlx/backend/rocm/device.cpp (the map, device(), record_stream_error, clear_all_encoders, Device::clear_encoders), event.hip (leak the pool), a new test tests/rocm_device_map_concurrency.rs, and a LOCAL_FIXES entry.
Out of scope: the unlocked lazy init in Device::get_rocblas_handle() (device.cpp:175-245), which sets rocblas_initialized_ = true before rocblas_create_handle; it needs its own issue. Eager device construction, per-key locking, and the CUDA overlay.
Constraints: the hit path is one uncontended shared_lock per call, on every primitive; no extra HIP calls. The miss path must still call hipSetDevice and ensure_device_flags before construction, and must touch only the requested index. Lock order: device map, then encoders_mtx_, never the reverse.
Edge cases: a negative or out-of-range index behaves as today (the Device constructor or HIP fails, nothing is inserted, the next call retries); concurrent first-use of one index from several threads builds one Device.
Error handling: unchanged. No new messages.
Acceptance Criteria
No unlocked access to the device map remains: grep -n "get_devices" src/lib/mlx-cpp/patches-rocm returns nothing, and every lookup goes through the shared or unique lock.
The device map and the event pool's map and mutex are heap-allocated and never destroyed.
tests/rocm_device_map_concurrency.rs exists, gated #![cfg(feature = "rocm")], skipping on other backends, with: streams_created_while_others_evaluate (8 threads, 16 rounds, behind a Barrier: each thread creates and installs a new thread-local stream with new_thread_local_generation_stream, evaluates arange_f32(0, 256, 1) * 2 + 1 and checks every element exactly, so device() lookups overlap other threads' stream creation and evaluation on GPU 0); and second_gpu_first_stream_while_first_gpu_evaluates (skips with a message unless gpu_device_count() >= 2: 8 threads keep evaluating on GPU 0 while one thread runs new_thread_local_stream_on_gpu(1), installs it and evaluates the same check, all exact). The module doc states that the insert race can only run on a multi-GPU host and what was measured where.
Both new tests, tests/rocm_jit_module_concurrency.rs and tests/rocm_gpu_faults.rs pass on gfx1151, and the process exits 0 with no hipEventDestroy error at teardown.
LOCAL_FIXES gets the next item number (37 once fix(rocm): lock the JIT module cache and the HIP event pool #2196 has merged), naming the functions, the reachability above, the upstream comparison and why eager construction was rejected, and ending: "Applies to the fork; kept in mlxcelverse under the 2026-10-06 fork policy, not proposed there (lablup/mlxcel#)." No upstream marking.
Verification
cargo test --features rocm --test rocm_device_map_concurrency
foriin$(seq 20);do cargo test --features rocm --test rocm_device_map_concurrency -q ||break;done
cargo test --features rocm --test rocm_jit_module_concurrency --test rocm_gpu_faults
cargo fmt --check
cargo clippy --features rocm --all-targets -- -W warnings
make verify-rocm
Pass: all 20 loops succeed; on a single-GPU host the second-GPU test prints its skip line. Record in the PR whether a multi-GPU host was available.
Part of #1801
Problem / Background
PR #2196 (#2183) locked two process-globals on the ROCm launch path: the JIT module cache (LOCAL_FIXES item 35) and
HipEventPoolinevent.hip(item 36, whose unlocked version failed 20 of 20 eight-thread runs). Its review found the same pattern inrocm::device(): an unlockedfindandtry_emplaceon a process-globalstd::unordered_map<int, Device>. The review also found that the event pool's map is a destructible function-static, so aHipEventreleased during static teardown pushes into a destroyed map; the JIT cache in item 35 is leaked to avoid exactly this.Current Behavior
All in
src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/on main (84d3a7b):device.cpp:1483-1486:get_devices()returns a function-staticstd::unordered_map<int, Device>, never locked.device.cpp:1557-1571:device(mlx::core::Device)doesdevices.find(index)and, on a miss,hipSetDevice,ensure_device_flags,devices.try_emplace(index, index).device.cpp:1579-1584:get_command_encoder(Stream)callsdevice()first. It runs for every primitive ingpu::eval(eval.cpp:60), infinalize,synchronizeandnew_stream(eval.cpp:33-44,122-128), androcm::device(s.device)is also called directly insideeval_gpuof qmm, SDPA, flash attention, conv, compiled and fp8 convert. So the map is read on every launch, on whichever thread evaluates: the server's scheduler, embedding, rerank and audio workers each evaluate on their own thread-local stream.device.cpp:591-603:record_stream_errorcallsget_devices().find()from the event error path (event.hip:434), unlocked.device.cpp:1586-1591:clear_all_encoders()(fromgpu::clear_streams,eval.cpp:130-132) iterates the map unlocked, andDevice::clear_encoders()(device.cpp:454-456) clearsencoders_without takingencoders_mtx_, whichget_command_encoderandfind_encoderdo take.event.hip:HipEventPool::cache_forholdsstatic std::map<std::pair<int, int>, std::vector<HipEventHandle>> cache. After fix(rocm): lock the JIT module cache and the HIP event pool #2196 it is locked by astatic std::mutex mu; both are still destroyed at exit, and the map's destructor runshipEventDestroyon every pooled handle, which fails on a faulted device (item 7).Reachability
Latent; not reachable on a single-GPU host or in any shipped code path today.
mlx::core::new_stream(patches-rocm/mlx/stream.cpp:67-78) holds theall_streams()std::unique_lockwhile it callsgpu::new_stream->get_command_encoder->device(). So two first inserts never run concurrently, on any host.new_streambefore any thread can hold a stream. Every later call is afindon a map that no longer changes.findwalking the same buckets). That is a data race and undefined behavior, even though libstdc++ does not rehash at the second insert.--main-gpuother than 0 is rejected (src/cli/ggml_compat_args.rs:633), andnew_stream_on_gpu,set_default_gpu_deviceandnew_thread_local_stream_on_gpu(src/lib/mlxcel-core/src/streams.rs:420-453) have no callers outside tests. It becomes reachable as soon as the multi-GPU path in feat: single-node multi-GPU tensor parallelism to load and serve models larger than one GPU (CUDA/DGX) #486 or feat: multi-GPU tensor-parallel runtime with per-rank device placement and cross-GPU collectives #488 places a stream on GPU 1.Proposed Solution
Make the device map safe in the same way upstream's is, keeping the fork's lazy, per-device construction.
Upstream at 81ba1c6a (
mlx/backend/cuda/device.cpp:581-597) leaks astd::vector<Device>and fills it in a magic-static initializer that constructs every device, so lookups are read-only after one thread-safe init. Eager construction is rejected here:Device::Devicecallsmake_current()(hipSetDevice), and the fork must never touch a GPU it was not asked for, because a context or queue on the other GPU of a multi-GPU host wedges the discrete GPU's queue over a TB5 link (eval.cpp:18-23,device.cppensure_device_flags).In
device.cpp:get_devices()with a file-local leaked map and lock, matching item 35:static auto* devices = new std::unordered_map<int, Device>;andstatic auto* mtx = new std::shared_mutex;in an anonymous-namespace accessor. Remove the forward declaration atdevice.cpp:591, so nothing reaches the map without the lock.device(): take astd::shared_lockandfind; on a hit, returnit->second. On a miss, take astd::unique_lock,findagain, and only if still missing runhipSetDevice(index),ensure_device_flags(index)andtry_emplace(index, index). Return the reference after releasing the lock:unordered_mapnodes are never erased, so references stay valid. If theDeviceconstructor throws, nothing is inserted and the unwind releases the lock.record_stream_error: under the shared lock, look up theDevice*and release the lock before callingfind_encoder. Keep the leakedfallbackpath unchanged.clear_all_encoders: copy theDevice*values under the shared lock, release it, then callclear_encoders()on each. Never call into aDevicewith the map lock held (except construction indevice()), so an encoder destructor that reachesrecord_stream_errorordevice()cannot deadlock on a non-recursive lock.Device::clear_encoders(): underencoders_mtx_, swapencoders_into a local map, then let the local map destroy the encoders after the lock is released, so~CommandEncodernever runs whileencoders_mtx_is held.In
event.hip, if #2196 has merged without it, leak the pool:static auto* cache = new std::map<...>;andstatic auto* mu = new std::mutex;, soreleaseduring static teardown is safe and nohipEventDestroyruns at exit. If #2196 has not merged when this is implemented, rebase onto it first; do not re-implement its lock.Scope
In scope:
patches-rocm/mlx/backend/rocm/device.cpp(the map,device(),record_stream_error,clear_all_encoders,Device::clear_encoders),event.hip(leak the pool), a new testtests/rocm_device_map_concurrency.rs, and a LOCAL_FIXES entry.Out of scope: the unlocked lazy init in
Device::get_rocblas_handle()(device.cpp:175-245), which setsrocblas_initialized_ = truebeforerocblas_create_handle; it needs its own issue. Eager device construction, per-key locking, and the CUDA overlay.Implementation Notes
std::shared_mutexand double-checked insert from item 35 (jit_module.cppin fix(rocm): lock the JIT module cache and the HIP event pool #2196); the barrier and thread-local stream harness fromtests/rocm_jit_module_concurrency.rs.shared_lockper call, on every primitive; no extra HIP calls. The miss path must still callhipSetDeviceandensure_device_flagsbefore construction, and must touch only the requested index. Lock order: device map, thenencoders_mtx_, never the reverse.Deviceconstructor or HIP fails, nothing is inserted, the next call retries); concurrent first-use of one index from several threads builds oneDevice.Acceptance Criteria
grep -n "get_devices" src/lib/mlx-cpp/patches-rocmreturns nothing, and every lookup goes through the shared or unique lock.tests/rocm_device_map_concurrency.rsexists, gated#![cfg(feature = "rocm")], skipping on other backends, with:streams_created_while_others_evaluate(8 threads, 16 rounds, behind aBarrier: each thread creates and installs a new thread-local stream withnew_thread_local_generation_stream, evaluatesarange_f32(0, 256, 1) * 2 + 1and checks every element exactly, sodevice()lookups overlap other threads' stream creation and evaluation on GPU 0); andsecond_gpu_first_stream_while_first_gpu_evaluates(skips with a message unlessgpu_device_count() >= 2: 8 threads keep evaluating on GPU 0 while one thread runsnew_thread_local_stream_on_gpu(1), installs it and evaluates the same check, all exact). The module doc states that the insert race can only run on a multi-GPU host and what was measured where.tests/rocm_jit_module_concurrency.rsandtests/rocm_gpu_faults.rspass on gfx1151, and the process exits 0 with nohipEventDestroyerror at teardown.Verification
Pass: all 20 loops succeed; on a single-GPU host the second-GPU test prints its skip line. Record in the PR whether a multi-GPU host was available.