Skip to content
Merged
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
60 changes: 55 additions & 5 deletions src/flagsparse/sparse_operations/_spmv_csr_benchmark.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,61 @@
"""Shared CSR SpMV measurements; setup and diagnostics never change the denominator."""

import torch
import triton
import triton.language as tl

from . import _common as common
from .spmv_csr import flagsparse_spmv_csr_run, _spmv_device_context


def _prepared_row_tile_op(prepared, x, out):
"""Bind a prepared real row-tile launch outside the event window."""
from . import _spmv_csr_kernels as kernels

config = prepared.config["row_tile"]
grid = (triton.cdiv(prepared.n_rows, config["rows_per_program"]),)
acc = tl.float32 if prepared.data.dtype == torch.float32 else tl.float64
avg_nnz = prepared.data.numel() / prepared.n_rows
fixed_steps = 0
if (
prepared.data.dtype == torch.float32
and prepared.max_row_nnz <= 32
) or (
prepared.data.dtype == torch.float64
and avg_nnz > 7.0
and prepared.max_row_nnz <= 32
):
fixed_steps = triton.cdiv(
prepared.max_row_nnz, config["lanes_per_row"]
)

def launch():
kernels.row_tile_real_kernel[grid](
prepared.data,
prepared.kernel_indices,
prepared.kernel_indptr,
x,
out,
prepared.kernel_indptr,
prepared.n_rows,
INDEXED=False,
R=config["rows_per_program"],
V=config["lanes_per_row"],
STAGES=config["loop_num_stages"],
ACC=acc,
POSITION_64=not (
prepared.data.dtype == torch.float32
and prepared.kernel_indptr.dtype == torch.int32
),
FIXED_STEPS=fixed_steps,
num_warps=config["num_warps"],
enable_fp_fusion=(prepared.data.dtype == torch.float32),
)
return out

return launch


def _filtered_avg_ms(times):
if not times:
return None
Expand Down Expand Up @@ -65,17 +115,17 @@ def measure_route(prepared, x, warmup=10, iters=50, timing=False):
# to the GPU event while the stream is empty. Keep the generic route
# for paths that genuinely perform runtime processing.
direct_row_tile = (
common._is_rocm_runtime()
and
prepared.alg == "row_tile"
and not prepared.transpose
and not x.is_conj()
and not prepared.data.is_conj()
and x.is_contiguous()
and prepared.data.dtype in (torch.float32, torch.float64)
)
if direct_row_tile:
from . import _spmv_csr_kernels as kernels

