Skip to content

fix(rocm): lock the device map in rocm::device() #2197

Description

@inureyes

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 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.
  • Nothing in production targets index 1: --main-gpu other than 0 is rejected (src/cli/ggml_compat_args.rs:633), and new_stream_on_gpu, set_default_gpu_device and new_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 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.cpp ensure_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.

Implementation Notes

  • Reuse: the leaked map plus std::shared_mutex and double-checked insert from item 35 (jit_module.cpp in fix(rocm): lock the JIT module cache and the HIP event pool #2196); the barrier and thread-local stream harness from tests/rocm_jit_module_concurrency.rs.
  • 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
for i in $(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.

Activity

  1. added
    type:bugBug fixes, error corrections, or issue resolutions
    area:coremlxcel-core: MLX FFI, primitives, KV cache, layers
    platform:linuxLinux (CUDA / packaging) specific
    on Oct 7, 2026
  2. added and removed on Oct 7, 2026
  3. added a commit that references this issue on Oct 7, 2026
    05368e2
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    area:coremlxcel-core: MLX FFI, primitives, KV cache, layersplatform:linuxLinux (CUDA / packaging) specificpriority:lowLow prioritystatus:doneCompletedtype:bugBug fixes, error corrections, or issue resolutions

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions