Skip to content

Destroy previous worker/scheduler streams on re-init - #9

Open
badnikhil wants to merge 1 commit into
ROCm:amd_mi350from
badnikhil:fix-stream-leak-on-reinit
Open

badnikhil wants to merge 1 commit into
ROCm:amd_mi350from
badnikhil:fix-stream-leak-on-reinit

Conversation

@badnikhil

Copy link
Copy Markdown

Problem

A second call to init_persistent_kernel() in the same process makes the next launch_persistent_kernel() hang forever. compile() already expects repeated inits for the online modes ("We will init for multiple times so the output directory should be permanent", python/mirage/mpk/persistent_kernel.py). Calling init_func again to change max_seq_length or reset event timing also hits the hang.

During the hang the worker kernel is running and the scheduler kernel never starts. The first launch prints 8 [SCHED_XCD] lines; the launch after the second init prints none. A runtime state dump taken during the same hang in a larger graph showed:

  • all 296 workers registered (worker_xcd_ready_count = 296)
  • scheduler queue heads at [1, 0, ...], which is only prepare_kernel's seed event
  • every worker queue head at 0
  • next_request_id still 0

The workers spin waiting for tasks that no scheduler dispatches.

Root cause

init_persistent_kernel() calls cudaStreamCreateWithFlags() (hipStreamCreateWithFlags) for worker_stream and scheduler_stream every time. It never destroys the pair from the previous init; only finalize_persistent_kernel() destroys streams.

HIP backs streams with a small pool of hardware queues (GPU_MAX_HW_QUEUES, default 4). Once the leaked streams use up the pool, the new scheduler stream can land on the same hardware queue as the new worker stream. Dispatches on one queue run serially. The scheduler kernel therefore waits behind the persistent worker kernel, and the worker kernel waits for the scheduler. This is the same serialization deadlock the split-mode comment in init_persistent_kernel() describes for rocprofv3 --pmc.

Raising the pool size with GPU_MAX_HW_QUEUES=8 makes the same sequence pass, which is consistent with this explanation.

Fix

  • In init_persistent_kernel(): if a stream pair already exists, destroy it before creating the new pair. global_runtime_config is a namespace-scope static, so the handles are null on the first init.
  • In finalize_persistent_kernel(): clear both handles after destroying them, so an init after a finalize (same launcher .so) does not destroy them twice.

The ownership model stays the same (init creates, finalize destroys). Launch behaviour on a single init is unchanged.

This change does not address the events and device buffers that init_persistent_kernel() also re-creates on each call. They leak memory but do not deadlock; I kept this PR to the hang.

Reproduction

The Qwen3 demo calls init once, so it cannot hit this. On MI300X its default path (USE_CK_FMHA=1) also needs gfx950-only builtins. The reproduction below is instead a standalone script using only upstream code:

timeout is needed because launch_func blocks in cudaStreamSynchronize while holding the GIL.

reinit_repro.py
"""Standalone reproducer: a second init_persistent_kernel in the same process deadlocks the next launch.

Graph: embedding -> rmsnorm -> argmax_partial -> argmax_reduce (Fleet's own ops, one decode iteration per launch).
Sequence: compile() (= init #1) -> launch 1 -> init_func (= init #2) -> launch 2.
With --no-reinit the second init is skipped (control: two launches on one init).
Each launch plants a different argmax in the embedding row it reads and checks output_tokens, so a passing
second launch really ran the kernels.

Usage (from an empty directory: compile() writes ./permanent_output_dir):
  MIRAGE_HOME=<repo> AMDGPU_TARGETS=gfx942 timeout 90 python -u reinit_repro.py [--no-reinit]
"""
import argparse
import time

import torch

import mirage as mi


def log(msg):
    print(f"[repro {time.strftime('%H:%M:%S')}] {msg}", flush=True)


ap = argparse.ArgumentParser()
ap.add_argument("--no-reinit", action="store_true", help="skip the second init_func (control)")
args = ap.parse_args()

HID, POS = 4096, 1024
MAX_SEQ = POS + 2  # prepare_next_batch stops at step + 2 >= max_seq_length: one iteration per launch
PLANTED = {0: 3000, 1: 1234}  # token -> argmax index planted in its embedding row
dev = "cuda"
i32, i64 = torch.int32, torch.int64
meta = dict(step=torch.zeros(1, dtype=i32, device=dev),
            tokens=torch.zeros(1, MAX_SEQ, dtype=i64, device=dev),
            input_tokens=torch.zeros(1, 1, dtype=i64, device=dev),
            output_tokens=torch.zeros(1, 1, dtype=i64, device=dev),
            num_new_tokens=torch.ones(1, dtype=i32, device=dev),
            prompt_lengths=torch.zeros(1, dtype=i32, device=dev),
            qo_indptr_buffer=torch.zeros(2, dtype=i32, device=dev),
            paged_kv_indptr_buffer=torch.zeros(2, dtype=i32, device=dev),
            paged_kv_indices_buffer=torch.zeros(1, dtype=i32, device=dev),
            paged_kv_last_page_len_buffer=torch.zeros(1, dtype=i32, device=dev))
workers, scheds = mi.get_configurations_from_gpu(0)
log(f"workers={workers} schedulers={scheds}")
mpk = mi.PersistentKernel(
    mode="offline", world_size=1, mpi_rank=0, num_workers=workers, num_local_schedulers=scheds,
    num_remote_schedulers=0, max_seq_length=MAX_SEQ, max_num_batched_requests=1, max_num_batched_tokens=1,
    max_num_pages=1, page_size=MAX_SEQ, eos_token_id=-1, meta_tensors=meta, profiler_tensor=None,
    trace_name="", spec_decode_config=None, use_cutlass_kernel=False)

emb = torch.randn(len(PLANTED), HID, device=dev).to(torch.bfloat16) * 0.01
for tok, idx in PLANTED.items():
    emb[tok, idx] = 1.0