timed_op = lambda: kernels.compute(
prepared, x, out, prepared.alg, prepared.config, plan=None
)
timed_op = _prepared_row_tile_op(prepared, x, out)
else:
timed_op = lambda: flagsparse_spmv_csr_run(prepared, x, out=out)
value, gpu_ms = event_benchmark(
Expand Down
116 changes: 107 additions & 9 deletions src/flagsparse/sparse_operations/_spmv_csr_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@
import triton
import triton.language as tl

from . import _common


@triton.jit
def _product(A, X, pos, col, mask, COMPLEX: tl.constexpr, ACC: tl.constexpr):
Expand Down Expand Up @@ -67,6 +69,53 @@ def row_tile_kernel(
_store_result(Y, row, tl.sum(acc, 1), tl.sum(imag, 1), valid, COMPLEX)


@triton.jit
def row_tile_real_kernel(
A,
CI,
RP,
X,
Y,
ROWS,
N,
INDEXED: tl.constexpr,
R: tl.constexpr,
V: tl.constexpr,
STAGES: tl.constexpr,
ACC: tl.constexpr,
POSITION_64: tl.constexpr,
FIXED_STEPS: tl.constexpr,
):
"""Real-only row tile without complex component bookkeeping."""
ridx = tl.program_id(0) * R + tl.arange(0, R)
if POSITION_64:
ridx = ridx.to(tl.int64)
valid = ridx < N
if INDEXED:
row = tl.load(ROWS + ridx, valid, 0)
else:
row = ridx
start = tl.load(RP + row, valid, 0)
end = tl.load(RP + row + 1, valid, 0)
if POSITION_64:
start = start.to(tl.int64)
end = end.to(tl.int64)
lane = tl.arange(0, V)
acc = tl.zeros((R, V), ACC)
if FIXED_STEPS > 0:
steps = FIXED_STEPS
else:
steps = tl.max(tl.cdiv(end - start, V), 0)
for step in tl.range(0, steps, num_stages=STAGES):
pos = start[:, None] + step * V + lane[None, :]
mask = valid[:, None] & (pos < end[:, None])
col = tl.load(CI + pos, mask, 0).to(tl.int64)
value = tl.load(A + pos, mask, 0).to(ACC)
vector = tl.load(X + col, mask, 0).to(ACC)
acc += value * vector
tl.store(Y + row, tl.sum(acc, 1), valid)


@triton.jit
def row_vector_kernel(
A,
Expand Down Expand Up @@ -401,19 +450,68 @@ def compute(prepared, x, y, alg, config, plan=None):
n = m if rows is None else rows.numel()
c = config["row_tile"]
if n:
row_tile_kernel[(triton.cdiv(n, c["rows_per_program"]),)](
grid = (triton.cdiv(n, c["rows_per_program"]),)
launch_args = (
*args,
prepared.kernel_indptr if rows is None else rows,
n,
INDEXED=rows is not None,
R=c["rows_per_program"],
V=c["lanes_per_row"],
STAGES=c["loop_num_stages"],
COMPLEX=complex_input,
ACC=acc,
num_warps=c["num_warps"],
enable_fp_fusion=False,
)
if not _common._is_rocm_runtime():
# Keep the pre-existing kernel for non-ROCm backends. The
# real-only specialization below is tuned for DCU wavefronts.
row_tile_kernel[grid](
*launch_args,
INDEXED=rows is not None,
R=c["rows_per_program"],
V=c["lanes_per_row"],
STAGES=c["loop_num_stages"],
COMPLEX=complex_input,
ACC=acc,
num_warps=c["num_warps"],
enable_fp_fusion=False,
)
elif complex_input:
row_tile_kernel[grid](
*launch_args,
INDEXED=rows is not None,
R=c["rows_per_program"],
V=c["lanes_per_row"],
STAGES=c["loop_num_stages"],
COMPLEX=True,
ACC=acc,
num_warps=c["num_warps"],
enable_fp_fusion=False,
)
else:
avg_nnz = prepared.data.numel() / m
fixed_steps = 0
if alg == "row_tile" and rows is None:
if (
prepared.data.dtype == torch.float32
and prepared.max_row_nnz <= 32
) or (
prepared.data.dtype == torch.float64
and avg_nnz > 7.0
and prepared.max_row_nnz <= 32
):
fixed_steps = triton.cdiv(
prepared.max_row_nnz, c["lanes_per_row"]
)
row_tile_real_kernel[grid](
*launch_args,
INDEXED=rows is not None,
R=c["rows_per_program"],
V=c["lanes_per_row"],
STAGES=c["loop_num_stages"],
ACC=acc,
POSITION_64=not (
prepared.data.dtype == torch.float32
and prepared.kernel_indptr.dtype == torch.int32
),
FIXED_STEPS=fixed_steps,
num_warps=c["num_warps"],
enable_fp_fusion=(prepared.data.dtype == torch.float32),
)
if alg in ("row_vector", "row_adaptive_split"):
rows = None if plan is None else plan["mid_rows"]
n = m if rows is None else rows.numel()
Expand Down
10 changes: 5 additions & 5 deletions src/flagsparse/sparse_operations/spmm_coo.py
Original file line number Diff line number Diff line change
Expand Up @@ -1243,7 +1243,7 @@ def _resolve_spmm_coo_launch_config(
# exactly the cost the sweep above measured away. The sweep constant holds
# for gfx936 too; an explicit ``block_nnz=`` argument still overrides it,
# since this whole branch only runs when the caller passed none.
block_nnz = 4
block_nnz = 4 if _is_rocm_runtime() else 256

# MetaX/MACA: the rowrun kernels unroll ``tl.static_range(0, BLOCK_NNZ)``, so
# BLOCK_NNZ multiplies the kernel's per-thread private memory. C550's driver caps
Expand Down Expand Up @@ -2366,7 +2366,7 @@ def _run_spmm_coo_canonical_route(
n_dense_cols,
output_dtype,
block_n=None,
block_nnz=256,
block_nnz=None,
out=None,
return_time=False,
route="rowrun",
Expand Down Expand Up @@ -2421,7 +2421,7 @@ def _run_spmm_coo_route(
B,
shape,
block_n=None,
block_nnz=256,
block_nnz=None,
out=None,
return_time=False,
return_meta=False,
Expand Down Expand Up @@ -2587,7 +2587,7 @@ def flagsparse_spmm_coo(
B,
shape,
block_n=None,
block_nnz=256,
block_nnz=None,
out=None,
return_time=False,
transpose=None,
Expand Down Expand Up @@ -2780,7 +2780,7 @@ def benchmark_spmm_coo_case(
warmup=20,
iters=200,
block_n=None,
block_nnz=256,
block_nnz=None,
run_cusparse=True,
route="rowrun",
compare_routes=False,
Expand Down
23 changes: 20 additions & 3 deletions src/flagsparse/sparse_operations/spmv_coo.py
Original file line number Diff line number Diff line change
Expand Up @@ -215,6 +215,7 @@ def _spmv_coo_seg_f32(
BLOCK_INNER: tl.constexpr,
SEG_IS_ROW: tl.constexpr,
HAS_BETA: tl.constexpr,
USE_MASKED_SELECT: tl.constexpr,
):
"""y[row] = alpha * sum(A[row] * x) + beta * y[row].

Expand Down Expand Up @@ -244,7 +245,9 @@ def _spmv_coo_seg_f32(
v = tl.load(data_ptr + offs, mask=m, other=0.0)
c = tl.load(col_ptr + offs, mask=m, other=0)
xv = tl.load(x_ptr + c, mask=m, other=0.0)
acc += tl.sum(tl.where(m, v * xv, 0.0))
# Masked loads already return zero, so the select only adds an extra
# instruction to the FP32 reduction.
acc += tl.sum(tl.where(m, v * xv, 0.0)) if USE_MASKED_SELECT else tl.sum(v * xv)
pos += BLOCK_INNER
out = alpha * acc
if HAS_BETA:
Expand All @@ -266,6 +269,7 @@ def _spmv_coo_seg_f64(
BLOCK_INNER: tl.constexpr,
SEG_IS_ROW: tl.constexpr,
HAS_BETA: tl.constexpr,
USE_MASKED_SELECT: tl.constexpr,
):
"""y[row] = alpha * sum(A[row] * x) + beta * y[row].

Expand Down Expand Up @@ -295,7 +299,9 @@ def _spmv_coo_seg_f64(
v = tl.load(data_ptr + offs, mask=m, other=0.0)
c = tl.load(col_ptr + offs, mask=m, other=0)
xv = tl.load(x_ptr + c, mask=m, other=0.0)
acc += tl.sum(tl.where(m, v * xv, 0.0))
# Masked loads return zero for inactive lanes, so no select is needed
# around the FP64 product before reduction.
acc += tl.sum(tl.where(m, v * xv, 0.0)) if USE_MASKED_SELECT else tl.sum(v * xv)
pos += BLOCK_INNER
out = alpha * acc
if HAS_BETA:
Expand Down Expand Up @@ -648,7 +654,17 @@ def _validate_x_coo(x, prepared):

def _triton_spmv_coo_kernel(prepared, x, block_size, num_warps, block_inner):
dtype = prepared.data.dtype
y = torch.zeros(prepared.n_rows, dtype=dtype, device=prepared.data.device)
# A run-compressed launch stores exactly one value for every non-empty row.
# When every row has a run, avoid a full output memset before the kernel;
# retain zero initialization for genuinely empty rows.
if (
_is_rocm_runtime()
and prepared.use_seg_kernel
and prepared.n_segs == prepared.n_rows
):
y = torch.empty(prepared.n_rows, dtype=dtype, device=prepared.data.device)
else:
y = torch.zeros(prepared.n_rows, dtype=dtype, device=prepared.data.device)
nnz = prepared.nnz
if nnz == 0:
return y
Expand Down Expand Up @@ -716,6 +732,7 @@ def _triton_spmv_coo_kernel(prepared, x, block_size, num_warps, block_inner):
BLOCK_INNER=block_inner,
SEG_IS_ROW=False,
HAS_BETA=False,
USE_MASKED_SELECT=not _is_rocm_runtime(),
num_warps=1,
)
return y
Expand Down
Loading
Loading