Fix unnecessary gather/scatter on CPU for non-contiguous Mamba statepool (#31754)
This commit is contained in:
@@ -488,11 +488,16 @@ class GDNAttnBackend(MambaAttnBackendBase):
|
|||||||
# slot layout, so they silently drop the write to the strided envelope
|
# slot layout, so they silently drop the write to the strided envelope
|
||||||
# pool. Run them on contiguous per-sequence copies (identity-indexed) and
|
# pool. Run them on contiguous per-sequence copies (identity-indexed) and
|
||||||
# scatter the result back. No-op for the default contiguous pool.
|
# scatter the result back. No-op for the default contiguous pool.
|
||||||
|
# CPU kernels (causal_conv1d_fwd_cpu, chunk_gated_delta_rule_cpu) use
|
||||||
|
# proper indexed writes and handle non-contiguous pools directly via
|
||||||
|
# cache_indices, so the gather/scatter round-trip is unnecessary on CPU.
|
||||||
# TODO(ch-wan): drop these .contiguous() copies by making the prefill conv
|
# TODO(ch-wan): drop these .contiguous() copies by making the prefill conv
|
||||||
# and chunk_gated_delta_rule kernels honor the pool's real slot stride +
|
# and chunk_gated_delta_rule kernels honor the pool's real slot stride +
|
||||||
# int64 indexing, like packed_decode / causal_conv1d_update already do.
|
# int64 indexing, like packed_decode / causal_conv1d_update already do.
|
||||||
needs_state_gather = (not is_target_verify) and (
|
needs_state_gather = (
|
||||||
not conv_states.is_contiguous() or not ssm_states.is_contiguous()
|
(not is_target_verify)
|
||||||
|
and (not is_cpu())
|
||||||
|
and (not conv_states.is_contiguous() or not ssm_states.is_contiguous())
|
||||||
)
|
)
|
||||||
if needs_state_gather:
|
if needs_state_gather:
|
||||||
conv_states_contig = conv_states[cache_indices].contiguous()
|
conv_states_contig = conv_states[cache_indices].contiguous()
|
||||||
|
|||||||
Reference in New Issue
Block a user