x = mpk.new_tensor(dims=(1, HID), dtype=mi.bfloat16, name="x", io_category="cuda_tensor")
xn = mpk.new_tensor(dims=(1, HID), dtype=mi.bfloat16, name="xn", io_category="cuda_tensor")
pv = mpk.new_tensor(dims=(1, HID // 1024), dtype=mi.bfloat16, name="pv", io_category="cuda_tensor")
pi = mpk.new_tensor(dims=(1, HID // 1024), dtype=mi.int64, name="pi", io_category="cuda_tensor")
mpk.embed_layer(input=mpk.attach_input(torch_tensor=meta["input_tokens"], name="input_token"),
                weight=mpk.attach_input(torch_tensor=emb, name="embed_tokens"), output=x,
                grid_dim=(1, 1, 1), block_dim=(256, 1, 1), input_source=1)
mpk.rmsnorm_layer(input=x, weight=mpk.attach_input(torch_tensor=torch.ones(HID, dtype=torch.bfloat16, device=dev),
                  name="norm"), output=xn, grid_dim=(1, 1, 1), block_dim=(256, 1, 1))
mpk.argmax_partial_layer(input=xn, output=(pv, pi), grid_dim=(HID // 1024, 1, 1), block_dim=(256, 1, 1))
mpk.argmax_reduce_layer(input=(pv, pi), output=mpk.attach_input(torch_tensor=meta["output_tokens"], name="output_token"),
                        grid_dim=(1, 1, 1), block_dim=(256, 1, 1))
t0 = time.time()
mpk.compile()  # calls init_func (init #1)
log(f"compiled + init #1 in {time.time() - t0:.0f}s")


def launch(n, tok):
    mpk.init_request_func()
    meta["step"].fill_(POS)
    meta["prompt_lengths"].fill_(POS)
    meta["tokens"][0, 0] = tok  # the admission path embeds tokens[0]
    torch.cuda.synchronize()
    log(f"launch {n}: start")
    t = time.time()
    mpk()
    torch.cuda.synchronize()
    out = int(meta["output_tokens"][0, 0])
    log(f"launch {n}: returned in {(time.time() - t) * 1e3:.1f} ms, output_token={out} expected={PLANTED[tok]}")
    assert out == PLANTED[tok]


launch(1, 0)
if not args.no_reinit:
    ptrs = [meta[k].data_ptr() for k in ("step", "tokens", "input_tokens", "output_tokens", "num_new_tokens",
                                         "prompt_lengths", "qo_indptr_buffer", "paged_kv_indptr_buffer",
                                         "paged_kv_indices_buffer", "paged_kv_last_page_len_buffer")]
    mpk.init_func(ptrs, 0, 0, workers, scheds, 0, MAX_SEQ, 1, -1)
    log("init #2 done")
launch(2, 1)
log("PASS")

Commands (MI300X, so AMDGPU_TARGETS=gfx942; deps/ cloned at json v3.12.0, cutlass v4.7.1, composable_kernel ac18460782fadcd24aa321394c63dc284c90593b):

git clone https://github.com/ROCm/fleet-chiplet-megakernel && cd fleet-chiplet-megakernel
AMDGPU_TARGETS=gfx942 MAX_JOBS=8 pip install -e . -v --no-build-isolation --no-deps
mkdir repro && cd repro   # compile() writes ./permanent_output_dir
export MIRAGE_HOME=$PWD/.. AMDGPU_TARGETS=gfx942
timeout 90 python -u reinit_repro.py                         # before: hangs (exit 124)
GPU_MAX_HW_QUEUES=8 timeout 90 python -u reinit_repro.py     # before: PASS
timeout 90 python -u reinit_repro.py --no-reinit             # before: PASS (control)

Before, default environment (the 8 [SCHED_XCD] lines of launch 1 omitted; none appear after launch 2: start):

[01:48:15] [repro 01:48:15] launch 1: returned in 17.3 ms, output_token=3000 expected=3000
[01:48:15] [repro 01:48:15] init #2 done
[01:48:15] [repro 01:48:15] launch 2: start
[01:49:19] == exit code 124 (124 = killed by the 90 s timeout)

Before, GPU_MAX_HW_QUEUES=8:

[01:49:45] [repro 01:49:45] init #2 done
[01:49:45] [repro 01:49:45] launch 2: start
[01:49:45] [SCHED_XCD] sched_id=2 xcd=0 workers_on_xcd=37 block=2
...        (8 schedulers)
[01:49:45] [repro 01:49:45] launch 2: returned in 15.8 ms, output_token=1234 expected=1234
[01:49:45] [repro 01:49:45] PASS

After, default environment:

[01:51:20] [repro 01:51:20] launch 1: returned in 15.1 ms, output_token=3000 expected=3000
[01:51:20] [repro 01:51:20] init #2 done
[01:51:20] [repro 01:51:20] launch 2: start
[01:51:20] [SCHED_XCD] sched_id=1 xcd=6 workers_on_xcd=37 block=1
...        (8 schedulers)
[01:51:20] [repro 01:51:20] launch 2: returned in 1.4 ms, output_token=1234 expected=1234
[01:51:20] [repro 01:51:20] PASS

Environment:

  • AMD Instinct MI300X (gfx942), virtual function, 296 workers / 8 schedulers
  • ROCm 7.2.4, HIP 7.2.53211
  • Ubuntu, Python 3.12.3, torch 2.14.0+rocm7.2
  • upstream amd_mi350 at 51dce4f

Testing

build sequence env result
51dce4f init, launch, init, launch default hang, killed at 90 s
51dce4f init, launch, init, launch GPU_MAX_HW_QUEUES=8 PASS
51dce4f init, launch, launch default PASS
this PR init, launch, init, launch default PASS (launch 2: 1.4 ms, correct token)
this PR init, launch, launch default PASS (no regression on the single-init path)
  • persistent_kernel.cuh is compiled by hipcc at compile() time, so every run above removed permanent_output_dir first. No pip install rebuild is involved.
  • The changed lines are clean under clang-format 15 (git clang-format --diff).
  • Tested on MI300X only. I have no MI350 to run on, but the code path is not architecture specific.

init_persistent_kernel() creates a new worker/scheduler stream pair on
every call and never destroys the previous pair. HIP backs streams with a
small pool of hardware queues (GPU_MAX_HW_QUEUES, default 4). Once the
leaked streams use it up, the new scheduler stream can land on the same
hardware queue as the new worker stream. Dispatches on one queue run
serially, so the scheduler kernel waits behind the persistent worker
kernel, which waits for the scheduler: the next launch never returns.

Destroy the previous pair before creating the new one, and clear the
handles in finalize_persistent_kernel() so an init after a finalize does
not destroy them a second time.

Reproduced on MI300X (gfx942), ROCm 7.2.4 with init -> launch -> init ->
launch: hangs at the default queue count, passes with
GPU_MAX_HW_QUEUES=8, passes with this change.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant