[Fix][AMD] MoRI EP: drop record_stream in TBO dispatch/combine (HSA out-of-resources) (#32746)
Co-authored-by: billishyahao <bill.he@amd.com> Co-authored-by: Duyi-Wang <duyi.wang@amd.com>
This commit is contained in:
co-authored by
billishyahao
Duyi-Wang
parent
4654b927eb
commit
66de161976
@@ -669,10 +669,15 @@ class _MoriEPDispatcherImplNormal(_MoriEPDispatcherImplBase):
|
|||||||
compute_stream = torch.cuda.current_stream()
|
compute_stream = torch.cuda.current_stream()
|
||||||
comm_stream = self._comm_stream # comm stream
|
comm_stream = self._comm_stream # comm stream
|
||||||
|
|
||||||
for t in (hidden_states, topk_weights, topk_ids):
|
# Deliberately no `record_stream(comm_stream)`: it would hold every
|
||||||
t.record_stream(comm_stream)
|
# dispatch/combine block in the allocator's deferred-free list until
|
||||||
if scale is not None:
|
# the comm event retires, and with `async_finish` the compute stream
|
||||||
scale.record_stream(comm_stream)
|
# never blocks, so `reserved` churns until ROCr runs out of scratch
|
||||||
|
# (HSA_STATUS_ERROR_OUT_OF_RESOURCES) -- the failure the non-EP DP
|
||||||
|
# TBO path already hit, see `_TBO_PERSIST_BUF` in dp_attention.py.
|
||||||
|
# Dropping it outright (rather than a keep-alive as in DeepseekV4
|
||||||
|
# op_combine_a) is safe because dispatch_b / combine_b wait_event on
|
||||||
|
# the compute stream in the same call, with no TBO yield between.
|
||||||
|
|
||||||
with torch.cuda.stream(comm_stream):
|
with torch.cuda.stream(comm_stream):
|
||||||
# if (previous_event) stream_wait(comm_stream, previous_event)
|
# if (previous_event) stream_wait(comm_stream, previous_event)
|
||||||
@@ -710,14 +715,6 @@ class _MoriEPDispatcherImplNormal(_MoriEPDispatcherImplBase):
|
|||||||
else:
|
else:
|
||||||
compute_stream.wait_stream(comm_stream)
|
compute_stream.wait_stream(comm_stream)
|
||||||
|
|
||||||
for t in (
|
|
||||||
packed_recv_hidden,
|
|
||||||
recv_topk_weights,
|
|
||||||
recv_scales,
|
|
||||||
recv_topk_ids,
|
|
||||||
):
|
|
||||||
if t is not None:
|
|
||||||
t.record_stream(comm_stream)
|
|
||||||
else:
|
else:
|
||||||
|
|
||||||
(
|
(
|
||||||
@@ -783,8 +780,7 @@ class _MoriEPDispatcherImplNormal(_MoriEPDispatcherImplBase):
|
|||||||
compute_stream = torch.cuda.current_stream()
|
compute_stream = torch.cuda.current_stream()
|
||||||
comm_stream = self._comm_stream
|
comm_stream = self._comm_stream
|
||||||
|
|
||||||
for t in (hidden_states, topk_ids, topk_weights):
|
# No `record_stream(comm_stream)` -- see `_dispatch_core`.
|
||||||
t.record_stream(comm_stream)
|
|
||||||
|
|
||||||
with torch.cuda.stream(comm_stream):
|
with torch.cuda.stream(comm_stream):
|
||||||
if previous_event is not None:
|
if previous_event is not None:
|
||||||
@@ -810,8 +806,6 @@ class _MoriEPDispatcherImplNormal(_MoriEPDispatcherImplBase):
|
|||||||
else:
|
else:
|
||||||
compute_stream.wait_stream(comm_stream)
|
compute_stream.wait_stream(comm_stream)
|
||||||
|
|
||||||
combined_hidden_states.record_stream(comm_stream)
|
|
||||||
|
|
||||||
else:
|
else:
|
||||||
combine_kwargs = self._combine_kwargs(hidden_states)
|
combine_kwargs = self._combine_kwargs(hidden_states)
|
||||||
combined_hidden_states = self.mori_op.combine(
|
combined_hidden_states = self.mori_op.combine(
|
||||||
|
|||||||
Reference in New Issue
Block a user