diff --git a/src/flagsparse/sparse_operations/_spmv_csr_benchmark.py b/src/flagsparse/sparse_operations/_spmv_csr_benchmark.py index f0ef777..21f4cd2 100644 --- a/src/flagsparse/sparse_operations/_spmv_csr_benchmark.py +++ b/src/flagsparse/sparse_operations/_spmv_csr_benchmark.py @@ -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 @@ -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( diff --git a/src/flagsparse/sparse_operations/_spmv_csr_kernels.py b/src/flagsparse/sparse_operations/_spmv_csr_kernels.py index 752cffd..21b0e0e 100644 --- a/src/flagsparse/sparse_operations/_spmv_csr_kernels.py +++ b/src/flagsparse/sparse_operations/_spmv_csr_kernels.py @@ -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): @@ -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, @@ -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() diff --git a/src/flagsparse/sparse_operations/spmm_coo.py b/src/flagsparse/sparse_operations/spmm_coo.py index 8dd2e9c..b41413b 100644 --- a/src/flagsparse/sparse_operations/spmm_coo.py +++ b/src/flagsparse/sparse_operations/spmm_coo.py @@ -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 @@ -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", @@ -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, @@ -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, @@ -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, diff --git a/src/flagsparse/sparse_operations/spmv_coo.py b/src/flagsparse/sparse_operations/spmv_coo.py index 4b5aea7..084af9e 100644 --- a/src/flagsparse/sparse_operations/spmv_coo.py +++ b/src/flagsparse/sparse_operations/spmv_coo.py @@ -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]. @@ -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: @@ -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]. @@ -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: @@ -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 @@ -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 diff --git a/src/flagsparse/sparse_operations/spmv_csr.py b/src/flagsparse/sparse_operations/spmv_csr.py index 30cb612..502e9fb 100644 --- a/src/flagsparse/sparse_operations/spmv_csr.py +++ b/src/flagsparse/sparse_operations/spmv_csr.py @@ -1500,40 +1500,76 @@ def _configure_spmv_route(prepared, alg, config=None): actual["row_tile"]["loop_num_stages"] = 2 source += ":stage2" avg_nnz_per_row = prepared.data.numel() / prepared.n_rows - if avg_nnz_per_row <= 4.0 and prepared.max_row_nnz <= 16: - # Uniform low-degree graphs benefit from twice as many rows per - # program and a four-lane reduction. Keep the tight max-row - # guard: heavy-tailed graphs with the same average regress here. - actual["row_tile"].update( - rows_per_program=128, - lanes_per_row=4, - loop_num_stages=( - 2 - if prepared.data.dtype in (torch.float32, torch.float64) - else 1 - ), - ) + maximum = prepared.max_row_nnz + if prepared.data.dtype == torch.float32: + if avg_nnz_per_row <= 4.0 and maximum <= 16: + actual["row_tile"].update( + rows_per_program=128, lanes_per_row=4, num_warps=2 + ) + source += ":uniform-short-128x4w2" + elif 8.0 <= avg_nnz_per_row <= 11.0 and maximum <= 24: + actual["row_tile"].update(rows_per_program=64, lanes_per_row=16) + source += ":regular-fp32-64x16" + elif avg_nnz_per_row >= 20.0: + actual["row_tile"].update( + rows_per_program=32, lanes_per_row=16, loop_num_stages=3 + ) + source += ":dense-fp32-32x16s3" + elif 15.0 <= avg_nnz_per_row <= 17.0 and maximum <= 36: + # Moderately long, tightly bounded rows benefit from fewer + # rows per program and a third load stage; this avoids the + # idle tail imposed by the 64x16 tile. + actual["row_tile"].update( + rows_per_program=32, lanes_per_row=8, loop_num_stages=3 + ) + source += ":compact-medium-fp32-32x8s3" + elif avg_nnz_per_row >= 11.0 and maximum <= 40: + actual["row_tile"].update(rows_per_program=64, lanes_per_row=16) + source += ":regular-medium-fp32-64x16" + elif avg_nnz_per_row >= 11.0: + actual["row_tile"].update( + rows_per_program=32, lanes_per_row=8, loop_num_stages=3 + ) + source += ":irregular-medium-fp32-32x8s3" + elif prepared.data.dtype == torch.float64: + if avg_nnz_per_row <= 4.0 and maximum <= 16: + actual["row_tile"].update(rows_per_program=64, lanes_per_row=4) + source += ":uniform-short-fp64-64x4" + elif avg_nnz_per_row <= 4.0: + actual["row_tile"].update(rows_per_program=32, lanes_per_row=4) + source += ":irregular-short-fp64-32x4" + elif avg_nnz_per_row <= 7.0 and maximum <= 16: + actual["row_tile"].update( + rows_per_program=32, lanes_per_row=4, loop_num_stages=3 + ) + source += ":compact-fp64-32x4s3" + elif avg_nnz_per_row < 9.0 and maximum <= 16: + actual["row_tile"]["loop_num_stages"] = 3 + source += ":regular-short-fp64-64x8s3" + elif avg_nnz_per_row >= 20.0: + actual["row_tile"].update( + rows_per_program=16, + lanes_per_row=8 if maximum <= 32 else 16, + loop_num_stages=3, + ) + source += ":dense-fp64-small-tile" + elif avg_nnz_per_row < 15.0 and maximum <= 40: + actual["row_tile"].update( + rows_per_program=16, lanes_per_row=8, loop_num_stages=3 + ) + source += ":regular-medium-fp64-16x8s3" + elif avg_nnz_per_row >= 12.0: + actual["row_tile"].update( + rows_per_program=32, lanes_per_row=8, loop_num_stages=3 + ) + source += ":irregular-medium-fp64-32x8s3" + elif avg_nnz_per_row >= 9.0: + actual["row_tile"].update(rows_per_program=32, lanes_per_row=16) + source += ":rowlen-fp64-32x16" + elif avg_nnz_per_row <= 4.0 and maximum <= 16: + actual["row_tile"].update(rows_per_program=128, lanes_per_row=4) source += ":uniform-short-128x4" - elif ( - prepared.data.dtype == torch.float64 - and 5.0 <= avg_nnz_per_row <= 7.0 - and prepared.max_row_nnz <= 16 - ): - source += ":compact-fp64-stage2" - elif ( - prepared.data.dtype == torch.float32 - and 8.0 <= avg_nnz_per_row <= 11.0 - and prepared.max_row_nnz <= 24 - ): - # Regular FP32 rows in this range fit in one 16-lane vector. - # Keeping 64 rows per program retains the graph-like occupancy - # that the 32x16 profile loses on the amazon/GL shapes. - actual["row_tile"].update(rows_per_program=64, lanes_per_row=16) - source += ":regular-fp32-64x16" elif avg_nnz_per_row >= 9.0: - # A 32x16 tile is faster on regular medium/long-row matrices. - # Rows below nine remain in the 64x8 layout: its extra row-level - # parallelism matters for the short, nearly uniform NACA shape. actual["row_tile"].update(rows_per_program=32, lanes_per_row=16) source += ":rowlen-32x16" prepared.alg_requested, prepared.alg = requested, resolved