Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 53 additions & 3 deletions aiter/ops/flydsl/kernels/pa_decode_plan.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,20 @@
from .pa_decode_reduce import MAX_CONTEXT_PARTITIONS


def _work_granularity_partitions(context_length: int, query_length: int) -> int:
"""Suggested MTP partition floor: about 256 query/compute-tile pairs.

The occupancy target alone leaves a few very long CTAs in large MTP
batches. Round the work-based floor to a power of two for the reducer.
Single-query decode retains its existing scheduling policy.
"""
if query_length == 1:
return 1
tiles = (context_length + 255) // 256
count = max(1, (tiles * query_length + 255) // 256)
return min(MAX_CONTEXT_PARTITIONS, 1 << (count - 1).bit_length())


@triton.jit
def _plan_pa_decode(
lengths,
Expand Down Expand Up @@ -99,6 +113,11 @@ class PADecodePlan:
def capacity(self) -> int:
return self.work_info.shape[0]

@property
def num_partitions(self) -> torch.Tensor:
"""GPU view of per-request counts, updated in place by plan refresh."""
return self.reduce_info[:, 1]

def validate(self, batch_size: int, num_kv_heads: int, device: torch.device):
if self.num_kv_heads != num_kv_heads:
raise ValueError("plan KV head count does not match the cache")
Expand All @@ -123,24 +142,36 @@ def plan_pa_decode(
context_lengths: torch.Tensor,
num_kv_heads: int,
*,
max_partitions: int = MAX_CONTEXT_PARTITIONS,
max_partitions: int | None = None,
workgroup_budget: int | None = None,
sliding_window: int = 0,
total_context_length: int | None = None,
query_length: int = 1,
plan: PADecodePlan | None = None,
) -> PADecodePlan:
"""Build/update a plan on the current stream without GPU-to-CPU readback.

Allocate once outside graph capture, then pass ``plan=...`` to refresh the
same metadata in place. Include this refresh in end-to-end measurements.
same metadata in place. A refresh inherits the existing plan's partition
limit when ``max_partitions`` is omitted; new plans default to 256.
``plan.num_partitions`` exposes the actual per-request counts on the GPU.
Pass the plan to ``pa_decode`` without a separate partition-count argument.
Include this refresh in end-to-end measurements.
The budget counts CTAs over all KV heads, with fused query positions.
This is opt-in: uniform or short-context workloads may favor static splits.

A positive ``sliding_window`` counts visible tokens including the query's
own position; 0 and -1 disable it. Context lengths include the MTP tokens,
so the planned range is the union of ``query_length`` causal windows.
Pass the same window and (when enabled) query length to ``pa_decode`` and
when refreshing the plan. Dense plans remain independent of query length.
when refreshing the plan. Dense plan metadata remains independent of query length.

A host-known ``total_context_length`` (sum of the batch's KV lengths) and
``query_length`` can increase the default budget for long MTP workloads.
Use total work, not the longest request times batch size, so one long
request does not inflate an otherwise short batch's launch. This hint is
only for initial buffer sizing; GPU lengths still determine every task.
Without it, or with an explicit budget, retain the existing budget policy.
"""
if not isinstance(sliding_window, int):
raise TypeError("sliding_window must be an int")
Expand All @@ -151,6 +182,8 @@ def plan_pa_decode(
raise TypeError("query_length must be an int")
if query_length < 1:
raise ValueError("query_length must be positive")
if total_context_length is not None and total_context_length < 0:
raise ValueError("total_context_length must be non-negative")
if context_lengths.device.type != "cuda" or context_lengths.dtype != torch.int32:
raise ValueError("context_lengths must be a CUDA int32 tensor")
if context_lengths.ndim != 1 or not context_lengths.is_contiguous():
Expand All @@ -160,6 +193,10 @@ def plan_pa_decode(
raise ValueError("plan supports batches in [1, 4096]")
if num_kv_heads < 1:
raise ValueError("num_kv_heads must be positive")
if plan is not None and not isinstance(plan, PADecodePlan):
raise TypeError("plan must be a PADecodePlan")
if max_partitions is None:
max_partitions = MAX_CONTEXT_PARTITIONS if plan is None else plan.max_partitions
if not 1 <= max_partitions <= MAX_CONTEXT_PARTITIONS:
raise ValueError(f"max_partitions must be in [1, {MAX_CONTEXT_PARTITIONS}]")
dev = context_lengths.device
Expand All @@ -168,6 +205,17 @@ def plan_pa_decode(
workgroup_budget = (
2 * torch.cuda.get_device_properties(dev).multi_processor_count
)
if total_context_length is not None:
average_context = (total_context_length + batch - 1) // batch
if sliding_window > 0:
# The union may start inside a tile; allow that alignment tail.
average_context = min(
average_context, sliding_window + query_length - 1 + 255
)
work_floor = _work_granularity_partitions(average_context, query_length)
workgroup_budget = max(
workgroup_budget, batch * num_kv_heads * work_floor
)
if workgroup_budget < 1:
raise ValueError("workgroup_budget must be positive")
capacity = min(
Expand All @@ -185,6 +233,8 @@ def plan_pa_decode(
else:
if workgroup_budget is not None:
raise ValueError("workgroup_budget is fixed when reusing a plan")
if total_context_length is not None:
raise ValueError("total_context_length is only used when creating a plan")
if max_partitions != plan.max_partitions:
raise ValueError("max_partitions must match the reused plan")
if sliding_window != plan.sliding_window:
Expand Down
141 changes: 90 additions & 51 deletions aiter/ops/flydsl/kernels/pa_decode_reduce.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,7 @@ def compile_pa_decode_ps_reduce(
sink_dtype_str: str,
use_sinks: bool,
use_work_plan: bool = False,
compact_work_plan: bool = False,
):
"""Build the partitioned-softmax reduction used by ``pa_decode``.

Expand Down Expand Up @@ -111,7 +112,12 @@ def compile_pa_decode_ps_reduce(
# two or eight independent wave pairs. A pair covers the two 64-element
# halves of the output vector, while its y-coordinate selects a disjoint
# contiguous range of partitions.
use_parallel_lds = head_size == 128 and max_context_partition_num > warp_size
compact_work_plan = compact_work_plan and use_work_plan
use_parallel_lds = (
head_size == 128
and max_context_partition_num > warp_size
and not compact_work_plan
)
parallel_groups = 1
if use_parallel_lds:
# Eight groups win from NP=128 onward on gfx950; two avoid excessive
Expand Down Expand Up @@ -587,64 +593,97 @@ def _sink_exp(sink_value, safe_max):
global_exp_sum, one_f
)

acc = zero_f
for chunk_idx in fx.range_constexpr(partitions_per_lane):
chunk_base = chunk_idx * warp_size
chunk_size = min(warp_size, max_context_partition_num - chunk_base)
if fx.const_expr(use_work_plan):
# Initialize in the enclosing constexpr-loop scope so the
# FlyDSL dynamic-if rewriter never observes a stale value
# from a previous unrolled chunk.
weight_local_i32 = zero_f.bitcast(fx.Int32)
if fx.Int32(chunk_base) < c_part_num:
if fx.const_expr(compact_work_plan):
# The plan limit can be 256 even when most rows have only a
# few partitions. Keep the four lane-local weights in registers
# and iterate only the actual count, including zero-part rows.
weights = fx.Vector.from_elements(
[
(scaled_sums[i] / safe_global_exp_sum).bitcast(fx.Int32)
for i in fx.range_constexpr(partitions_per_lane)
],
dtype=fx.Int32,
)
for part_idx, state in fx.range(0, c_part_num, 1, init=[zero_f]):
part_idx = fx.Int32(part_idx)
weight_local_i32 = weights[part_idx // c_warp_size]
weight_i32 = fx.Int32(
fx.rocdl.ds_bpermute(
T.i32,
(part_idx & c_wave_mask) * c_four,
weight_local_i32,
)
)
weight = weight_i32.bitcast(fx.Float32)
logits_offset = (
logits_seq_offset
+ kv_head_idx * stride_logits_head
+ part_idx * stride_logits_part
+ eqgs_idx * stride_logits_group
+ tid
)
part_logits = fx.Float32(logits[logits_offset])
reduced = yield [state[0] + part_logits * weight]
acc = fx.Float32(reduced)
else:
acc = zero_f
for chunk_idx in fx.range_constexpr(partitions_per_lane):
chunk_base = chunk_idx * warp_size
chunk_size = min(warp_size, max_context_partition_num - chunk_base)
if fx.const_expr(use_work_plan):
# Initialize in the enclosing constexpr-loop scope so the
# FlyDSL dynamic-if rewriter never observes a stale value
# from a previous unrolled chunk.
weight_local_i32 = zero_f.bitcast(fx.Int32)
if fx.Int32(chunk_base) < c_part_num:
weight_local_i32 = (
scaled_sums[chunk_idx] / safe_global_exp_sum
).bitcast(fx.Int32)
for part_lane in fx.range_constexpr(chunk_size):
part_idx = chunk_base + part_lane
c_part_idx = fx.Int32(part_idx)
if c_part_idx < c_part_num:
weight_i32 = fx.Int32(
fx.rocdl.ds_bpermute(
T.i32,
fx.Int32(part_lane) * c_four,
weight_local_i32,
)
)
weight = weight_i32.bitcast(fx.Float32)
logits_offset = (
logits_seq_offset
+ kv_head_idx * stride_logits_head
+ c_part_idx * stride_logits_part
+ eqgs_idx * stride_logits_group
+ tid
)
part_logits = fx.Float32(logits[logits_offset])
acc = acc + part_logits * weight
else:
weight_local_i32 = (
scaled_sums[chunk_idx] / safe_global_exp_sum
).bitcast(fx.Int32)
for part_lane in fx.range_constexpr(chunk_size):
part_idx = chunk_base + part_lane
c_part_idx = fx.Int32(part_idx)
if c_part_idx < c_part_num:
weight_i32 = fx.Int32(
fx.rocdl.ds_bpermute(
T.i32,
fx.Int32(part_lane) * c_four,
weight_local_i32,
)
)
weight = weight_i32.bitcast(fx.Float32)
logits_offset = (
logits_seq_offset
+ kv_head_idx * stride_logits_head
+ c_part_idx * stride_logits_part
+ eqgs_idx * stride_logits_group
+ tid
weight_i32 = fx.Int32(
fx.rocdl.ds_bpermute(
T.i32,
fx.Int32(part_lane) * c_four,
weight_local_i32,
)
part_logits = fx.Float32(logits[logits_offset])
acc = acc + part_logits * weight
else:
weight_local_i32 = (
scaled_sums[chunk_idx] / safe_global_exp_sum
).bitcast(fx.Int32)
for part_lane in fx.range_constexpr(chunk_size):
part_idx = chunk_base + part_lane
c_part_idx = fx.Int32(part_idx)
weight_i32 = fx.Int32(
fx.rocdl.ds_bpermute(
T.i32,
fx.Int32(part_lane) * c_four,
weight_local_i32,
)
)
weight = weight_i32.bitcast(fx.Float32)
logits_offset = (
logits_seq_offset
+ kv_head_idx * stride_logits_head
+ c_part_idx * stride_logits_part
+ eqgs_idx * stride_logits_group
+ tid
)
part_logits = fx.Float32(logits[logits_offset])
acc = acc + part_logits * weight
weight = weight_i32.bitcast(fx.Float32)
logits_offset = (
logits_seq_offset
+ kv_head_idx * stride_logits_head
+ c_part_idx * stride_logits_part
+ eqgs_idx * stride_logits_group
+ tid
)
part_logits = fx.Float32(logits[logits_offset])
acc = acc + part_logits * weight

query_idx = eqgs_idx // c_qgs
if fx.const_expr(use_parallel_lds):
Expand Down
Loading