From d4b3f107b36a590f068c4599622007937727ff83 Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Thu, 2 Jul 2026 09:41:48 +0800 Subject: [PATCH 01/13] spmv_csc --- conf/operators.yaml | 17 + ops_support.py | 6 + pytest.ini | 1 + run_flagsparse_pytest.py | 11 + src/flagsparse/__init__.py | 6 + src/flagsparse/sparse_operations/__init__.py | 4 + src/flagsparse/sparse_operations/spmv_csc.py | 570 ++++++++++++++++ tests/pytest/test_spmv_csc_accuracy.py | 220 +++++++ tests/test_spmv_csc.py | 647 +++++++++++++++++++ tools/ci/run_gpu_benchmark.py | 11 + 10 files changed, 1493 insertions(+) create mode 100644 src/flagsparse/sparse_operations/spmv_csc.py create mode 100644 tests/pytest/test_spmv_csc_accuracy.py create mode 100644 tests/test_spmv_csc.py diff --git a/conf/operators.yaml b/conf/operators.yaml index a51c456..cef3241 100644 --- a/conf/operators.yaml +++ b/conf/operators.yaml @@ -63,6 +63,23 @@ ops: stages: - beta: "1.0" + - id: spmv_csc + description: | + Computes sparse matrix-vector multiplication for a CSC matrix. + for: + - flagsparse_spmv_csc + - prepare_spmv_csc + labels: + - flagsparse + - sparse + - csc + - triton + - public-api + kind: + - SparseLinearAlg + stages: + - beta: "1.0" + - id: spmv_coo_tocsr description: | Computes COO SpMV through a COO-to-CSR preparation path. diff --git a/ops_support.py b/ops_support.py index 7c91f0c..2fd01e3 100644 --- a/ops_support.py +++ b/ops_support.py @@ -228,6 +228,11 @@ def registry(modules: dict[str, SourceModule]) -> tuple[ApiSpec, ...]: if "spmv_coo" in modules else ("non", "trans", "conj") ) + spmv_csc_ops = ( + op_names(modules["spmv_csc"], "SPMV_CSC_OP_NAMES") + if "spmv_csc" in modules + else ("non", "trans", "conj") + ) spmm_values = ( normalize_dtype_values(modules["spmm_csr"].get("SUPPORTED_SPMM_VALUE_DTYPES")) if "spmm_csr" in modules @@ -264,6 +269,7 @@ def registry(modules: dict[str, SourceModule]) -> tuple[ApiSpec, ...]: ), ApiSpec("spmv", "flagsparse_spmv_csr", "spmv_csr", "CSR", "triton", value_const="SUPPORTED_SPMV_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmv_ops, notes="op supports non/trans/conj; conj on real dtypes is transpose-equivalent"), ApiSpec("spmv", "flagsparse_spmv_coo", "spmv_coo", "COO", "triton", value_const="SUPPORTED_SPMV_COO_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmv_coo_ops, notes="COO path stores canonical row/col tensors and supports non/trans/conj"), + ApiSpec("spmv", "flagsparse_spmv_csc", "spmv_csc", "CSC", "triton", value_const="SUPPORTED_SPMV_CSC_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmv_csc_ops, notes="native CSC path supports non/trans/conj without CSR/COO conversion"), ApiSpec("spmv", "flagsparse_spmv_coo_tocsr", "spmv_csr", "COO->CSR", "triton", value_const="SUPPORTED_SPMV_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=("non",), notes="COO input is converted to CSR before compute"), ApiSpec("spmm", "flagsparse_spmm_csr", "spmm_csr", "CSR", "triton", value_const="SUPPORTED_SPMM_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmm_ops, notes="op supports non/trans/conj; conj on real dtypes is transpose-equivalent"), ApiSpec("spmm", "flagsparse_spmm_csr_opt", "spmm_csr", "CSR", "triton_opt", values=("float32", "float64"), index_const="SUPPORTED_INDEX_DTYPES", ops=("non",), notes="bucketed opt path only supports float32/float64"), diff --git a/pytest.ini b/pytest.ini index 378a20b..9c4910e 100644 --- a/pytest.ini +++ b/pytest.ini @@ -11,6 +11,7 @@ markers = alpha_spmm_alg1: AlphaSparse ALG1 SpMM accuracy (tests/pytest; manual unified-run opt-in) spmv_csr: CSR SpMV accuracy (tests/pytest) spmv_coo: COO SpMV accuracy (tests/pytest) + spmv_csc: CSC SpMV accuracy (tests/pytest) spmv_coo_tocsr: COO-to-CSR SpMV accuracy (tests/pytest) spmm_csr: CSR SpMM accuracy (tests/pytest) spmm_csr_opt: optimized CSR SpMM accuracy (tests/pytest) diff --git a/run_flagsparse_pytest.py b/run_flagsparse_pytest.py index ebaebe6..47b1346 100644 --- a/run_flagsparse_pytest.py +++ b/run_flagsparse_pytest.py @@ -162,6 +162,16 @@ class OperatorTestConfig: "--iters", "{iters}", ), + "spmv_csc": ( + "tests/test_spmv_csc.py", + "{input}", + "--csv-csc", + "{csv}", + "--warmup", + "{warmup}", + "--iters", + "{iters}", + ), "spmv_coo_tocsr": ( "tests/test_spmv_coo.py", "{input}", @@ -304,6 +314,7 @@ class OperatorTestConfig: "scatter": OperatorTestConfig("scatter", PERFORMANCE_COMMANDS["scatter"]), "spmv_csr": OperatorTestConfig("spmv_csr", PERFORMANCE_COMMANDS["spmv_csr"]), "spmv_coo": OperatorTestConfig("spmv_coo", PERFORMANCE_COMMANDS["spmv_coo"]), + "spmv_csc": OperatorTestConfig("spmv_csc", PERFORMANCE_COMMANDS["spmv_csc"]), "spmv_coo_tocsr": OperatorTestConfig( "spmv_coo_tocsr", PERFORMANCE_COMMANDS["spmv_coo_tocsr"] ), diff --git a/src/flagsparse/__init__.py b/src/flagsparse/__init__.py index 5466aae..608f6cb 100644 --- a/src/flagsparse/__init__.py +++ b/src/flagsparse/__init__.py @@ -16,6 +16,7 @@ "comprehensive_gather_test", "comprehensive_scatter_test", "PreparedCoo", + "PreparedCscSpmv", "PreparedAlphaSpmmAlg1", "PreparedCsrSpmv", "PreparedCsrSpmmOpt", @@ -38,6 +39,7 @@ "prepare_spmm_csr_opt_alg2_preprocess", "prepare_spmv_coo", "prepare_spmv_coo_tocsr", + "prepare_spmv_csc", "flagsparse_spmv_csr", "flagsparse_alpha_spmm_alg1", "flagsparse_alpha_spmm_alg1_tle", @@ -51,6 +53,7 @@ "alpha_spmm_alg1_tle_opt2_unavailable_reason", "flagsparse_spmv_coo", "flagsparse_spmv_coo_tocsr", + "flagsparse_spmv_csc", "flagsparse_spmm_csr_opt", "flagsparse_spmm_csr_run", "flagsparse_spmm_csr_opt_alg1", @@ -136,6 +139,7 @@ "comprehensive_gather_test", "comprehensive_scatter_test", "PreparedCoo", + "PreparedCscSpmv", "PreparedAlphaSpmmAlg1", "PreparedCsrSpmv", "PreparedCsrSpmmOpt", @@ -158,6 +162,7 @@ "prepare_spmm_csr_opt_alg2_preprocess", "prepare_spmv_coo", "prepare_spmv_coo_tocsr", + "prepare_spmv_csc", "flagsparse_spmv_csr", "flagsparse_alpha_spmm_alg1", "flagsparse_alpha_spmm_alg1_tle", @@ -171,6 +176,7 @@ "alpha_spmm_alg1_tle_opt2_unavailable_reason", "flagsparse_spmv_coo", "flagsparse_spmv_coo_tocsr", + "flagsparse_spmv_csc", "flagsparse_spmm_csr_opt", "flagsparse_spmm_csr_run", "flagsparse_spmm_csr_opt_alg1", diff --git a/src/flagsparse/sparse_operations/__init__.py b/src/flagsparse/sparse_operations/__init__.py index d54af97..5c3ce9e 100644 --- a/src/flagsparse/sparse_operations/__init__.py +++ b/src/flagsparse/sparse_operations/__init__.py @@ -64,6 +64,7 @@ prepare_spmm_csr_opt_alg2_preprocess, ) from .spmv_coo import PreparedCoo, flagsparse_spmv_coo, prepare_spmv_coo +from .spmv_csc import PreparedCscSpmv, flagsparse_spmv_csc, prepare_spmv_csc from .spmv_csr import ( PreparedCsrSpmv, flagsparse_spmv_coo_tocsr, @@ -111,6 +112,7 @@ "PreparedCoo", "PreparedAlphaSpmmAlg1", "PreparedCsrSpmv", + "PreparedCscSpmv", "PreparedCsrSpmmOpt", "PreparedCsrSpmmRoute", "PreparedCsrSpmmOptAlg2", @@ -159,6 +161,7 @@ "flagsparse_spmm_csr_opt_alg2_preprocess", "flagsparse_spmv_coo", "flagsparse_spmv_coo_tocsr", + "flagsparse_spmv_csc", "flagsparse_spmv_csr", "flagsparse_spsm_coo", "flagsparse_spsm_csr", @@ -202,6 +205,7 @@ "prepare_spmm_csr_opt_alg2_preprocess", "prepare_spmv_coo", "prepare_spmv_coo_tocsr", + "prepare_spmv_csc", "prepare_spmv_csr", "resolve_spmm_csr_algorithm", "SPMM_CSR_ALGORITHMS", diff --git a/src/flagsparse/sparse_operations/spmv_csc.py b/src/flagsparse/sparse_operations/spmv_csc.py new file mode 100644 index 0000000..013d24d --- /dev/null +++ b/src/flagsparse/sparse_operations/spmv_csc.py @@ -0,0 +1,570 @@ +"""Native CSC SpMV kernels and public helpers.""" + +from ._common import * + +import triton +import triton.language as tl + + +SUPPORTED_SPMV_CSC_VALUE_DTYPES = ( + torch.float32, + torch.float64, + torch.complex64, + torch.complex128, +) + +SPMV_CSC_OP_NON = 0 +SPMV_CSC_OP_TRANS = 1 +SPMV_CSC_OP_CONJ_TRANS = 2 +SPMV_CSC_OP_NAMES = { + SPMV_CSC_OP_NON: "non", + SPMV_CSC_OP_TRANS: "trans", + SPMV_CSC_OP_CONJ_TRANS: "conj", +} +_SPMV_CSC_OP_NAME_TO_CODE = { + name: code for code, name in SPMV_CSC_OP_NAMES.items() +} + + +def _spmv_csc_dtype_error_message(): + return "CSC SpMV supports float32, float64, complex64, and complex128" + + +def _normalize_spmv_csc_op(op=None, transpose=False): + if op is None: + return SPMV_CSC_OP_TRANS if bool(transpose) else SPMV_CSC_OP_NON + if isinstance(op, str): + token = op.strip().lower() + if token not in _SPMV_CSC_OP_NAME_TO_CODE: + raise ValueError("op must be one of: 0=non, 1=trans, 2=conj") + return _SPMV_CSC_OP_NAME_TO_CODE[token] + try: + op_code = int(op) + except (TypeError, ValueError) as exc: + raise ValueError("op must be one of: 0=non, 1=trans, 2=conj") from exc + if op_code not in SPMV_CSC_OP_NAMES: + raise ValueError("op must be one of: 0=non, 1=trans, 2=conj") + return op_code + + +def _spmv_csc_op_to_name(op): + return SPMV_CSC_OP_NAMES[_normalize_spmv_csc_op(op)] + + +def _spmv_csc_op_transposes(op): + return _normalize_spmv_csc_op(op) in ( + SPMV_CSC_OP_TRANS, + SPMV_CSC_OP_CONJ_TRANS, + ) + + +def _normalize_spmv_csc_index_fallback_policy(index_fallback_policy): + policy = str(index_fallback_policy).lower() + if policy not in ("auto", "strict"): + raise ValueError("index_fallback_policy must be 'auto' or 'strict'") + return policy + + +class PreparedCscSpmv: + """Prepared CSC metadata for repeated SpMV calls.""" + + __slots__ = ( + "data", + "kernel_indices", + "kernel_indptr", + "shape", + "n_rows", + "n_cols", + "nnz", + "block_nnz", + "max_segments", + "col_lengths", + "max_col_nnz", + "op", + "transpose", + "index_fallback_policy", + "index_fallback_applied", + "index_fallback_reason", + ) + + def __init__( + self, + data, + kernel_indices, + kernel_indptr, + shape, + n_rows, + n_cols, + block_nnz, + max_segments, + max_col_nnz, + col_lengths=None, + op=None, + transpose=False, + index_fallback_policy="auto", + index_fallback_applied=False, + index_fallback_reason=None, + ): + self.data = data + self.kernel_indices = kernel_indices + self.kernel_indptr = kernel_indptr + self.shape = (int(shape[0]), int(shape[1])) + self.n_rows = int(n_rows) + self.n_cols = int(n_cols) + self.nnz = int(data.numel()) + self.block_nnz = int(block_nnz) + self.max_segments = int(max_segments) + if col_lengths is None: + col_lengths = kernel_indptr[1:] - kernel_indptr[:-1] + self.col_lengths = col_lengths + self.max_col_nnz = int(max_col_nnz) + self.op = _normalize_spmv_csc_op(op, transpose=transpose) + self.transpose = _spmv_csc_op_transposes(self.op) + self.index_fallback_policy = str(index_fallback_policy).lower() + self.index_fallback_applied = bool(index_fallback_applied) + self.index_fallback_reason = index_fallback_reason + + +@triton.jit +def _spmv_csc_non_real_kernel( + data_ptr, + indices_ptr, + indptr_ptr, + x_ptr, + y_ptr, + n_cols, + BLOCK_NNZ: tl.constexpr, +): + col = tl.program_id(0) + seg = tl.program_id(1) + if col >= n_cols: + return + start = tl.load(indptr_ptr + col) + end = tl.load(indptr_ptr + col + 1) + offs = start + seg * BLOCK_NNZ + tl.arange(0, BLOCK_NNZ) + mask = offs < end + rows = tl.load(indices_ptr + offs, mask=mask, other=0) + vals = tl.load(data_ptr + offs, mask=mask, other=0.0) + x_val = tl.load(x_ptr + col) + tl.atomic_add(y_ptr + rows, vals * x_val, mask=mask, sem="relaxed") + + +@triton.jit +def _spmv_csc_non_complex_kernel( + data_ri_ptr, + indices_ptr, + indptr_ptr, + x_ri_ptr, + y_ri_ptr, + n_cols, + BLOCK_NNZ: tl.constexpr, +): + col = tl.program_id(0) + seg = tl.program_id(1) + if col >= n_cols: + return + start = tl.load(indptr_ptr + col) + end = tl.load(indptr_ptr + col + 1) + offs = start + seg * BLOCK_NNZ + tl.arange(0, BLOCK_NNZ) + mask = offs < end + rows = tl.load(indices_ptr + offs, mask=mask, other=0) + a_re = tl.load(data_ri_ptr + offs * 2, mask=mask, other=0.0) + a_im = tl.load(data_ri_ptr + offs * 2 + 1, mask=mask, other=0.0) + x_re = tl.load(x_ri_ptr + col * 2) + x_im = tl.load(x_ri_ptr + col * 2 + 1) + prod_re = a_re * x_re - a_im * x_im + prod_im = a_re * x_im + a_im * x_re + tl.atomic_add(y_ri_ptr + rows * 2, prod_re, mask=mask, sem="relaxed") + tl.atomic_add(y_ri_ptr + rows * 2 + 1, prod_im, mask=mask, sem="relaxed") + + +@triton.jit +def _spmv_csc_trans_real_kernel( + data_ptr, + indices_ptr, + indptr_ptr, + x_ptr, + y_ptr, + n_cols, + BLOCK_NNZ: tl.constexpr, + MAX_SEGMENTS: tl.constexpr, +): + col = tl.program_id(0) + if col >= n_cols: + return + start = tl.load(indptr_ptr + col) + end = tl.load(indptr_ptr + col + 1) + acc = tl.load(data_ptr + start, mask=start < end, other=0.0) * 0 + for seg in range(MAX_SEGMENTS): + offs = start + seg * BLOCK_NNZ + tl.arange(0, BLOCK_NNZ) + mask = offs < end + rows = tl.load(indices_ptr + offs, mask=mask, other=0) + vals = tl.load(data_ptr + offs, mask=mask, other=0.0) + x_vals = tl.load(x_ptr + rows, mask=mask, other=0.0) + acc = acc + tl.sum(tl.where(mask, vals * x_vals, 0.0)) + tl.store(y_ptr + col, acc) + + +@triton.jit +def _spmv_csc_trans_complex_kernel( + data_ri_ptr, + indices_ptr, + indptr_ptr, + x_ri_ptr, + y_ri_ptr, + n_cols, + BLOCK_NNZ: tl.constexpr, + MAX_SEGMENTS: tl.constexpr, + CONJ: tl.constexpr, +): + col = tl.program_id(0) + if col >= n_cols: + return + start = tl.load(indptr_ptr + col) + end = tl.load(indptr_ptr + col + 1) + acc_re = tl.load(data_ri_ptr + start * 2, mask=start < end, other=0.0) * 0 + acc_im = tl.load(data_ri_ptr + start * 2 + 1, mask=start < end, other=0.0) * 0 + for seg in range(MAX_SEGMENTS): + offs = start + seg * BLOCK_NNZ + tl.arange(0, BLOCK_NNZ) + mask = offs < end + rows = tl.load(indices_ptr + offs, mask=mask, other=0) + a_re = tl.load(data_ri_ptr + offs * 2, mask=mask, other=0.0) + a_im_raw = tl.load(data_ri_ptr + offs * 2 + 1, mask=mask, other=0.0) + if CONJ: + a_im = -a_im_raw + else: + a_im = a_im_raw + x_re = tl.load(x_ri_ptr + rows * 2, mask=mask, other=0.0) + x_im = tl.load(x_ri_ptr + rows * 2 + 1, mask=mask, other=0.0) + prod_re = a_re * x_re - a_im * x_im + prod_im = a_re * x_im + a_im * x_re + acc_re = acc_re + tl.sum(tl.where(mask, prod_re, 0.0)) + acc_im = acc_im + tl.sum(tl.where(mask, prod_im, 0.0)) + tl.store(y_ri_ptr + col * 2, acc_re) + tl.store(y_ri_ptr + col * 2 + 1, acc_im) + + +def _prepare_spmv_csc_matrix(data, indices, indptr, shape): + if not all(torch.is_tensor(t) for t in (data, indices, indptr)): + raise TypeError("data, indices, indptr must all be torch.Tensor") + if data.ndim != 1 or indices.ndim != 1 or indptr.ndim != 1: + raise ValueError("data, indices, indptr must be 1D tensors") + n_rows, n_cols = int(shape[0]), int(shape[1]) + if indptr.numel() != n_cols + 1: + raise ValueError( + f"indptr length must be n_cols+1={n_cols + 1}, got {indptr.numel()}" + ) + if data.numel() != indices.numel(): + raise ValueError("data and indices must have the same length (nnz)") + if not all(t.is_cuda for t in (data, indices, indptr)): + raise ValueError("data, indices, indptr must be CUDA tensors") + if not all(t.device == data.device for t in (indices, indptr)): + raise ValueError("data, indices, indptr must be on the same CUDA device") + if data.dtype not in SUPPORTED_SPMV_CSC_VALUE_DTYPES: + raise TypeError(_spmv_csc_dtype_error_message()) + if indices.dtype not in SUPPORTED_INDEX_DTYPES: + raise TypeError("indices dtype must be torch.int32 or torch.int64") + if indptr.dtype not in SUPPORTED_INDEX_DTYPES: + raise TypeError("indptr dtype must be torch.int32 or torch.int64") + data = data.contiguous() + indices = indices.contiguous() + indptr = indptr.contiguous() + if indptr.numel() > 0: + if int(indptr[0].item()) != 0: + raise ValueError("indptr must start at zero") + if int(indptr[-1].item()) != data.numel(): + raise ValueError("indptr[-1] must equal nnz") + if indptr.numel() > 1 and torch.any(indptr[1:] < indptr[:-1]).item(): + raise ValueError("indptr must be non-decreasing") + if data.numel() > 0: + min_index = int(indices.min().item()) + max_index = int(indices.max().item()) + if min_index < 0 or max_index >= n_rows: + raise IndexError("indices out of range for n_rows") + col_lengths = indptr[1:] - indptr[:-1] + max_col_nnz = int(col_lengths.max().item()) if n_cols > 0 else 0 + return data, indices, indptr, n_rows, n_cols, col_lengths, max_col_nnz + + +def prepare_spmv_csc( + data, + indices, + indptr, + shape, + block_nnz=256, + max_segments=None, + transpose=False, + op=None, + index_fallback_policy="auto", +): + index_fallback_policy = _normalize_spmv_csc_index_fallback_policy( + index_fallback_policy + ) + op_code = _normalize_spmv_csc_op(op, transpose=transpose) + if op is not None and bool(transpose) and op_code == SPMV_CSC_OP_NON: + raise ValueError("transpose=True conflicts with op=non") + data, indices, indptr, n_rows, n_cols, col_lengths, max_col_nnz = ( + _prepare_spmv_csc_matrix(data, indices, indptr, shape) + ) + block_nnz_use = int(block_nnz) + if block_nnz_use <= 0: + raise ValueError("block_nnz must be positive") + if max_segments is None: + max_segments_use = max((max_col_nnz + block_nnz_use - 1) // block_nnz_use, 1) + while max_segments_use > 2048 and block_nnz_use < 65536: + block_nnz_use *= 2 + max_segments_use = max( + (max_col_nnz + block_nnz_use - 1) // block_nnz_use, + 1, + ) + else: + max_segments_use = max(1, int(max_segments)) + return PreparedCscSpmv( + data=data, + kernel_indices=indices, + kernel_indptr=indptr, + shape=shape, + n_rows=n_rows, + n_cols=n_cols, + block_nnz=block_nnz_use, + max_segments=max_segments_use, + max_col_nnz=max_col_nnz, + col_lengths=col_lengths, + op=op_code, + index_fallback_policy=index_fallback_policy, + ) + + +def _validate_spmv_csc_x(x, prepared, op_code): + if x is None or not torch.is_tensor(x): + raise TypeError("x must be a torch.Tensor") + if x.ndim != 1: + raise ValueError("x must be a 1D tensor") + if not x.is_cuda: + raise ValueError("x must be a CUDA tensor") + if x.dtype != prepared.data.dtype: + raise TypeError("x dtype must match sparse matrix dtype") + expected = prepared.n_rows if _spmv_csc_op_transposes(op_code) else prepared.n_cols + if x.numel() != expected: + raise ValueError(f"x length must be {expected}, got {x.numel()}") + if x.device != prepared.data.device: + raise ValueError("x must be on the same device as sparse matrix data") + return x.contiguous() + + +def _triton_spmv_csc_kernel(prepared, x, op_code): + dtype = prepared.data.dtype + trans = _spmv_csc_op_transposes(op_code) + out_len = prepared.n_cols if trans else prepared.n_rows + y = torch.zeros(out_len, dtype=dtype, device=prepared.data.device) + if prepared.nnz == 0: + return y + if not trans: + grid = (prepared.n_cols, prepared.max_segments) + if _is_complex_dtype(dtype): + data_ri = torch.view_as_real(prepared.data).reshape(-1) + x_ri = torch.view_as_real(x).reshape(-1) + y_ri = torch.zeros(out_len * 2, dtype=data_ri.dtype, device=y.device) + _spmv_csc_non_complex_kernel[grid]( + data_ri, + prepared.kernel_indices, + prepared.kernel_indptr, + x_ri, + y_ri, + prepared.n_cols, + BLOCK_NNZ=prepared.block_nnz, + ) + y.copy_(torch.view_as_complex(y_ri.reshape(out_len, 2))) + return y + _spmv_csc_non_real_kernel[grid]( + prepared.data, + prepared.kernel_indices, + prepared.kernel_indptr, + x, + y, + prepared.n_cols, + BLOCK_NNZ=prepared.block_nnz, + ) + return y + grid = (prepared.n_cols,) + if _is_complex_dtype(dtype): + data_ri = torch.view_as_real(prepared.data).reshape(-1) + x_ri = torch.view_as_real(x).reshape(-1) + y_ri = torch.empty(out_len * 2, dtype=data_ri.dtype, device=y.device) + _spmv_csc_trans_complex_kernel[grid]( + data_ri, + prepared.kernel_indices, + prepared.kernel_indptr, + x_ri, + y_ri, + prepared.n_cols, + BLOCK_NNZ=prepared.block_nnz, + MAX_SEGMENTS=prepared.max_segments, + CONJ=(op_code == SPMV_CSC_OP_CONJ_TRANS), + ) + y.copy_(torch.view_as_complex(y_ri.reshape(out_len, 2))) + return y + _spmv_csc_trans_real_kernel[grid]( + prepared.data, + prepared.kernel_indices, + prepared.kernel_indptr, + x, + y, + prepared.n_cols, + BLOCK_NNZ=prepared.block_nnz, + MAX_SEGMENTS=prepared.max_segments, + ) + return y + + +def _spmv_csc_uses_int64_indices(prepared): + return ( + prepared.kernel_indices.dtype == torch.int64 + or prepared.kernel_indptr.dtype == torch.int64 + ) + + +def _spmv_csc_int32_fallback_blocker(prepared): + if prepared.nnz > _INDEX_LIMIT_INT32: + return f"nnz {prepared.nnz} cannot fit int32" + if prepared.kernel_indices.numel() > 0: + max_row = int(prepared.kernel_indices.max().item()) + if max_row > _INDEX_LIMIT_INT32: + return f"row index {max_row} cannot fit int32" + if prepared.kernel_indptr.numel() > 0: + max_ptr = int(prepared.kernel_indptr[-1].item()) + if max_ptr > _INDEX_LIMIT_INT32: + return f"indptr offset {max_ptr} cannot fit int32" + return None + + +def _spmv_csc_prepared_with_int32_indices(prepared, reason): + blocker = _spmv_csc_int32_fallback_blocker(prepared) + if blocker is not None: + raise RuntimeError(f"int32 fallback is unsafe: {blocker}") from reason + return PreparedCscSpmv( + data=prepared.data, + kernel_indices=prepared.kernel_indices.to(torch.int32).contiguous(), + kernel_indptr=prepared.kernel_indptr.to(torch.int32).contiguous(), + shape=prepared.shape, + n_rows=prepared.n_rows, + n_cols=prepared.n_cols, + block_nnz=prepared.block_nnz, + max_segments=prepared.max_segments, + max_col_nnz=prepared.max_col_nnz, + col_lengths=prepared.col_lengths, + op=prepared.op, + index_fallback_policy=prepared.index_fallback_policy, + index_fallback_applied=True, + index_fallback_reason=str(reason), + ) + + +def _run_spmv_csc_prepared_with_fallback(prepared, x, op_code): + try: + return _triton_spmv_csc_kernel(prepared, x, op_code) + except RuntimeError as exc: + if ( + prepared.index_fallback_policy != "auto" + or not _spmv_csc_uses_int64_indices(prepared) + ): + raise + fallback_prepared = _spmv_csc_prepared_with_int32_indices(prepared, exc) + return _triton_spmv_csc_kernel(fallback_prepared, x, op_code) + + +def flagsparse_spmv_csc( + data=None, + indices=None, + indptr=None, + x=None, + shape=None, + block_nnz=256, + max_segments=None, + out=None, + return_time=False, + return_meta=False, + prepared=None, + transpose=None, + op=None, + index_fallback_policy="auto", +): + """CSC SpMV using native Triton CSC kernels.""" + op_explicit = op is not None + op_code = _normalize_spmv_csc_op( + op, + transpose=False if transpose is None else bool(transpose), + ) + if ( + op_explicit + and transpose is not None + and bool(transpose) != _spmv_csc_op_transposes(op_code) + ): + raise ValueError("transpose conflicts with op") + if prepared is None: + if any(arg is None for arg in (data, indices, indptr, shape)): + raise ValueError( + "data, indices, indptr, and shape are required when prepared is not provided" + ) + prepared = prepare_spmv_csc( + data, + indices, + indptr, + shape, + block_nnz=block_nnz, + max_segments=max_segments, + op=op_code, + index_fallback_policy=index_fallback_policy, + ) + else: + if op_explicit and op_code != prepared.op: + raise ValueError( + f"op={_spmv_csc_op_to_name(op_code)} does not match prepared.op={_spmv_csc_op_to_name(prepared.op)}" + ) + if ( + not op_explicit + and transpose is not None + and bool(transpose) != prepared.transpose + ): + raise ValueError( + f"transpose={bool(transpose)} does not match prepared.transpose={prepared.transpose}" + ) + if not op_explicit: + op_code = prepared.op + x = _validate_spmv_csc_x(x, prepared, op_code) + do_timing = bool(return_time or return_meta) + if do_timing: + torch.cuda.synchronize() + t0 = time.perf_counter() + y = _run_spmv_csc_prepared_with_fallback(prepared, x, op_code) + if do_timing: + torch.cuda.synchronize() + compute_ms = (time.perf_counter() - t0) * 1000.0 + op_total_ms = compute_ms + else: + compute_ms = None + op_total_ms = None + if out is not None: + if not out.is_cuda: + raise ValueError("out must be a CUDA tensor") + if out.device != y.device: + raise ValueError("out must be on the same CUDA device as the result") + if out.shape != y.shape or out.dtype != y.dtype: + raise ValueError("out shape/dtype must match result") + out.copy_(y) + y = out + if return_meta: + meta = { + "op": _spmv_csc_op_to_name(op_code), + "symbolic_ms": 0.0 if do_timing else None, + "compute_ms": compute_ms, + "op_total_ms": op_total_ms, + "index_fallback_applied": prepared.index_fallback_applied, + "index_fallback_reason": prepared.index_fallback_reason, + } + if return_time: + return y, op_total_ms, meta + return y, meta + if return_time: + return y, op_total_ms + return y diff --git a/tests/pytest/test_spmv_csc_accuracy.py b/tests/pytest/test_spmv_csc_accuracy.py new file mode 100644 index 0000000..35ca609 --- /dev/null +++ b/tests/pytest/test_spmv_csc_accuracy.py @@ -0,0 +1,220 @@ +import importlib + +import pytest +import torch + +from flagsparse import flagsparse_spmv_csc, prepare_spmv_csc +from tests.pytest.accuracy_utils import close_tolerances +from tests.pytest.param_shapes import SPMV_MN_SHAPES + + +spmv_csc_mod = importlib.import_module("flagsparse.sparse_operations.spmv_csc") +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") + + +def _value_dtype_cases(): + cases = [ + ("float32", torch.float32), + ("float64", torch.float64), + ("complex64", torch.complex64), + ("complex128", torch.complex128), + ] + return [(name, dtype) for name, dtype in cases if dtype is not None] + + +def _random_values(shape, dtype, device): + if dtype in (torch.float32, torch.float64): + return torch.randn(shape, dtype=dtype, device=device) + if dtype == torch.complex64: + return torch.complex( + torch.randn(shape, dtype=torch.float32, device=device), + torch.randn(shape, dtype=torch.float32, device=device), + ) + if dtype == torch.complex128: + return torch.complex( + torch.randn(shape, dtype=torch.float64, device=device), + torch.randn(shape, dtype=torch.float64, device=device), + ) + raise TypeError(f"unsupported dtype: {dtype}") + + +def _reference_dtype(dtype): + if dtype == torch.float32: + return torch.float64 + if dtype == torch.complex64: + return torch.complex128 + return dtype + + +def _random_csc_mn(M, N, dtype, index_dtype, device): + denom = max(M * N, 1) + p = min(0.25, max(0.06, 32.0 / denom)) + mask = torch.rand(M, N, device=device) < p + if int(mask.sum().item()) == 0: + mask[0, 0] = True + dense = torch.where( + mask, + _random_values((M, N), dtype, device), + torch.zeros((), dtype=dtype, device=device), + ) + rows, cols = torch.nonzero(mask, as_tuple=True) + order = torch.argsort(cols * max(1, M) + rows) + rows = rows[order] + cols = cols[order] + data = dense[rows, cols].contiguous() + col_counts = torch.bincount(cols, minlength=N) + indptr = torch.zeros(N + 1, dtype=torch.int64, device=device) + indptr[1:] = torch.cumsum(col_counts, dim=0) + return data, rows.to(index_dtype).contiguous(), indptr.to(index_dtype), dense + + +def _make_x(length, dtype, device): + return _random_values((length,), dtype, device) + + +def _op_transposes(op): + return op in ("trans", "conj") + + +def _apply_dense_op(dense, op): + if op == "non": + return dense + if op == "trans": + return dense.t() + if op == "conj": + return dense.conj().t() + raise ValueError(f"unsupported op: {op}") + + +def _assert_close(actual, expected, dtype): + rtol, atol = close_tolerances(dtype) + ref_dtype = _reference_dtype(dtype) + assert torch.allclose( + actual.to(ref_dtype), expected.to(ref_dtype), rtol=rtol, atol=atol + ) + + +@pytest.mark.spmv_csc +@pytest.mark.parametrize("M, N", SPMV_MN_SHAPES) +@pytest.mark.parametrize( + "name,dtype", _value_dtype_cases(), ids=[c[0] for c in _value_dtype_cases()] +) +@pytest.mark.parametrize( + "index_dtype", [torch.int32, torch.int64], ids=["int32", "int64"] +) +@pytest.mark.parametrize("op", ["non", "trans", "conj"], ids=["non", "trans", "conj"]) +def test_spmv_csc_matches_dense_reference(M, N, name, dtype, index_dtype, op): + device = torch.device("cuda") + data, indices, indptr, dense = _random_csc_mn(M, N, dtype, index_dtype, device) + x_len = M if _op_transposes(op) else N + x = _make_x(x_len, dtype, device) + ref_dtype = _reference_dtype(dtype) + ref = (_apply_dense_op(dense, op).to(ref_dtype) @ x.to(ref_dtype)).to(dtype) + + out = flagsparse_spmv_csc( + data, + indices, + indptr, + x, + shape=(M, N), + op=op, + index_fallback_policy="auto", + ) + _assert_close(out, ref, dtype) + + +@pytest.mark.spmv_csc +@pytest.mark.parametrize("op", ["non", "trans", "conj"], ids=["non", "trans", "conj"]) +def test_spmv_csc_prepared_path_matches_dense_reference(op): + device = torch.device("cuda") + M, N = 8, 10 + dtype = torch.complex64 + data, indices, indptr, dense = _random_csc_mn(M, N, dtype, torch.int32, device) + prepared = prepare_spmv_csc(data, indices, indptr, (M, N), op=op) + x_len = M if _op_transposes(op) else N + x = _make_x(x_len, dtype, device) + ref_dtype = _reference_dtype(dtype) + ref = (_apply_dense_op(dense, op).to(ref_dtype) @ x.to(ref_dtype)).to(dtype) + + out = flagsparse_spmv_csc(x=x, prepared=prepared) + _assert_close(out, ref, dtype) + + +@pytest.mark.spmv_csc +def test_spmv_csc_prepared_transpose_mismatch_rejected(): + device = torch.device("cuda") + data, indices, indptr, _dense = _random_csc_mn( + 8, 10, torch.float32, torch.int32, device + ) + prepared = prepare_spmv_csc(data, indices, indptr, (8, 10), transpose=True) + x = torch.randn(10, dtype=torch.float32, device=device) + with pytest.raises(ValueError, match="does not match prepared.transpose"): + flagsparse_spmv_csc(x=x, prepared=prepared, transpose=False) + + +@pytest.mark.spmv_csc +def test_spmv_csc_prepared_op_mismatch_rejected(): + device = torch.device("cuda") + data, indices, indptr, _dense = _random_csc_mn( + 8, 10, torch.complex64, torch.int32, device + ) + prepared = prepare_spmv_csc(data, indices, indptr, (8, 10), op="conj") + x = _make_x(8, torch.complex64, device) + with pytest.raises(ValueError, match="does not match prepared.op"): + flagsparse_spmv_csc(x=x, prepared=prepared, op="trans") + + +@pytest.mark.spmv_csc +def test_spmv_csc_int64_auto_fallback_to_int32(monkeypatch): + device = torch.device("cuda") + data, indices, indptr, dense = _random_csc_mn( + 12, 9, torch.float32, torch.int64, device + ) + x = torch.randn(9, dtype=torch.float32, device=device) + ref = dense.to(torch.float64) @ x.to(torch.float64) + state = {"forced_once": False} + original = spmv_csc_mod._triton_spmv_csc_kernel + + def fail_int64_once(prepared, x_in, op_code): + if prepared.kernel_indices.dtype == torch.int64 and not state["forced_once"]: + state["forced_once"] = True + raise RuntimeError("forced int64 launch failure") + return original(prepared, x_in, op_code) + + monkeypatch.setattr(spmv_csc_mod, "_triton_spmv_csc_kernel", fail_int64_once) + out = flagsparse_spmv_csc( + data, + indices, + indptr, + x, + shape=(12, 9), + index_fallback_policy="auto", + ) + assert state["forced_once"] + _assert_close(out, ref.to(torch.float32), torch.float32) + + +@pytest.mark.spmv_csc +def test_spmv_csc_int64_strict_no_fallback(monkeypatch): + device = torch.device("cuda") + data, indices, indptr, _dense = _random_csc_mn( + 12, 9, torch.float32, torch.int64, device + ) + x = torch.randn(9, dtype=torch.float32, device=device) + original = spmv_csc_mod._triton_spmv_csc_kernel + + def fail_int64(prepared, x_in, op_code): + if prepared.kernel_indices.dtype == torch.int64: + raise RuntimeError("forced int64 launch failure") + return original(prepared, x_in, op_code) + + monkeypatch.setattr(spmv_csc_mod, "_triton_spmv_csc_kernel", fail_int64) + with pytest.raises(RuntimeError, match="forced int64 launch failure"): + flagsparse_spmv_csc( + data, + indices, + indptr, + x, + shape=(12, 9), + index_fallback_policy="strict", + ) diff --git a/tests/test_spmv_csc.py b/tests/test_spmv_csc.py new file mode 100644 index 0000000..8f056cb --- /dev/null +++ b/tests/test_spmv_csc.py @@ -0,0 +1,647 @@ +"""Native CSC SpMV benchmark and correctness script.""" + +import argparse +import csv +import glob +import math +import os +import sys +from pathlib import Path + +import torch + +_PROJECT_ROOT = Path(__file__).resolve().parents[1] +_SRC_ROOT = _PROJECT_ROOT / "src" +if str(_SRC_ROOT) not in sys.path: + sys.path.insert(0, str(_SRC_ROOT)) + +import flagsparse as fs + +try: + import cupy as cp + import cupyx.scipy.sparse as cpx_sparse +except ImportError: + cp = None + cpx_sparse = None + + +VALUE_DTYPES = (torch.float32, torch.float64, torch.complex64, torch.complex128) +INDEX_DTYPES = (torch.int32, torch.int64) +OPS = ("non", "trans", "conj") +TEST_SIZES = ((64, 96), (160, 1024), (128, 256)) +WARMUP = 10 +ITERS = 50 + + +def _dtype_name(dtype): + return str(dtype).replace("torch.", "") + + +DTYPE_MAP = { + "float32": torch.float32, + "float64": torch.float64, + "complex64": torch.complex64, + "complex128": torch.complex128, +} +INDEX_DTYPE_MAP = {"int32": torch.int32, "int64": torch.int64} + + +def _parse_csv_tokens(value, mapping, option_name): + tokens = [token.strip().lower() for token in str(value).split(",") if token.strip()] + if not tokens: + raise ValueError(f"{option_name} must not be empty") + invalid = [token for token in tokens if token not in mapping] + if invalid: + raise ValueError( + f"unsupported {option_name}: {', '.join(invalid)}; allowed: {', '.join(mapping)}" + ) + return [mapping[token] for token in tokens] + + +def _parse_ops(value): + token = "non,trans,conj" if value is None else str(value).strip().lower() + if token == "all": + return list(OPS) + ops = [item.strip().lower() for item in token.split(",") if item.strip()] + invalid = [op for op in ops if op not in OPS] + if not ops or invalid: + raise ValueError(f"unsupported --ops: {', '.join(invalid or ops)}") + return ops + + +def _random_values(shape, dtype, device): + if dtype in (torch.float32, torch.float64): + return torch.randn(shape, dtype=dtype, device=device) + if dtype == torch.complex64: + return torch.complex( + torch.randn(shape, dtype=torch.float32, device=device), + torch.randn(shape, dtype=torch.float32, device=device), + ) + if dtype == torch.complex128: + return torch.complex( + torch.randn(shape, dtype=torch.float64, device=device), + torch.randn(shape, dtype=torch.float64, device=device), + ) + raise TypeError(f"unsupported dtype: {dtype}") + + +def _reference_dtype(dtype): + if dtype == torch.float32: + return torch.float64 + if dtype == torch.complex64: + return torch.complex128 + return dtype + + +def _reference_tolerance(dtype): + if dtype in (torch.float32, torch.complex64): + return 1.3e-6, 1e-3 + if dtype in (torch.float64, torch.complex128): + return 1e-7, 1e-5 + return 1e-6, 1e-5 + + +def _op_transposes(op): + return op in ("trans", "conj") + + +def _x_size_for_op(shape, op): + return int(shape[0]) if _op_transposes(op) else int(shape[1]) + + +def _out_size_for_op(shape, op): + return int(shape[1]) if _op_transposes(op) else int(shape[0]) + + +def _dense_to_csc(dense, index_dtype): + rows, cols = dense.nonzero(as_tuple=True) + if rows.numel() == 0: + data = dense.new_empty((0,)) + indices = torch.empty(0, dtype=index_dtype, device=dense.device) + indptr = torch.zeros(int(dense.shape[1]) + 1, dtype=index_dtype, device=dense.device) + return data, indices, indptr + order = torch.argsort(cols * max(1, int(dense.shape[0])) + rows) + rows = rows[order] + cols = cols[order] + data = dense[rows, cols].contiguous() + col_counts = torch.bincount(cols, minlength=int(dense.shape[1])) + indptr = torch.zeros(int(dense.shape[1]) + 1, dtype=torch.int64, device=dense.device) + indptr[1:] = torch.cumsum(col_counts, dim=0) + return data, rows.to(index_dtype).contiguous(), indptr.to(index_dtype) + + +def _csc_col_indices(indptr): + counts = indptr[1:].to(torch.int64) - indptr[:-1].to(torch.int64) + return torch.repeat_interleave( + torch.arange(indptr.numel() - 1, dtype=torch.int64, device=indptr.device), + counts, + ) + + +def _csc_to_torch_coo(data, indices, indptr, shape): + cols = _csc_col_indices(indptr) + row = indices.to(torch.int64) + return torch.sparse_coo_tensor( + torch.stack([row, cols]), + data, + size=shape, + device=data.device, + dtype=data.dtype, + ).coalesce() + + +def _pytorch_reference(data, indices, indptr, x, shape, dtype, op): + ref_dtype = _reference_dtype(dtype) + A = _csc_to_torch_coo( + data.to(ref_dtype), + indices, + indptr, + shape, + ) + x_ref = x.to(ref_dtype) + if op == "non": + out = torch.sparse.mm(A, x_ref.unsqueeze(1)).squeeze(1) + elif op == "trans": + out = torch.sparse.mm(A.transpose(0, 1), x_ref.unsqueeze(1)).squeeze(1) + elif op == "conj": + out = torch.sparse.mm(A.conj().transpose(0, 1), x_ref.unsqueeze(1)).squeeze(1) + else: + raise ValueError(f"unsupported op: {op}") + return out.to(dtype) + + +def _allclose_error_ratio(actual, expected, atol, rtol): + if expected.numel() == 0: + return 0.0 + diff = torch.abs(actual - expected).to(torch.float64) + denom = atol + rtol * torch.abs(expected).to(torch.float64) + return float(torch.max(diff / denom).item()) + + +def _cuda_event_benchmark(op, warmup, iters): + out = None + count = max(1, int(iters)) + for _ in range(max(0, int(warmup))): + out = op() + torch.cuda.synchronize() + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(count): + out = op() + end.record() + torch.cuda.synchronize() + return out, start.elapsed_time(end) / count + + +def _time_flagsparse_csc(data, indices, indptr, x, shape, op, warmup, iters, timing=False): + prepared = fs.prepare_spmv_csc(data, indices, indptr, shape, op=op) + out, gpu_ms = _cuda_event_benchmark( + lambda: fs.flagsparse_spmv_csc(x=x, prepared=prepared), + warmup, + iters, + ) + return { + "out": out, + "ms": gpu_ms, + "gpu_ms": gpu_ms, + "process_cpu_ms": 0.0, + "process_gpu_ms": 0.0 if timing else None, + "compute_ms": gpu_ms if timing else None, + } + + +def _time_pytorch(data, indices, indptr, x, shape, op, warmup, iters): + A = _csc_to_torch_coo(data, indices, indptr, shape) + if op == "non": + fn = lambda: torch.sparse.mm(A, x.unsqueeze(1)).squeeze(1) + elif op == "trans": + At = A.transpose(0, 1) + fn = lambda: torch.sparse.mm(At, x.unsqueeze(1)).squeeze(1) + else: + AH = A.conj().transpose(0, 1) + fn = lambda: torch.sparse.mm(AH, x.unsqueeze(1)).squeeze(1) + _, ms = _cuda_event_benchmark(fn, warmup, iters) + return ms + + +def _time_cusparse(data, indices, indptr, x, shape, op, warmup, iters): + if cp is None or cpx_sparse is None: + return None + if data.dtype not in (torch.float32, torch.float64, torch.complex64, torch.complex128): + return None + data_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(data)) + ind_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(indices.to(torch.int64))) + ptr_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(indptr.to(torch.int64))) + x_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(x)) + A = cpx_sparse.csc_matrix((data_cp, ind_cp, ptr_cp), shape=shape) + if op == "non": + fn = lambda: A @ x_cp + elif op == "trans": + fn = lambda: A.T @ x_cp + else: + fn = lambda: A.conj().T @ x_cp + for _ in range(max(0, int(warmup))): + _ = fn() + cp.cuda.runtime.deviceSynchronize() + start = cp.cuda.Event() + end = cp.cuda.Event() + count = max(1, int(iters)) + start.record() + for _ in range(count): + _ = fn() + end.record() + end.synchronize() + return cp.cuda.get_elapsed_time(start, end) / count + + +def _fmt(v): + return "N/A" if v is None else f"{v:.4f}" + + +def _fmt_err(v): + return "N/A" if v is None else f"{v:.2e}" + + +def _spd(base, other): + if base is None or other is None or other <= 0: + return "N/A" + return f"{base / other:.2f}x" + + +def _status(ok): + return "PASS" if ok else "FAIL" + + +def _header(timing=False): + split = f" {'ProcGPU':>9} {'Compute':>9}" if timing else "" + return ( + f"{'Matrix':<28} {'Op':>5} {'Out':>7} {'N_rows':>7} {'N_cols':>7} {'NNZ':>10} " + f"{'CSC(ms)':>9} {'CSCGPU':>9} {'CPUProc':>9}{split} " + f"{'PT(ms)':>9} {'CU(ms)':>9} {'CSC/PT':>8} {'CSC/CU':>8} " + f"{'Err':>10} {'Status':>6}" + ) + + +def _sep(timing=False): + return "-" * (150 if timing else 130) + + +def _print_row(row, timing=False): + name = str(row["matrix"])[:27] + if len(str(row["matrix"])) > 27: + name += "..." + split = ( + f" {_fmt(row.get('process_gpu_ms')):>9} {_fmt(row.get('compute_ms')):>9}" + if timing + else "" + ) + print( + f"{name:<28} {row['op']:>5} {row['out_size']:>7} {row['n_rows']:>7} {row['n_cols']:>7} {row['nnz']:>10} " + f"{_fmt(row['csc_ms']):>9} {_fmt(row['csc_gpu_ms']):>9} {_fmt(row['process_cpu_ms']):>9}{split} " + f"{_fmt(row['pytorch_ms']):>9} {_fmt(row['cusparse_ms']):>9} " + f"{_spd(row['pytorch_ms'], row['csc_ms']):>8} {_spd(row['cusparse_ms'], row['csc_ms']):>8} " + f"{_fmt_err(row['err']):>10} {row['status']:>6}" + ) + + +def _run_one_case( + data, + indices, + indptr, + shape, + dtype, + index_dtype, + op, + matrix_name, + warmup, + iters, + timing=False, + run_cusparse=True, +): + indices = indices.to(index_dtype).contiguous() + indptr = indptr.to(index_dtype).contiguous() + x = _random_values((_x_size_for_op(shape, op),), dtype, data.device) + atol, rtol = _reference_tolerance(dtype) + csc = _time_flagsparse_csc(data, indices, indptr, x, shape, op, warmup, iters, timing=timing) + y_ref = _pytorch_reference(data, indices, indptr, x, shape, dtype, op) + err = _allclose_error_ratio(csc["out"], y_ref, atol, rtol) + pt_ms = None + cu_ms = None + try: + pt_ms = _time_pytorch(data, indices, indptr, x, shape, op, warmup, iters) + except Exception: + pass + if run_cusparse: + try: + cu_ms = _time_cusparse(data, indices, indptr, x, shape, op, warmup, iters) + except Exception: + pass + ok = (not math.isnan(err)) and err <= 1.0 + return { + "matrix": matrix_name, + "value_dtype": _dtype_name(dtype), + "index_dtype": _dtype_name(index_dtype), + "op": op, + "out_size": _out_size_for_op(shape, op), + "n_rows": int(shape[0]), + "n_cols": int(shape[1]), + "nnz": int(data.numel()), + "csc_ms": csc["ms"], + "csc_gpu_ms": csc["gpu_ms"], + "process_cpu_ms": csc["process_cpu_ms"], + "process_gpu_ms": csc["process_gpu_ms"], + "compute_ms": csc["compute_ms"], + "pytorch_ms": pt_ms, + "cusparse_ms": cu_ms, + "err": err, + "status": _status(ok), + } + + +def _mtx_value_for_dtype(raw_value, dtype): + if dtype in (torch.complex64, torch.complex128): + return complex(raw_value) + return float(raw_value.real if isinstance(raw_value, complex) else raw_value) + + +def load_mtx_to_csc_torch(path, dtype=torch.float32, device=None): + device = torch.device("cuda" if device is None else device) + with open(path, "r", encoding="utf-8") as handle: + lines = handle.readlines() + mm_field = "real" + mm_symmetry = "general" + header = None + data_lines = [] + for line in lines: + stripped = line.strip() + if stripped.startswith("%%MatrixMarket"): + parts = stripped.split() + if len(parts) >= 5: + mm_field = parts[3].lower() + mm_symmetry = parts[4].lower() + continue + if stripped.startswith("%"): + continue + if header is None and stripped: + parts = stripped.split() + header = (int(parts[0]), int(parts[1]), int(parts[2]) if len(parts) > 2 else 0) + continue + if stripped: + data_lines.append(stripped) + if header is None: + raise ValueError(f"Cannot parse .mtx header: {path}") + n_rows, n_cols, nnz = header + entries = {} + + def add_entry(r, c, value): + key = (int(r), int(c)) + entries[key] = entries.get(key, 0.0) + value + + is_pattern = mm_field == "pattern" + is_complex = mm_field == "complex" + is_symmetric = mm_symmetry == "symmetric" + is_skew = mm_symmetry == "skew-symmetric" + is_hermitian = mm_symmetry == "hermitian" + for line in data_lines[:nnz]: + parts = line.split() + if len(parts) < 2: + continue + r = int(parts[0]) - 1 + c = int(parts[1]) - 1 + if not (0 <= r < n_rows and 0 <= c < n_cols): + continue + if is_pattern: + value = 1.0 + elif is_complex: + value = complex(float(parts[2]), float(parts[3])) + else: + value = float(parts[2]) + add_entry(r, c, value) + if r != c: + if is_symmetric and 0 <= c < n_rows and 0 <= r < n_cols: + add_entry(c, r, value) + elif is_skew and 0 <= c < n_rows and 0 <= r < n_cols: + add_entry(c, r, -value) + elif is_hermitian and 0 <= c < n_rows and 0 <= r < n_cols: + add_entry(c, r, value.conjugate() if isinstance(value, complex) else value) + sorted_entries = sorted(entries.items(), key=lambda item: (item[0][1], item[0][0])) + rows = [key[0] for key, _ in sorted_entries] + cols = [key[1] for key, _ in sorted_entries] + vals = [_mtx_value_for_dtype(value, dtype) for _, value in sorted_entries] + data = torch.tensor(vals, dtype=dtype, device=device) + indices = torch.tensor(rows, dtype=torch.int64, device=device) + col_tensor = torch.tensor(cols, dtype=torch.int64, device=device) + col_counts = torch.bincount(col_tensor, minlength=n_cols) if col_tensor.numel() else torch.zeros(n_cols, dtype=torch.int64, device=device) + indptr = torch.zeros(n_cols + 1, dtype=torch.int64, device=device) + indptr[1:] = torch.cumsum(col_counts, dim=0) + return data, indices, indptr, (n_rows, n_cols) + + +def run_synthetic(value_dtypes=None, index_dtypes=None, ops=None, warmup=WARMUP, iters=ITERS, timing=False, run_cusparse=True): + if not torch.cuda.is_available(): + print("CUDA is not available.") + return + device = torch.device("cuda") + value_dtypes = VALUE_DTYPES if value_dtypes is None else value_dtypes + index_dtypes = INDEX_DTYPES if index_dtypes is None else index_dtypes + ops = OPS if ops is None else ops + print("=" * 140) + print("FLAGSPARSE SpMV CSC BENCHMARK (native CSC Triton)") + print("=" * 140) + print("Timing policy: csc_ms = process_cpu_ms + csc_gpu_ms; CSC v1 has no process phase.") + for dtype in value_dtypes: + for index_dtype in index_dtypes: + for op in ops: + print(_sep(timing)) + print(f"dtype: {_dtype_name(dtype)} | index_dtype: {_dtype_name(index_dtype)} | op: {op}") + print(_sep(timing)) + print(_header(timing)) + print(_sep(timing)) + for m, n in TEST_SIZES: + dense = _random_values((m, n), dtype, device) + dense *= (torch.rand(m, n, device=device) < 0.1).to(dtype=dtype) + data, indices, indptr = _dense_to_csc(dense, index_dtype) + row = _run_one_case( + data, + indices, + indptr, + (m, n), + dtype, + index_dtype, + op, + f"{m}x{n}", + warmup, + iters, + timing=timing, + run_cusparse=run_cusparse, + ) + _print_row(row, timing=timing) + print(_sep(timing)) + print() + + +def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, ops=None, warmup=WARMUP, iters=ITERS, timing=False, run_cusparse=True): + if not torch.cuda.is_available(): + print("CUDA is not available.") + return + device = torch.device("cuda") + value_dtypes = VALUE_DTYPES if value_dtypes is None else value_dtypes + index_dtypes = INDEX_DTYPES if index_dtypes is None else index_dtypes + ops = OPS if ops is None else ops + rows = [] + for dtype in value_dtypes: + for index_dtype in index_dtypes: + for op in ops: + print(_sep(timing)) + print(f"Value dtype: {_dtype_name(dtype)} | Index dtype: {_dtype_name(index_dtype)} | op: {op}") + print(_sep(timing)) + print(_header(timing)) + print(_sep(timing)) + for path in mtx_paths: + try: + data, indices, indptr, shape = load_mtx_to_csc_torch(path, dtype=dtype, device=device) + row = _run_one_case( + data, + indices, + indptr, + shape, + dtype, + index_dtype, + op, + os.path.basename(path), + warmup, + iters, + timing=timing, + run_cusparse=run_cusparse, + ) + except Exception as exc: + row = { + "matrix": os.path.basename(path), + "value_dtype": _dtype_name(dtype), + "index_dtype": _dtype_name(index_dtype), + "op": op, + "out_size": "ERR", + "n_rows": "ERR", + "n_cols": "ERR", + "nnz": "ERR", + "csc_ms": None, + "csc_gpu_ms": None, + "process_cpu_ms": None, + "process_gpu_ms": None, + "compute_ms": None, + "pytorch_ms": None, + "cusparse_ms": None, + "err": None, + "status": "ERROR", + "error": str(exc), + } + rows.append(row) + _print_row(row, timing=timing) + print(_sep(timing)) + fieldnames = [ + "matrix", + "value_dtype", + "index_dtype", + "op", + "out_size", + "n_rows", + "n_cols", + "nnz", + "csc_ms", + "csc_gpu_ms", + "process_cpu_ms", + "process_gpu_ms", + "compute_ms", + "pytorch_ms", + "cusparse_ms", + "err", + "status", + "error", + ] + if not timing: + fieldnames = [ + field + for field in fieldnames + if field not in ("process_gpu_ms", "compute_ms") + ] + with open(csv_path, "w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fieldnames, extrasaction="ignore") + writer.writeheader() + for row in rows: + writer.writerow({key: ("" if value is None else value) for key, value in row.items()}) + print(f"Wrote {len(rows)} rows to {csv_path}") + + +def main(): + parser = argparse.ArgumentParser(description="Native CSC SpMV benchmark/test.") + parser.add_argument("mtx", nargs="*", help=".mtx files or directories") + parser.add_argument("--synthetic", action="store_true") + parser.add_argument("--csv-csc", type=str, default=None, metavar="FILE") + parser.add_argument("--dtypes", default="float32,float64,complex64,complex128") + parser.add_argument("--index-dtypes", default="int32,int64") + parser.add_argument("--ops", default="non,trans,conj") + parser.add_argument("--warmup", type=int, default=WARMUP) + parser.add_argument("--iters", type=int, default=ITERS) + parser.add_argument("--timing", action="store_true") + parser.add_argument("--no-cusparse", action="store_true") + args = parser.parse_args() + try: + value_dtypes = _parse_csv_tokens(args.dtypes, DTYPE_MAP, "--dtypes") + index_dtypes = _parse_csv_tokens(args.index_dtypes, INDEX_DTYPE_MAP, "--index-dtypes") + ops = _parse_ops(args.ops) + except ValueError as exc: + parser.error(str(exc)) + if args.synthetic: + run_synthetic( + value_dtypes=value_dtypes, + index_dtypes=index_dtypes, + ops=ops, + warmup=args.warmup, + iters=args.iters, + timing=args.timing, + run_cusparse=not args.no_cusparse, + ) + return + paths = [] + for path in args.mtx: + if os.path.isfile(path) and path.endswith(".mtx"): + paths.append(path) + elif os.path.isdir(path): + paths.extend(sorted(glob.glob(os.path.join(path, "*.mtx")))) + if args.csv_csc: + if not paths: + paths = sorted(glob.glob("*.mtx")) + if not paths: + print("No .mtx files found for --csv-csc") + return + run_csv( + paths, + args.csv_csc, + value_dtypes=value_dtypes, + index_dtypes=index_dtypes, + ops=ops, + warmup=args.warmup, + iters=args.iters, + timing=args.timing, + run_cusparse=not args.no_cusparse, + ) + return + if not paths: + print("No .mtx files. Use --synthetic or --csv-csc with inputs.") + return + run_csv( + paths, + "spmv_csc_results.csv", + value_dtypes=value_dtypes, + index_dtypes=index_dtypes, + ops=ops, + warmup=args.warmup, + iters=args.iters, + timing=args.timing, + run_cusparse=not args.no_cusparse, + ) + + +if __name__ == "__main__": + main() diff --git a/tools/ci/run_gpu_benchmark.py b/tools/ci/run_gpu_benchmark.py index 1442da3..27b40b8 100644 --- a/tools/ci/run_gpu_benchmark.py +++ b/tools/ci/run_gpu_benchmark.py @@ -27,6 +27,7 @@ def _parse_args() -> argparse.Namespace: "scatter", "spmv", "spmv-coo", + "spmv-csc", "spmm", "spmm-coo", "spsv", @@ -93,6 +94,15 @@ def _command_specs( "tests/test_spmv_coo.py", "--synthetic", ], + "spmv-csc": [ + "tests/test_spmv_csc.py", + "--synthetic", + "--warmup", + str(args.warmup), + "--iters", + str(args.iters), + *no_cusparse, + ], "spmm": [ "tests/test_spmm.py", "--synthetic", @@ -135,6 +145,7 @@ def _command_specs( "scatter", "spmv", "spmv-coo", + "spmv-csc", "spmm", "spmm-coo", "spsv", From c6ed11005f2cd4520b95cfab8128c7d17a994fc6 Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Thu, 2 Jul 2026 10:14:05 +0800 Subject: [PATCH 02/13] spmv_csc --- tests/test_spmv_csc.py | 101 +++++++++++++++++++++++++++++++---------- 1 file changed, 78 insertions(+), 23 deletions(-) diff --git a/tests/test_spmv_csc.py b/tests/test_spmv_csc.py index 8f056cb..2ca948d 100644 --- a/tests/test_spmv_csc.py +++ b/tests/test_spmv_csc.py @@ -303,6 +303,9 @@ def _print_row(row, timing=False): f"{_spd(row['pytorch_ms'], row['csc_ms']):>8} {_spd(row['cusparse_ms'], row['csc_ms']):>8} " f"{_fmt_err(row['err']):>10} {row['status']:>6}" ) + error = row.get("error") + if error: + print(f" error: {str(error)[:240]}") def _run_one_case( @@ -323,9 +326,48 @@ def _run_one_case( indptr = indptr.to(index_dtype).contiguous() x = _random_values((_x_size_for_op(shape, op),), dtype, data.device) atol, rtol = _reference_tolerance(dtype) - csc = _time_flagsparse_csc(data, indices, indptr, x, shape, op, warmup, iters, timing=timing) - y_ref = _pytorch_reference(data, indices, indptr, x, shape, dtype, op) - err = _allclose_error_ratio(csc["out"], y_ref, atol, rtol) + base_row = { + "matrix": matrix_name, + "value_dtype": _dtype_name(dtype), + "index_dtype": _dtype_name(index_dtype), + "op": op, + "out_size": _out_size_for_op(shape, op), + "n_rows": int(shape[0]), + "n_cols": int(shape[1]), + "nnz": int(data.numel()), + "csc_ms": None, + "csc_gpu_ms": None, + "process_cpu_ms": 0.0, + "process_gpu_ms": 0.0 if timing else None, + "compute_ms": None, + "pytorch_ms": None, + "cusparse_ms": None, + "err": None, + "status": "ERROR", + "error": None, + } + try: + csc = _time_flagsparse_csc( + data, indices, indptr, x, shape, op, warmup, iters, timing=timing + ) + except Exception as exc: + base_row["error"] = f"flagsparse_spmv_csc failed: {exc}" + return base_row + base_row.update( + { + "csc_ms": csc["ms"], + "csc_gpu_ms": csc["gpu_ms"], + "process_cpu_ms": csc["process_cpu_ms"], + "process_gpu_ms": csc["process_gpu_ms"], + "compute_ms": csc["compute_ms"], + } + ) + try: + y_ref = _pytorch_reference(data, indices, indptr, x, shape, dtype, op) + err = _allclose_error_ratio(csc["out"], y_ref, atol, rtol) + except Exception as exc: + base_row["error"] = f"reference failed after CSC run: {exc}" + return base_row pt_ms = None cu_ms = None try: @@ -338,25 +380,16 @@ def _run_one_case( except Exception: pass ok = (not math.isnan(err)) and err <= 1.0 - return { - "matrix": matrix_name, - "value_dtype": _dtype_name(dtype), - "index_dtype": _dtype_name(index_dtype), - "op": op, - "out_size": _out_size_for_op(shape, op), - "n_rows": int(shape[0]), - "n_cols": int(shape[1]), - "nnz": int(data.numel()), - "csc_ms": csc["ms"], - "csc_gpu_ms": csc["gpu_ms"], - "process_cpu_ms": csc["process_cpu_ms"], - "process_gpu_ms": csc["process_gpu_ms"], - "compute_ms": csc["compute_ms"], - "pytorch_ms": pt_ms, - "cusparse_ms": cu_ms, - "err": err, - "status": _status(ok), - } + base_row.update( + { + "pytorch_ms": pt_ms, + "cusparse_ms": cu_ms, + "err": err, + "status": _status(ok), + "error": None if ok else "correctness check failed", + } + ) + return base_row def _mtx_value_for_dtype(raw_value, dtype): @@ -481,7 +514,18 @@ def run_synthetic(value_dtypes=None, index_dtypes=None, ops=None, warmup=WARMUP, print() -def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, ops=None, warmup=WARMUP, iters=ITERS, timing=False, run_cusparse=True): +def run_csv( + mtx_paths, + csv_path, + value_dtypes=None, + index_dtypes=None, + ops=None, + warmup=WARMUP, + iters=ITERS, + timing=False, + run_cusparse=True, + fail_fast=False, +): if not torch.cuda.is_available(): print("CUDA is not available.") return @@ -516,6 +560,8 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, ops=None, run_cusparse=run_cusparse, ) except Exception as exc: + if fail_fast: + raise row = { "matrix": os.path.basename(path), "value_dtype": _dtype_name(dtype), @@ -536,6 +582,8 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, ops=None, "status": "ERROR", "error": str(exc), } + if fail_fast and row.get("status") == "ERROR": + raise RuntimeError(row.get("error") or "CSC SpMV case failed") rows.append(row) _print_row(row, timing=timing) print(_sep(timing)) @@ -565,6 +613,9 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, ops=None, for field in fieldnames if field not in ("process_gpu_ms", "compute_ms") ] + csv_parent = Path(csv_path).parent + if str(csv_parent) not in ("", "."): + csv_parent.mkdir(parents=True, exist_ok=True) with open(csv_path, "w", newline="", encoding="utf-8") as handle: writer = csv.DictWriter(handle, fieldnames=fieldnames, extrasaction="ignore") writer.writeheader() @@ -585,6 +636,7 @@ def main(): parser.add_argument("--iters", type=int, default=ITERS) parser.add_argument("--timing", action="store_true") parser.add_argument("--no-cusparse", action="store_true") + parser.add_argument("--fail-fast", action="store_true") args = parser.parse_args() try: value_dtypes = _parse_csv_tokens(args.dtypes, DTYPE_MAP, "--dtypes") @@ -601,6 +653,7 @@ def main(): iters=args.iters, timing=args.timing, run_cusparse=not args.no_cusparse, + fail_fast=args.fail_fast, ) return paths = [] @@ -625,6 +678,7 @@ def main(): iters=args.iters, timing=args.timing, run_cusparse=not args.no_cusparse, + fail_fast=args.fail_fast, ) return if not paths: @@ -640,6 +694,7 @@ def main(): iters=args.iters, timing=args.timing, run_cusparse=not args.no_cusparse, + fail_fast=args.fail_fast, ) From b1bbd5f1eb87881cdff4141ded9760f556b26de8 Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Thu, 2 Jul 2026 16:33:27 +0800 Subject: [PATCH 03/13] spmm_coo_opt --- src/flagsparse/__init__.py | 16 + src/flagsparse/sparse_operations/__init__.py | 20 +- src/flagsparse/sparse_operations/spmm_coo.py | 656 ++++++++++++++++++ tests/test_spmm_coo.py | 664 ++++++++++++++++--- 4 files changed, 1258 insertions(+), 98 deletions(-) diff --git a/src/flagsparse/__init__.py b/src/flagsparse/__init__.py index 608f6cb..bfcd42d 100644 --- a/src/flagsparse/__init__.py +++ b/src/flagsparse/__init__.py @@ -16,6 +16,7 @@ "comprehensive_gather_test", "comprehensive_scatter_test", "PreparedCoo", + "PreparedCooSpmmRoute", "PreparedCscSpmv", "PreparedAlphaSpmmAlg1", "PreparedCsrSpmv", @@ -32,6 +33,7 @@ "prepare_alpha_spmm_alg1_tle_opt", "prepare_alpha_spmm_alg1_tle_opt2", "prepare_spmm_csr_opt", + "prepare_spmm_coo_route", "prepare_spmm_csr_route", "prepare_spmm_csr_opt_alg1", "prepare_spmm_csr_opt_alg1_preprocess", @@ -91,6 +93,7 @@ "flagsparse_sddmm_csr", "flagsparse_spmm_csr", "flagsparse_spmm_coo", + "flagsparse_spmm_coo_run", "benchmark_spmm_case", "benchmark_spmm_opt_case", "benchmark_spmm_opt_alg2_case", @@ -99,9 +102,14 @@ "comprehensive_spmm_test", "comprehensive_spsm_test", "benchmark_spmv_case", + "list_spmm_coo_algorithms", "list_spmm_csr_algorithms", + "resolve_spmm_coo_algorithm", "resolve_spmm_csr_algorithm", + "SPMM_COO_ALGORITHMS", "SPMM_CSR_ALGORITHMS", + "SpmmCooAlgorithm", + "SpmmCooAlgorithmUnavailable", "SpmmCsrAlgorithm", "SpmmCsrAlgorithmUnavailable", "create_csr_matrix", @@ -139,6 +147,7 @@ "comprehensive_gather_test", "comprehensive_scatter_test", "PreparedCoo", + "PreparedCooSpmmRoute", "PreparedCscSpmv", "PreparedAlphaSpmmAlg1", "PreparedCsrSpmv", @@ -155,6 +164,7 @@ "prepare_alpha_spmm_alg1_tle_opt", "prepare_alpha_spmm_alg1_tle_opt2", "prepare_spmm_csr_opt", + "prepare_spmm_coo_route", "prepare_spmm_csr_route", "prepare_spmm_csr_opt_alg1", "prepare_spmm_csr_opt_alg1_preprocess", @@ -214,6 +224,7 @@ "flagsparse_sddmm_csr", "flagsparse_spmm_csr", "flagsparse_spmm_coo", + "flagsparse_spmm_coo_run", "benchmark_spmm_case", "benchmark_spmm_opt_case", "benchmark_spmm_opt_alg2_case", @@ -222,9 +233,14 @@ "comprehensive_spmm_test", "benchmark_spmv_case", "comprehensive_spsm_test", + "list_spmm_coo_algorithms", "list_spmm_csr_algorithms", + "resolve_spmm_coo_algorithm", "resolve_spmm_csr_algorithm", + "SPMM_COO_ALGORITHMS", "SPMM_CSR_ALGORITHMS", + "SpmmCooAlgorithm", + "SpmmCooAlgorithmUnavailable", "SpmmCsrAlgorithm", "SpmmCsrAlgorithmUnavailable", } diff --git a/src/flagsparse/sparse_operations/__init__.py b/src/flagsparse/sparse_operations/__init__.py index 5c3ce9e..d097a8b 100644 --- a/src/flagsparse/sparse_operations/__init__.py +++ b/src/flagsparse/sparse_operations/__init__.py @@ -33,7 +33,17 @@ ) from .sddmm_csr import SDDMMPrepared, benchmark_sddmm_case, flagsparse_sddmm_csr, prepare_sddmm_csr from .spgemm_csr import SpGEMMPrepared, benchmark_spgemm_case, flagsparse_spgemm_csr, prepare_spgemm_csr -from .spmm_coo import flagsparse_spmm_coo +from .spmm_coo import ( + PreparedCooSpmmRoute, + SPMM_COO_ALGORITHMS, + SpmmCooAlgorithm, + SpmmCooAlgorithmUnavailable, + flagsparse_spmm_coo, + flagsparse_spmm_coo_run, + list_spmm_coo_algorithms, + prepare_spmm_coo_route, + resolve_spmm_coo_algorithm, +) from .spmm_csr import ( PreparedCsrSpmmOpt, PreparedCsrSpmmRoute, @@ -110,6 +120,7 @@ __all__ = [ "PreparedCoo", + "PreparedCooSpmmRoute", "PreparedAlphaSpmmAlg1", "PreparedCsrSpmv", "PreparedCscSpmv", @@ -123,6 +134,8 @@ "FlagSparseDnVecDescr", "SpmmCsrAlgorithm", "SpmmCsrAlgorithmUnavailable", + "SpmmCooAlgorithm", + "SpmmCooAlgorithmUnavailable", "FlagSparseSpMatDescr", "FlagSparseSpSVDescr", "FlagSparseSpSVHandle", @@ -152,6 +165,7 @@ "flagsparse_sddmm_csr", "flagsparse_spgemm_csr", "flagsparse_spmm_coo", + "flagsparse_spmm_coo_run", "flagsparse_spmm_csr", "flagsparse_spmm_csr_run", "flagsparse_spmm_csr_opt", @@ -183,6 +197,7 @@ "flagsparse_spsv_solve_coo", "flagsparse_spsv_solve_csr", "list_spmm_csr_algorithms", + "list_spmm_coo_algorithms", "prepare_sddmm_csr", "build_alpha_spmm_alg1_tle_opt_meta", "build_alpha_spmm_alg1_tle_opt2_meta", @@ -201,6 +216,7 @@ "prepare_spmm_csr_opt_alg1", "prepare_spmm_csr_opt_alg1_preprocess", "prepare_spmm_csr_route", + "prepare_spmm_coo_route", "prepare_spmm_csr_opt_alg2", "prepare_spmm_csr_opt_alg2_preprocess", "prepare_spmv_coo", @@ -208,6 +224,8 @@ "prepare_spmv_csc", "prepare_spmv_csr", "resolve_spmm_csr_algorithm", + "resolve_spmm_coo_algorithm", + "SPMM_COO_ALGORITHMS", "SPMM_CSR_ALGORITHMS", "pytorch_index_gather", "pytorch_index_scatter", diff --git a/src/flagsparse/sparse_operations/spmm_coo.py b/src/flagsparse/sparse_operations/spmm_coo.py index 1b527fa..6c63602 100644 --- a/src/flagsparse/sparse_operations/spmm_coo.py +++ b/src/flagsparse/sparse_operations/spmm_coo.py @@ -1,5 +1,7 @@ """Native COO SpMM kernels, route helpers, and internal benchmark entry points.""" +from dataclasses import dataclass + from ._common import * from .spmm_csr import ( SUPPORTED_SPMM_VALUE_DTYPES, @@ -286,6 +288,62 @@ def _seg_starts_from_sorted_rows(row_i32, nnz, device): ) +@dataclass(frozen=True) +class SpmmCooAlgorithm: + """Registered COO SpMM route for the route-based run API.""" + + name: str + display_name: str + supported_ops: tuple + supported_dtypes: tuple + run: object + + +class SpmmCooAlgorithmUnavailable(RuntimeError): + """Raised when a registered COO SpMM algorithm is unavailable.""" + + +class PreparedCooSpmmRoute: + """Matrix-level COO SpMM route preparation shared by registered algorithms.""" + + __slots__ = ( + "data", + "row", + "col", + "shape", + "n_rows", + "n_cols", + "seg_starts", + "row_lengths", + "n_segs", + "nnz", + "max_row_nnz", + "avg_nnz_per_row", + "output_dtype", + "compute_dtype", + "op", + "alg", + ) + + def __init__(self, data, row, col, shape, seg_starts, row_lengths, output_dtype, compute_dtype, op, alg): + self.data = data + self.row = row + self.col = col + self.shape = (int(shape[0]), int(shape[1])) + self.n_rows = int(shape[0]) + self.n_cols = int(shape[1]) + self.seg_starts = seg_starts + self.row_lengths = row_lengths + self.n_segs = int(row_lengths.numel()) + self.nnz = int(data.numel()) + self.max_row_nnz = int(row_lengths.max().item()) if row_lengths.numel() else 0 + self.avg_nnz_per_row = float(self.nnz) / float(max(1, self.n_rows)) + self.output_dtype = output_dtype + self.compute_dtype = compute_dtype + self.op = str(op) + self.alg = str(alg) + + @triton.jit def _spmm_coo_rowrun_real_kernel( data_ptr, @@ -394,6 +452,100 @@ def _spmm_coo_rowrun_complex_kernel( mask=mask_n, ) + +@triton.jit +def _spmm_coo_alg1_process_count_kernel(row_lengths_ptr, counts_ptr, n_segs, BLOCK_M: tl.constexpr): + pid = tl.program_id(0) + offs = pid * BLOCK_M + tl.arange(0, BLOCK_M) + mask = offs < n_segs + lengths = tl.load(row_lengths_ptr + offs, mask=mask, other=0) + b0 = mask & (lengths <= 32) + b1 = mask & (lengths > 32) & (lengths <= 128) + b2 = mask & (lengths > 128) & (lengths <= 512) + b3 = mask & (lengths > 512) & (lengths <= 2048) + b4 = mask & (lengths > 2048) + tl.atomic_add(counts_ptr + 0, tl.sum(tl.where(b0, 1, 0)), sem="relaxed") + tl.atomic_add(counts_ptr + 1, tl.sum(tl.where(b1, 1, 0)), sem="relaxed") + tl.atomic_add(counts_ptr + 2, tl.sum(tl.where(b2, 1, 0)), sem="relaxed") + tl.atomic_add(counts_ptr + 3, tl.sum(tl.where(b3, 1, 0)), sem="relaxed") + tl.atomic_add(counts_ptr + 4, tl.sum(tl.where(b4, 1, 0)), sem="relaxed") + + +@triton.jit +def _spmm_coo_alg1_process_compact_kernel( + row_lengths_ptr, + offsets_ptr, + write_counts_ptr, + segs_flat_ptr, + n_segs, + BLOCK_M: tl.constexpr, +): + pid = tl.program_id(0) + offs = pid * BLOCK_M + tl.arange(0, BLOCK_M) + mask = offs < n_segs + lengths = tl.load(row_lengths_ptr + offs, mask=mask, other=0) + bucket = tl.full([BLOCK_M], 4, tl.int32) + bucket = tl.where(lengths <= 2048, 3, bucket) + bucket = tl.where(lengths <= 512, 2, bucket) + bucket = tl.where(lengths <= 128, 1, bucket) + bucket = tl.where(lengths <= 32, 0, bucket) + for b in tl.static_range(0, 5): + in_bucket = mask & (bucket == b) + local = tl.cumsum(tl.where(in_bucket, 1, 0), 0) - 1 + n_bucket = tl.sum(tl.where(in_bucket, 1, 0)) + base = tl.load(offsets_ptr + b) + tl.atomic_add(write_counts_ptr + b, n_bucket, sem="relaxed") + tl.store(segs_flat_ptr + base + local, offs, mask=in_bucket) + + +@triton.jit +def _spmm_coo_alg1_bucket_real_kernel( + data_ptr, + row_ptr, + col_ptr, + b_ptr, + c_ptr, + seg_starts_ptr, + bucket_segs_ptr, + n_bucket_segs, + n_dense_cols, + stride_bk, + stride_bn, + stride_cm, + stride_cn, + BLOCK_N: tl.constexpr, + BLOCK_NNZ: tl.constexpr, + ACC_DTYPE: tl.constexpr, +): + bucket_pos = tl.program_id(0) + pid_n = tl.program_id(1) + if bucket_pos >= n_bucket_segs: + return + + seg = tl.load(bucket_segs_ptr + bucket_pos) + offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) + mask_n = offs_n < n_dense_cols + start = tl.load(seg_starts_ptr + seg) + end = tl.load(seg_starts_ptr + seg + 1) + row_nnz = end - start + row_id = tl.load(row_ptr + start) + acc = tl.zeros([BLOCK_N], dtype=ACC_DTYPE) + + for chunk_start in tl.range(0, row_nnz, BLOCK_NNZ): + for kk in tl.static_range(0, BLOCK_NNZ): + idx = start + chunk_start + kk + valid = idx < end + a_val = tl.load(data_ptr + idx, mask=valid, other=0.0) + a_col = tl.load(col_ptr + idx, mask=valid, other=0) + b_vals = tl.load( + b_ptr + a_col * stride_bk + offs_n * stride_bn, + mask=mask_n & valid, + other=0.0, + ) + acc = acc + a_val.to(ACC_DTYPE) * b_vals.to(ACC_DTYPE) + + tl.store(c_ptr + row_id * stride_cm + offs_n * stride_cn, acc, mask=mask_n) + + @triton.jit def _spmm_coo_atomic_real_kernel( data_ptr, @@ -835,6 +987,510 @@ def _triton_spmm_coo_impl( ) +def _normalize_spmm_coo_alg(alg): + token = "auto" if alg is None else str(alg).strip().lower() + aliases = { + "rowrun": "coo_rowrun", + "coo": "coo_rowrun", + "base": "coo_rowrun", + "coo_base": "coo_rowrun", + "atomic": "coo_atomic", + "alg1": "spmm_coo_alg1", + "coo_alg1": "spmm_coo_alg1", + } + if token == "auto": + return "auto" + return aliases.get(token, token) + + +def _prepare_spmm_coo_matrix(data, row, col, shape): + if len(shape) != 2: + raise ValueError("shape must be a 2-tuple: (n_rows, n_cols)") + if data.ndim != 1 or row.ndim != 1 or col.ndim != 1: + raise ValueError("data, row, and col must be 1D tensors") + n_rows, n_cols = int(shape[0]), int(shape[1]) + if n_rows < 0 or n_cols < 0: + raise ValueError("shape dimensions must be non-negative") + if data.numel() != row.numel() or data.numel() != col.numel(): + raise ValueError("data, row, and col must have the same length (nnz)") + if not all(t.is_cuda for t in (data, row, col)): + raise ValueError("data, row, and col must be CUDA tensors") + if not all(t.device == data.device for t in (row, col)): + raise ValueError("data, row, and col must be on the same CUDA device") + if data.dtype not in SUPPORTED_SPMM_VALUE_DTYPES: + raise TypeError( + "data dtype must be one of: float16, bfloat16, float32, float64, complex64, complex128" + ) + if row.dtype not in SUPPORTED_INDEX_DTYPES: + raise TypeError("row dtype must be torch.int32 or torch.int64") + if col.dtype not in SUPPORTED_INDEX_DTYPES: + raise TypeError("col dtype must be torch.int32 or torch.int64") + nnz = int(data.numel()) + if nnz > _INDEX_LIMIT_INT32: + raise ValueError("nnz exceeds the int32 range supported by the Triton COO kernel") + if nnz > 0: + min_row = int(row.min().item()) + max_row = int(row.max().item()) + min_col = int(col.min().item()) + max_col = int(col.max().item()) + if min_row < 0 or max_row >= n_rows: + raise IndexError("row indices out of range for n_rows") + if min_col < 0 or max_col >= n_cols: + raise IndexError("col indices out of range for n_cols") + if max_row > _INDEX_LIMIT_INT32: + raise ValueError("row indices exceed the int32 range supported by the Triton kernel") + if max_col > _INDEX_LIMIT_INT32: + raise ValueError("column indices exceed the int32 range supported by the Triton kernel") + kernel_row = row.contiguous().to(torch.int32) if row.dtype == torch.int64 else row.contiguous() + kernel_col = col.contiguous().to(torch.int32) if col.dtype == torch.int64 else col.contiguous() + return data.contiguous(), kernel_row, kernel_col, (n_rows, n_cols) + + +def _validate_spmm_coo_route_runtime_inputs(prepared, B, dense_layout): + if B is None or not torch.is_tensor(B): + raise TypeError("B must be a torch.Tensor") + if B.ndim != 2: + raise ValueError("B must be a 2D dense tensor") + if not B.is_cuda: + raise ValueError("B must be a CUDA tensor") + if B.device != prepared.data.device: + raise ValueError("B must be on the same CUDA device as sparse matrix data") + if B.dtype != prepared.output_dtype: + raise TypeError("B dtype must match sparse matrix dtype") + if int(B.shape[0]) != prepared.n_cols: + raise ValueError(f"B.shape[0] must be n_cols={prepared.n_cols}, got {B.shape[0]}") + B_compute = B if prepared.compute_dtype == prepared.output_dtype else B.to(prepared.compute_dtype) + return _materialize_dense_layout(B_compute, dense_layout) + + +def _run_spmm_coo_rowrun_route(prepared, B, *, timing=False, diagnostics=False, dense_layout="row"): + dense_layout = _normalize_dense_layout(dense_layout) + B = _validate_spmm_coo_route_runtime_inputs(prepared, B, dense_layout) + launch = _resolve_spmm_coo_launch_config(int(B.shape[1]), prepared.nnz) + compute_ms = None + if timing: + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + C = _triton_spmm_coo_rowrun_impl( + prepared.data, + prepared.row, + prepared.col, + B, + prepared.n_rows, + int(B.shape[1]), + block_n=launch["block_n"], + block_nnz=launch["block_nnz"], + output_dtype=prepared.output_dtype, + dense_layout=dense_layout, + seg_starts=prepared.seg_starts, + ) + if timing: + end.record() + torch.cuda.synchronize() + compute_ms = start.elapsed_time(end) + meta = { + "alg": "coo_rowrun", + "display_name": "COORowRun", + "op": prepared.op, + "process_cpu_ms": 0.0, + "process_gpu_ms": 0.0 if timing else None, + "compute_ms": compute_ms, + "dense_layout": dense_layout, + "b_stride": tuple(int(v) for v in B.stride()), + "c_stride": tuple(int(v) for v in C.stride()), + "output_layout": _dense_layout_name(C), + } + if diagnostics: + meta["diagnostics"] = { + "launch_config_scope": "matrix", + "launch_config_count": 1, + "bucket_count": 0, + "long_row_count": 0, + "long_part_count": 0, + "launch_version": "coo_rowrun_v1", + "block_n": launch["block_n"], + "block_nnz": launch["block_nnz"], + "warp_size": launch["heuristic_warp_size"], + "factor": launch["heuristic_factor"], + "grid_m": prepared.n_segs, + "grid_n": triton.cdiv(int(B.shape[1]), launch["block_n"]), + "dense_layout": dense_layout, + "b_stride": tuple(int(v) for v in B.stride()), + "c_stride": tuple(int(v) for v in C.stride()), + "output_layout": _dense_layout_name(C), + } + return C, meta + + +def _run_spmm_coo_atomic_route(prepared, B, *, timing=False, diagnostics=False, dense_layout="row"): + dense_layout = _normalize_dense_layout(dense_layout) + B = _validate_spmm_coo_route_runtime_inputs(prepared, B, dense_layout) + launch = _resolve_spmm_coo_launch_config(int(B.shape[1]), prepared.nnz) + compute_ms = None + if timing: + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + C = _triton_spmm_coo_atomic_impl( + prepared.data, + prepared.row, + prepared.col, + B, + prepared.n_rows, + int(B.shape[1]), + block_n=launch["block_n"], + block_nnz=launch["block_nnz"], + output_dtype=prepared.output_dtype, + dense_layout=dense_layout, + ) + if timing: + end.record() + torch.cuda.synchronize() + compute_ms = start.elapsed_time(end) + meta = { + "alg": "coo_atomic", + "display_name": "COOAtomic", + "op": prepared.op, + "process_cpu_ms": 0.0, + "process_gpu_ms": 0.0 if timing else None, + "compute_ms": compute_ms, + "dense_layout": dense_layout, + "b_stride": tuple(int(v) for v in B.stride()), + "c_stride": tuple(int(v) for v in C.stride()), + "output_layout": _dense_layout_name(C), + } + if diagnostics: + meta["diagnostics"] = { + "launch_config_scope": "matrix", + "launch_config_count": 1, + "bucket_count": 0, + "long_row_count": 0, + "long_part_count": 0, + "launch_version": "coo_atomic_v1", + "block_n": launch["block_n"], + "block_nnz": launch["block_nnz"], + "warp_size": launch["heuristic_warp_size"], + "factor": launch["heuristic_factor"], + "grid_m": prepared.nnz, + "grid_n": int(B.shape[1]), + "dense_layout": dense_layout, + "b_stride": tuple(int(v) for v in B.stride()), + "c_stride": tuple(int(v) for v in C.stride()), + "output_layout": _dense_layout_name(C), + } + return C, meta + + +def _spmm_coo_alg1_build_bucket_descriptors(segs_flat, counts, offsets): + torch.cuda.synchronize() + t0 = time.perf_counter() + counts_cpu = counts.detach().cpu().tolist() + offsets_cpu = offsets.detach().cpu().tolist() + buckets = [] + for bucket_id, count in enumerate(counts_cpu): + offset = int(offsets_cpu[bucket_id]) + count = int(count) + buckets.append( + { + "bucket_id": bucket_id, + "rows": segs_flat.narrow(0, offset, count), + "count": count, + "block_nnz": (32, 128, 512, 2048, 2048)[bucket_id], + } + ) + process_cpu_ms = (time.perf_counter() - t0) * 1000.0 + return buckets, process_cpu_ms + + +def _run_spmm_coo_alg1_route(prepared, B, *, timing=False, diagnostics=False, dense_layout="row"): + if prepared.output_dtype not in (torch.float32, torch.float64): + raise TypeError("spmm_coo_alg1 only supports float32 and float64") + dense_layout = _normalize_dense_layout(dense_layout) + B = _validate_spmm_coo_route_runtime_inputs(prepared, B, dense_layout) + n_dense_cols = int(B.shape[1]) + device = prepared.data.device + bucket_count = 5 + counts = torch.zeros((bucket_count,), dtype=torch.int64, device=device) + offsets = torch.empty_like(counts) + write_counts = torch.zeros_like(counts) + segs_flat = torch.empty((prepared.n_segs,), dtype=torch.int32, device=device) + block_m = 256 + grid = (triton.cdiv(prepared.n_segs, block_m),) + process_gpu_ms = None + if timing: + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + if prepared.n_segs > 0: + _spmm_coo_alg1_process_count_kernel[grid]( + prepared.row_lengths, + counts, + prepared.n_segs, + BLOCK_M=block_m, + num_warps=4, + num_stages=1, + ) + offsets[0] = 0 + if bucket_count > 1: + offsets[1:] = torch.cumsum(counts[:-1], dim=0) + if prepared.n_segs > 0: + _spmm_coo_alg1_process_compact_kernel[grid]( + prepared.row_lengths, + offsets, + write_counts, + segs_flat, + prepared.n_segs, + BLOCK_M=block_m, + num_warps=4, + num_stages=1, + ) + if timing: + end.record() + torch.cuda.synchronize() + process_gpu_ms = start.elapsed_time(end) + else: + torch.cuda.synchronize() + buckets, process_cpu_ms = _spmm_coo_alg1_build_bucket_descriptors(segs_flat, counts, offsets) + + C_compute = _zeros_dense_layout((prepared.n_rows, n_dense_cols), prepared.compute_dtype, device, dense_layout) + launch = _resolve_spmm_coo_launch_config(n_dense_cols, prepared.nnz) + acc_dtype = tl.float64 if prepared.compute_dtype == torch.float64 else tl.float32 + compute_ms = None + if timing: + compute_start = torch.cuda.Event(enable_timing=True) + compute_end = torch.cuda.Event(enable_timing=True) + compute_start.record() + for bucket in buckets: + n_bucket = int(bucket["count"]) + if n_bucket <= 0: + continue + block_nnz = int(bucket["block_nnz"]) + grid_bucket = (n_bucket, triton.cdiv(n_dense_cols, launch["block_n"])) + _spmm_coo_alg1_bucket_real_kernel[grid_bucket]( + prepared.data, + prepared.row, + prepared.col, + B, + C_compute, + prepared.seg_starts, + bucket["rows"], + n_bucket, + n_dense_cols, + B.stride(0), + B.stride(1), + C_compute.stride(0), + C_compute.stride(1), + BLOCK_N=launch["block_n"], + BLOCK_NNZ=block_nnz, + ACC_DTYPE=acc_dtype, + ) + if timing: + compute_end.record() + torch.cuda.synchronize() + compute_ms = compute_start.elapsed_time(compute_end) + if prepared.compute_dtype != prepared.output_dtype: + C = C_compute.to(prepared.output_dtype) + if dense_layout == "col": + C_out = _empty_dense_layout((prepared.n_rows, n_dense_cols), prepared.output_dtype, device, dense_layout) + C_out.copy_(C) + C = C_out + else: + C = C_compute + counts_cpu = [int(bucket["count"]) for bucket in buckets] + meta = { + "alg": "spmm_coo_alg1", + "display_name": "COOAlg1", + "op": prepared.op, + "process_cpu_ms": process_cpu_ms, + "process_gpu_ms": process_gpu_ms, + "compute_ms": compute_ms, + "dense_layout": dense_layout, + "b_stride": tuple(int(v) for v in B.stride()), + "c_stride": tuple(int(v) for v in C.stride()), + "output_layout": _dense_layout_name(C), + } + if diagnostics: + meta["diagnostics"] = { + "launch_config_scope": "bucket", + "launch_config_count": sum(1 for count in counts_cpu if count > 0), + "bucket_count": sum(1 for count in counts_cpu if count > 0), + "bucket_counts": "|".join(str(v) for v in counts_cpu), + "long_row_count": counts_cpu[-1], + "long_part_count": counts_cpu[-1], + "launch_version": "spmm_coo_alg1_v1", + "block_n": launch["block_n"], + "block_nnz": "32|128|512|2048|2048", + "warp_size": launch["heuristic_warp_size"], + "factor": launch["heuristic_factor"], + "grid_m": prepared.n_segs, + "grid_n": triton.cdiv(n_dense_cols, launch["block_n"]), + "dense_layout": dense_layout, + "b_stride": tuple(int(v) for v in B.stride()), + "c_stride": tuple(int(v) for v in C.stride()), + "output_layout": _dense_layout_name(C), + } + return C, meta + + +SPMM_COO_ALGORITHMS = { + "coo_rowrun": SpmmCooAlgorithm( + name="coo_rowrun", + display_name="COORowRun", + supported_ops=tuple(SPMM_COO_OP_NAMES.values()), + supported_dtypes=SUPPORTED_SPMM_VALUE_DTYPES, + run=_run_spmm_coo_rowrun_route, + ), + "coo_atomic": SpmmCooAlgorithm( + name="coo_atomic", + display_name="COOAtomic", + supported_ops=tuple(SPMM_COO_OP_NAMES.values()), + supported_dtypes=SUPPORTED_SPMM_VALUE_DTYPES, + run=_run_spmm_coo_atomic_route, + ), + "spmm_coo_alg1": SpmmCooAlgorithm( + name="spmm_coo_alg1", + display_name="COOAlg1", + supported_ops=tuple(SPMM_COO_OP_NAMES.values()), + supported_dtypes=(torch.float32, torch.float64), + run=_run_spmm_coo_alg1_route, + ), +} + + +def resolve_spmm_coo_algorithm(alg, op, dtype): + token = _normalize_spmm_coo_alg(alg) + if token == "auto": + token = "coo_rowrun" + if token not in SPMM_COO_ALGORITHMS: + supported = ", ".join(sorted(SPMM_COO_ALGORITHMS)) + raise ValueError(f"unsupported COO SpMM algorithm {alg!r}; supported: auto, {supported}") + algorithm = SPMM_COO_ALGORITHMS[token] + op_name = _spmm_coo_op_to_name(op) + if op_name not in algorithm.supported_ops: + raise ValueError(f"algorithm {token!r} does not support op {op_name!r}") + if dtype not in algorithm.supported_dtypes: + raise TypeError(f"algorithm {token!r} does not support dtype {dtype}") + return algorithm + + +def list_spmm_coo_algorithms(op=None, dtype=None): + op_name = None if op is None else _spmm_coo_op_to_name(op) + names = [] + for name, algorithm in SPMM_COO_ALGORITHMS.items(): + if op_name is not None and op_name not in algorithm.supported_ops: + continue + if dtype is not None and dtype not in algorithm.supported_dtypes: + continue + names.append(name) + return tuple(names) + + +def prepare_spmm_coo_route(data, row, col, shape, *, op="non", alg="auto"): + """Prepare matrix-level canonical COO metadata for registered SpMM algorithms.""" + op_code = _normalize_spmm_coo_op(op) + op_name = _spmm_coo_op_to_name(op_code) + data, row, col, shape = _materialize_spmm_coo_op(data, row, col, shape, op_code) + output_dtype = data.dtype + compute_dtype = _spmm_coo_compute_dtype(output_dtype) + data, row, col, shape = _prepare_spmm_coo_matrix(data, row, col, shape) + data_compute = data if compute_dtype == output_dtype else data.to(compute_dtype) + canonical_data, canonical_row, canonical_col = _coalesce_coo_entries(data_compute, row, col, shape) + canonical_data, canonical_row, canonical_col = _sort_coo_lex_inplace( + canonical_data, + canonical_row, + canonical_col, + shape[1], + ) + canonical_row = canonical_row.to(torch.int32) + canonical_col = canonical_col.to(torch.int32) + seg_starts = _seg_starts_from_sorted_rows(canonical_row, int(canonical_data.numel()), canonical_data.device) + if seg_starts is None: + row_lengths = torch.empty((0,), dtype=torch.int32, device=canonical_data.device) + else: + row_lengths = (seg_starts[1:] - seg_starts[:-1]).contiguous() + resolved_alg = _normalize_spmm_coo_alg(alg) + if resolved_alg != "auto": + resolve_spmm_coo_algorithm(resolved_alg, op_name, output_dtype) + return PreparedCooSpmmRoute( + canonical_data, + canonical_row, + canonical_col, + shape, + seg_starts, + row_lengths, + output_dtype, + compute_dtype, + op_name, + resolved_alg, + ) + + +def flagsparse_spmm_coo_run( + prepared, + B, + *, + alg=None, + dense_layout="auto", + return_time=False, + return_meta=False, + timing=False, + diagnostics=False, +): + """Run a registered COO SpMM algorithm with CSR-style timing metadata.""" + if not isinstance(prepared, PreparedCooSpmmRoute): + raise TypeError("prepared must be a PreparedCooSpmmRoute instance") + alg_name = prepared.alg if alg is None else _normalize_spmm_coo_alg(alg) + algorithm = resolve_spmm_coo_algorithm(alg_name, prepared.op, prepared.output_dtype) + dense_layout = _normalize_dense_layout(dense_layout) + start = torch.cuda.Event(enable_timing=True) if (return_time or return_meta) else None + end = torch.cuda.Event(enable_timing=True) if (return_time or return_meta) else None + if start is not None: + torch.cuda.synchronize() + start.record() + C, route_meta = algorithm.run( + prepared, + B, + timing=bool(timing), + diagnostics=bool(diagnostics), + dense_layout=dense_layout, + ) + if end is not None: + end.record() + torch.cuda.synchronize() + gpu_ms = start.elapsed_time(end) + else: + gpu_ms = None + process_cpu_ms = float(route_meta.get("process_cpu_ms", 0.0) or 0.0) + operator_ms = (process_cpu_ms + float(gpu_ms)) if gpu_ms is not None else None + meta = None + if return_meta: + meta = { + "alg": algorithm.name, + "display_name": algorithm.display_name, + "op": prepared.op, + "operator_ms": operator_ms, + "gpu_ms": gpu_ms, + "process_cpu_ms": process_cpu_ms, + "dense_layout": route_meta.get("dense_layout", dense_layout), + "b_stride": route_meta.get("b_stride"), + "c_stride": route_meta.get("c_stride"), + "output_layout": route_meta.get("output_layout"), + } + if timing: + meta["process_gpu_ms"] = route_meta.get("process_gpu_ms") + meta["compute_ms"] = route_meta.get("compute_ms") + if diagnostics and "diagnostics" in route_meta: + meta["diagnostics"] = route_meta["diagnostics"] + if return_time and return_meta: + return C, operator_ms, meta + if return_time: + return C, operator_ms + if return_meta: + return C, meta + return C + + def _run_spmm_coo_canonical_route( canonical_data, canonical_row, diff --git a/tests/test_spmm_coo.py b/tests/test_spmm_coo.py index b951440..356d0f5 100644 --- a/tests/test_spmm_coo.py +++ b/tests/test_spmm_coo.py @@ -40,6 +40,74 @@ CSV_INDEX_DTYPES = [torch.int32, torch.int64] OP_NAMES = tuple(ast_ops.SPMM_COO_OP_NAMES.values()) LAYOUT_NAMES = ("row", "col") +PERF_FIELDS = [ + "matrix", + "dtype", + "index_dtype", + "op", + "layout", + "alg", + "n_rows", + "n_cols", + "nnz", + "dense_cols", + "b_stride", + "c_stride", + "ms", + "gpu_ms", + "process_cpu_ms", + "torch_ms", + "cusparse_ms", + "torch_vs_alg_speedup", + "cusparse_vs_alg_speedup", + "err_vs_torch", + "err_vs_cusparse", + "status", + "reason", + "cusparse_reason", +] +TIMING_FIELDS = ["process_gpu_ms", "compute_ms"] +DIAG_FIELDS = [ + "matrix", + "dtype", + "index_dtype", + "op", + "layout", + "alg", + "launch_config_scope", + "launch_config_count", + "bucket_count", + "long_row_count", + "long_part_count", + "num_warps", + "num_stages", + "block_n", + "block_nnz", + "warp_size", + "factor", + "block_rows", + "block_cols", + "grid_m", + "grid_n", + "launch_version", + "dense_layout", + "b_stride", + "c_stride", + "output_layout", + "bucket_counts", +] +BEST_FIELDS = [ + "matrix", + "dtype", + "index_dtype", + "op", + "layout", + "best_alg", + "best_ms", + "best_gpu_ms", + "best_torch_speedup", + "best_cusparse_speedup", +] TEST_CASES = [ (512, 512, 4096, 16), (1024, 1024, 16384, 32), @@ -83,6 +151,56 @@ def _speedup_ratio(other_ms, triton_ms): return other_ms / triton_ms +def _ratio(numerator, denominator): + if numerator is None or denominator is None or denominator <= 0: + return None + return float(numerator) / float(denominator) + + +def _parse_algs(value, route="rowrun"): + if value is None: + if route == "atomic": + return ["coo_atomic"] + if route == "compare": + return ["coo_rowrun", "coo_atomic"] + return ["auto"] + value = str(value).strip().lower() + if value in ("auto", "all"): + return [value] + allowed = set(ast.SPMM_COO_ALGORITHMS) + aliases = { + "rowrun": "coo_rowrun", + "atomic": "coo_atomic", + "alg1": "spmm_coo_alg1", + "coo_alg1": "spmm_coo_alg1", + } + names = [aliases.get(token.strip().lower(), token.strip().lower()) for token in value.split(",") if token.strip()] + if not names: + raise ValueError("--alg must not be empty") + invalid = [name for name in names if name not in allowed] + if invalid: + raise ValueError( + f"unsupported --alg: {', '.join(invalid)}; allowed: auto,all,{','.join(sorted(allowed))}" + ) + return names + + +def _expand_algs(alg_names, op, dtype): + expanded = [] + for alg in alg_names: + if alg == "all": + expanded.extend(ast.list_spmm_coo_algorithms(op=op, dtype=dtype)) + elif alg == "auto": + expanded.append("auto") + else: + expanded.append(alg) + deduped = [] + for alg in expanded: + if alg not in deduped: + deduped.append(alg) + return deduped + + def _parse_csv_tokens(value, mapping, option_name): tokens = [token.strip().lower() for token in str(value).split(",") if token.strip()] if not tokens: @@ -233,6 +351,50 @@ def _scaled_allclose_error(candidate, reference, value_dtype=None): denom = atol + rtol * torch.abs(reference) return float(torch.max(diff / denom).item()) + +def _error_profile(candidate, reference, dtype): + if candidate is None or reference is None: + return {"global_err": None, "status": "SKIP"} + if candidate.numel() == 0: + return {"global_err": 0.0, "status": "PASS"} + err = _scaled_allclose_error(candidate, reference, dtype) + return {"global_err": err, "status": "PASS" if err <= 1.0 else "FAIL"} + + +def _write_csv(path, rows, fields): + with open(path, "w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fields, extrasaction="ignore") + writer.writeheader() + for row in rows: + writer.writerow({key: ("" if value is None else value) for key, value in row.items()}) + + +def _best_rows(rows): + groups = {} + for row in rows: + if row.get("status") != "PASS" or row.get("ms") is None: + continue + key = (row["matrix"], row["dtype"], row["index_dtype"], row["op"], row["layout"]) + groups.setdefault(key, []).append(row) + best = [] + for (matrix, dtype, index_dtype, op, layout), group in sorted(groups.items()): + selected = min(group, key=lambda item: item["ms"]) + best.append( + { + "matrix": matrix, + "dtype": dtype, + "index_dtype": index_dtype, + "op": op, + "layout": layout, + "best_alg": selected["alg"], + "best_ms": selected["ms"], + "best_gpu_ms": selected["gpu_ms"], + "best_torch_speedup": selected["torch_vs_alg_speedup"], + "best_cusparse_speedup": selected["cusparse_vs_alg_speedup"], + } + ) + return best + def load_mtx_to_coo_torch(file_path, dtype=torch.float32, device=None): if device is None: device = torch.device("cuda" if torch.cuda.is_available() else "cpu") @@ -465,6 +627,318 @@ def _cuda_event_benchmark(op, warmup, iters): return out, start.elapsed_time(end) / count +def _time_coo_algorithm(prepared, B, alg, warmup, iters, timing=False, diagnose=False, layout="row"): + out, gpu_ms = _cuda_event_benchmark( + lambda: ast.flagsparse_spmm_coo_run(prepared, B, alg=alg, dense_layout=layout), + warmup, + iters, + ) + _, meta = ast.flagsparse_spmm_coo_run( + prepared, + B, + alg=alg, + dense_layout=layout, + return_meta=True, + timing=bool(timing), + diagnostics=bool(diagnose), + ) + process_cpu_ms = float(meta.get("process_cpu_ms", 0.0) or 0.0) + row = { + "alg": meta.get("alg", alg), + "ms": process_cpu_ms + gpu_ms, + "gpu_ms": gpu_ms, + "process_cpu_ms": process_cpu_ms, + "process_gpu_ms": None, + "compute_ms": None, + "dense_layout": meta.get("dense_layout", layout), + "b_stride": meta.get("b_stride"), + "c_stride": meta.get("c_stride"), + "output_layout": meta.get("output_layout"), + "diagnostics": meta.get("diagnostics", {}), + "out": out, + } + if timing: + row["process_gpu_ms"] = meta.get("process_gpu_ms") + row["compute_ms"] = meta.get("compute_ms") + if row["process_gpu_ms"] is None: + row["process_gpu_ms"] = 0.0 + if row["compute_ms"] is None: + row["compute_ms"] = gpu_ms + return row + + +def _skip_alg_row( + path, + dtype, + index_dtype_name, + op, + layout, + alg, + shape, + nnz, + dense_cols, + b_stride, + torch_ms, + cusparse_ms, + reason, + timing, + cusparse_reason="", +): + n_rows, n_cols = shape + row = { + "matrix": os.path.basename(path), + "dtype": _dtype_name(dtype), + "index_dtype": index_dtype_name, + "op": op, + "layout": layout, + "alg": alg, + "n_rows": n_rows, + "n_cols": n_cols, + "nnz": int(nnz), + "dense_cols": dense_cols, + "b_stride": b_stride, + "c_stride": "", + "ms": None, + "gpu_ms": None, + "process_cpu_ms": None, + "torch_ms": torch_ms, + "cusparse_ms": cusparse_ms, + "torch_vs_alg_speedup": None, + "cusparse_vs_alg_speedup": None, + "err_vs_torch": None, + "err_vs_cusparse": None, + "status": "SKIP", + "reason": reason, + "cusparse_reason": cusparse_reason or "", + } + if timing: + row["process_gpu_ms"] = None + row["compute_ms"] = None + return row + + +def _time_cusparse_coo(prepared_case, ref_C, dtype, warmup, iters, layout="row"): + if dtype not in (torch.float32, torch.float64, torch.complex64, torch.complex128): + return None, None, "dtype not supported by CuPy/cuSPARSE reference" + try: + import cupy as cp + import cupyx.scipy.sparse as cpx + except Exception as exc: + return None, None, f"CuPy/cuSPARSE unavailable: {exc}" + try: + data_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(prepared_case["cusparse_data"])) + row_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(prepared_case["cusparse_row"].to(torch.int64))) + col_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(prepared_case["cusparse_col"].to(torch.int64))) + B_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(prepared_case["native_B"])) + A_coo = cpx.coo_matrix( + (data_cp, (row_cp, col_cp)), + shape=(prepared_case["n_rows"], prepared_case["n_cols"]), + ) + + def _run(rhs): + return A_coo @ rhs + + try: + out_cp, ms = _cupy_event_benchmark(_run, B_cp, warmup, iters) + reason = "" + except Exception: + if layout != "col": + raise + B_cp = cp.asfortranarray(B_cp) + out_cp, ms = _cupy_event_benchmark(_run, B_cp, warmup, iters) + reason = "used cp.asfortranarray fallback for col-major B" + out = torch.utils.dlpack.from_dlpack(out_cp.toDlpack()) + del ref_C + return out, ms, reason + except Exception as exc: + return None, None, str(exc) + + +def _cupy_event_benchmark(op, arg, warmup, iters): + import cupy as cp + + out = None + for _ in range(max(0, int(warmup))): + out = op(arg) + cp.cuda.runtime.deviceSynchronize() + start = cp.cuda.Event() + end = cp.cuda.Event() + count = max(1, int(iters)) + start.record() + for _ in range(count): + out = op(arg) + end.record() + end.synchronize() + return out, cp.cuda.get_elapsed_time(start, end) / count + + +def run_one_alg_case( + path, + dtype, + index_dtype_name, + index_dtype, + op, + layout, + alg_names, + dense_cols, + warmup, + iters, + run_cusparse, + timing, + diagnose, +): + device = torch.device("cuda") + data, row, col, shape = load_mtx_to_coo_torch(path, dtype=dtype, device=device) + row = row.to(index_dtype) + col = col.to(index_dtype) + n_rows, n_cols = shape + b_rows = n_rows if ast_ops._spmm_coo_op_transposes(op) else n_cols + B = _materialize_dense_layout_for_test( + _build_dense_matrix(b_rows, dense_cols, dtype, device), + layout, + ) + b_stride = _stride_string(B) + case = _prepare_canonical_case(data, row, col, shape, B, op=op, layout=layout) + ref, pytorch_op, _torch_format, pytorch_reason = _build_pytorch_reference( + data, + row, + col, + shape, + B, + prepared=case, + op=op, + layout=layout, + ) + torch_ms = None + try: + _torch_out, torch_ms = _cuda_event_benchmark(pytorch_op, warmup, iters) + except Exception as exc: + pytorch_reason = str(exc) if pytorch_reason is None else f"{pytorch_reason}; timing: {exc}" + + cusparse_out = None + cusparse_ms = None + cusparse_reason = "" + if run_cusparse: + cusparse_out, cusparse_ms, cusparse_reason = _time_cusparse_coo( + case, + ref, + dtype, + warmup, + iters, + layout=layout, + ) + + rows = [] + diag_rows = [] + try: + prepared = ast.prepare_spmm_coo_route(data, row, col, shape, op=op, alg="auto") + except Exception as exc: + for alg in _expand_algs(alg_names, op, dtype): + rows.append( + _skip_alg_row( + path, + dtype, + index_dtype_name, + op, + layout, + alg, + shape, + data.numel(), + dense_cols, + b_stride, + torch_ms, + cusparse_ms, + f"prepare: {exc}", + timing, + cusparse_reason=cusparse_reason, + ) + ) + return rows, diag_rows + + for alg in _expand_algs(alg_names, op, dtype): + try: + ast.resolve_spmm_coo_algorithm(alg, op, dtype) + result = _time_coo_algorithm( + prepared, + B, + alg, + warmup, + iters, + timing=timing, + diagnose=diagnose, + layout=layout, + ) + except (ast.SpmmCooAlgorithmUnavailable, ValueError, TypeError) as exc: + rows.append( + _skip_alg_row( + path, + dtype, + index_dtype_name, + op, + layout, + alg, + shape, + data.numel(), + dense_cols, + b_stride, + torch_ms, + cusparse_ms, + str(exc), + timing, + cusparse_reason=cusparse_reason, + ) + ) + continue + out = result.pop("out") + diagnostics = result.pop("diagnostics") + torch_profile = _error_profile(out, ref, dtype) + cusparse_profile = _error_profile(out, cusparse_out, dtype) + row_out = { + "matrix": os.path.basename(path), + "dtype": _dtype_name(dtype), + "index_dtype": index_dtype_name, + "op": op, + "layout": layout, + "alg": result["alg"], + "n_rows": n_rows, + "n_cols": n_cols, + "nnz": int(data.numel()), + "dense_cols": dense_cols, + "b_stride": _stride_string(B), + "c_stride": _stride_string(out), + "ms": result["ms"], + "gpu_ms": result["gpu_ms"], + "process_cpu_ms": result["process_cpu_ms"], + "torch_ms": torch_ms, + "cusparse_ms": cusparse_ms, + "torch_vs_alg_speedup": _ratio(torch_ms, result["ms"]), + "cusparse_vs_alg_speedup": _ratio(cusparse_ms, result["ms"]), + "err_vs_torch": torch_profile["global_err"], + "err_vs_cusparse": cusparse_profile["global_err"], + "status": torch_profile["status"], + "reason": pytorch_reason or "", + "cusparse_reason": cusparse_reason or "", + } + if timing: + row_out["process_gpu_ms"] = result["process_gpu_ms"] + row_out["compute_ms"] = result["compute_ms"] + rows.append(row_out) + if diagnose: + diag = { + "matrix": os.path.basename(path), + "dtype": _dtype_name(dtype), + "index_dtype": index_dtype_name, + "op": op, + "layout": layout, + "alg": result["alg"], + } + for field in DIAG_FIELDS: + if field not in diag: + diag[field] = diagnostics.get(field) + diag_rows.append(diag) + return rows, diag_rows + + def _prepare_spmm_coo_timing_base(data, row, col, B, shape, op="non", layout="row"): layout = _normalize_layout_name(layout) op_code = ast_ops._normalize_spmm_coo_op(op) @@ -1529,13 +2003,13 @@ def run_all_dtypes_export_csv( op_names=None, layout_names=None, timing=False, + alg_names=None, + diagnose=False, ): - route = _normalize_route(route) - if route == "compare": - raise ValueError("CSV export only supports route='rowrun' or route='atomic'") - selected_route = _selected_route(route) + alg_names = _parse_algs(None, route=route) if alg_names is None else alg_names csv_path = _normalize_csv_path(csv_path) rows = [] + diag_rows = [] value_dtypes = CSV_VALUE_DTYPES if value_dtypes is None else value_dtypes index_dtypes = CSV_INDEX_DTYPES if index_dtypes is None else index_dtypes op_names = ["non"] if op_names is None else op_names @@ -1545,81 +2019,45 @@ def run_all_dtypes_export_csv( for op_name in op_names: for layout_name in layout_names: print("=" * 150) - _print_spmm_coo_mtx_header(value_dtype, index_dtype, route, layout=layout_name, timing=timing) - results = run_mtx_batch( - paths, - value_dtype=value_dtype, - index_dtype=index_dtype, - warmup=warmup, - iters=iters, - run_cusparse=run_cusparse, - n_dense_cols=n_dense_cols, - block_n=block_n, - block_nnz=block_nnz, - route=selected_route, - op=op_name, - layout=layout_name, - timing=timing, - on_result=lambda entry: _print_spmm_coo_mtx_row(entry, timing=timing), + print( + f"COO SpMM registry | dtype={_dtype_name(value_dtype)} index={_dtype_name(index_dtype)} " + f"op={op_name} layout={layout_name} alg={','.join(alg_names)}" ) - print("-" * (226 if timing else 202)) - for entry in results: - n_rows, n_cols = entry["shape"] - status = entry.get("status") - if status == "UNKNOWN": - status = ( - "PASS" - if (entry.get("triton_ok_pt") or entry.get("triton_ok_cu")) - else "FAIL" + for path in paths: + case_rows, case_diag_rows = run_one_alg_case( + path, + value_dtype, + _dtype_name(index_dtype), + index_dtype, + op_name, + layout_name, + alg_names, + n_dense_cols, + warmup, + iters, + run_cusparse, + timing, + diagnose, + ) + rows.extend(case_rows) + diag_rows.extend(case_diag_rows) + for row in case_rows: + print( + f"{row['matrix']:<32} {row['alg']:<16} {row['status']:<5} " + f"ms={_fmt_ms(row['ms'])} torch={_fmt_ms(row['torch_ms'])} " + f"err={_fmt_err(row['err_vs_torch'])} reason={row.get('reason') or row.get('cusparse_reason') or ''}" ) - rows.append({ - "matrix": os.path.basename(entry["path"]), - "op": entry.get("op", op_name), - "layout": entry.get("layout", layout_name), - "value_dtype": _dtype_name(value_dtype), - "index_dtype": _dtype_name(index_dtype), - "n_rows": n_rows, - "n_cols": n_cols, - "nnz": entry["nnz"], - "b_stride": entry.get("b_stride"), - "c_stride": entry.get("c_stride"), - "triton_ms": entry.get("triton_ms"), - "triton_gpu_ms": entry.get("triton_gpu_ms"), - "process_cpu_ms": entry.get("process_cpu_ms"), - "process_gpu_ms": entry.get("process_gpu_ms"), - "compute_ms": entry.get("compute_ms"), - "cusparse_ms": entry.get("cusparse_ms"), - "pytorch_ms": entry.get("pytorch_ms"), - "triton_speedup_vs_cusparse": _speedup_ratio( - entry.get("cusparse_ms"), entry.get("triton_ms") - ), - "triton_speedup_vs_pytorch": _speedup_ratio( - entry.get("pytorch_ms"), entry.get("triton_ms") - ), - "pt_status": _status_label(entry.get("triton_ok_pt")), - "cu_status": _status_label(entry.get("triton_ok_cu")), - "status": status, - "err_pt": entry.get("err_pt"), - "err_cu": entry.get("err_cu"), - "error": entry.get("error"), - "cusparse_reason": entry.get("cusparse_reason"), - }) - fieldnames = [ - "matrix", "op", "layout", "value_dtype", "index_dtype", "n_rows", "n_cols", "nnz", - "b_stride", "c_stride", - "triton_ms", "triton_gpu_ms", "process_cpu_ms", - "cusparse_ms", "pytorch_ms", - "triton_speedup_vs_cusparse", "triton_speedup_vs_pytorch", - "pt_status", "cu_status", "status", "err_pt", "err_cu", "error", "cusparse_reason", - ] + fieldnames = list(PERF_FIELDS) if timing: - insert_at = fieldnames.index("cusparse_ms") - fieldnames[insert_at:insert_at] = ["process_gpu_ms", "compute_ms"] - with open(csv_path, "w", newline="", encoding="utf-8") as handle: - writer = csv.DictWriter(handle, fieldnames=fieldnames, extrasaction="ignore") - writer.writeheader() - for row in rows: - writer.writerow({key: ("" if value is None else value) for key, value in row.items()}) + fieldnames += TIMING_FIELDS + _write_csv(csv_path, rows, fieldnames) + best_path = csv_path[:-4] + ".best.csv" + _write_csv(best_path, _best_rows(rows), BEST_FIELDS) + if diagnose: + diag_path = csv_path[:-4] + ".diagnose.csv" + _write_csv(diag_path, diag_rows, DIAG_FIELDS) + print(f"Wrote {len(diag_rows)} diagnose rows to {diag_path}") + print(f"Wrote best rows to {best_path}") print(f"Wrote {len(rows)} rows to {csv_path}") def run_api_validation_checks(): @@ -2043,7 +2481,9 @@ def main(): parser.add_argument("--block-n", type=int, default=DEFAULT_BLOCK_N, help="Output column tile override (default: auto from dense-column heuristic)") parser.add_argument("--block-nnz", type=int, default=DEFAULT_BLOCK_NNZ, help="COO nnz tile width override (default: 256)") parser.add_argument("--route", default="rowrun", choices=["rowrun", "atomic", "compare"], help="Native COO route to benchmark/test (default: rowrun)") + parser.add_argument("--alg", default=None, help="COO SpMM algorithm: auto, all, or comma-separated registered names") parser.add_argument("--timing", action="store_true", help="Add process_gpu_ms/compute_ms split timing columns") + parser.add_argument("--diagnose", action="store_true", help="Write algorithm diagnostics to .diagnose.csv") parser.add_argument("--warmup", type=int, default=10, help="Warmup runs") parser.add_argument("--iters", type=int, default=50, help="Timing iterations") parser.add_argument("--no-cusparse", action="store_true", help="Skip cuSPARSE baseline") @@ -2070,6 +2510,7 @@ def main(): try: op_names = _parse_op_names(args.op) layout_names = _layout_names(args.layout) + alg_names = _parse_algs(args.alg, route=args.route) except ValueError as exc: parser.error(str(exc)) @@ -2108,7 +2549,7 @@ def main(): if not paths: print("No .mtx files found. Specify files or a directory.") return - if args.route == "compare": + if args.route == "compare" and args.alg is None: print("CSV export only supports --route rowrun or --route atomic.") return csv_path = _normalize_csv_path(args.csv) @@ -2122,7 +2563,7 @@ def main(): csv_layout_names = _layout_names(args.layout) except ValueError as exc: parser.error(str(exc)) - print(f"GPU: {torch.cuda.get_device_name(0)} | Files: {len(paths)} | DenseN: {args.dense_cols} | Route: {args.route} | CSV: {csv_path}") + print(f"GPU: {torch.cuda.get_device_name(0)} | Files: {len(paths)} | DenseN: {args.dense_cols} | Alg: {args.alg or args.route} | CSV: {csv_path}") print(f"dtypes: {args.dtypes} | index_dtypes: {args.index_dtypes} | ops: {args.op} | layouts: {args.layout}") run_all_dtypes_export_csv( paths, @@ -2139,6 +2580,8 @@ def main(): op_names=csv_op_names, layout_names=csv_layout_names, timing=args.timing, + alg_names=alg_names, + diagnose=args.diagnose, ) return @@ -2149,29 +2592,56 @@ def main(): print( f"dtype: {args.dtype} index_dtype: {args.index_dtype} dense_cols: {args.dense_cols} " f"op: {args.op} layout: {args.layout} warmup: {args.warmup} iters: {args.iters} block_n: {_fmt_launch_value(args.block_n)} " - f"block_nnz: {_fmt_launch_value(args.block_nnz)} route: {args.route}" + f"block_nnz: {_fmt_launch_value(args.block_nnz)} alg: {args.alg or args.route}" ) print() + if args.route == "compare" and args.alg is None: + for op_name in op_names: + for layout_name in layout_names: + results = run_mtx_batch( + paths, + value_dtype=value_dtype, + index_dtype=index_dtype, + warmup=args.warmup, + iters=args.iters, + run_cusparse=not args.no_cusparse, + n_dense_cols=args.dense_cols, + block_n=args.block_n, + block_nnz=args.block_nnz, + route=args.route, + op=op_name, + layout=layout_name, + timing=args.timing, + ) + print_mtx_results(results, value_dtype, index_dtype, route=args.route, layout=layout_name, timing=args.timing) + print_compare_results(results, value_dtype, index_dtype) + return + for op_name in op_names: for layout_name in layout_names: - results = run_mtx_batch( - paths, - value_dtype=value_dtype, - index_dtype=index_dtype, - warmup=args.warmup, - iters=args.iters, - run_cusparse=not args.no_cusparse, - n_dense_cols=args.dense_cols, - block_n=args.block_n, - block_nnz=args.block_nnz, - route=args.route, - op=op_name, - layout=layout_name, - timing=args.timing, - ) - print_mtx_results(results, value_dtype, index_dtype, route=args.route, layout=layout_name, timing=args.timing) - if args.route == "compare": - print_compare_results(results, value_dtype, index_dtype) + for path in paths: + rows, _diag_rows = run_one_alg_case( + path, + value_dtype, + _dtype_name(index_dtype), + index_dtype, + op_name, + layout_name, + alg_names, + args.dense_cols, + args.warmup, + args.iters, + not args.no_cusparse, + args.timing, + args.diagnose, + ) + for row in rows: + print( + f"{row['matrix']:<32} {row['alg']:<16} {row['status']:<5} " + f"ms={_fmt_ms(row['ms'])} gpu={_fmt_ms(row['gpu_ms'])} " + f"torch={_fmt_ms(row['torch_ms'])} err={_fmt_err(row['err_vs_torch'])} " + f"reason={row.get('reason') or row.get('cusparse_reason') or ''}" + ) if __name__ == "__main__": From 2378bff4d1f671638ebab541032103e4d8de641a Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Thu, 2 Jul 2026 18:46:28 +0800 Subject: [PATCH 04/13] spmm_coo_opt --- tests/test_spmm_coo.py | 52 +++++++++++++++++++++++++++++++++++------- 1 file changed, 44 insertions(+), 8 deletions(-) diff --git a/tests/test_spmm_coo.py b/tests/test_spmm_coo.py index 356d0f5..51b8ca2 100644 --- a/tests/test_spmm_coo.py +++ b/tests/test_spmm_coo.py @@ -786,18 +786,36 @@ def run_one_alg_case( run_cusparse, timing, diagnose, + progress=False, ): + matrix_name = os.path.basename(path) + + def _start(stage): + if progress: + print(f"[{matrix_name}] {stage} ...", flush=True) + return time.perf_counter() + + def _done(stage, started_at, extra=""): + if progress: + suffix = f" {extra}" if extra else "" + print(f"[{matrix_name}] {stage} done in {(time.perf_counter() - started_at) * 1000.0:.2f} ms{suffix}", flush=True) + device = torch.device("cuda") + stage_t0 = _start("load mtx") data, row, col, shape = load_mtx_to_coo_torch(path, dtype=dtype, device=device) row = row.to(index_dtype) col = col.to(index_dtype) n_rows, n_cols = shape + _done("load mtx", stage_t0, f"shape=({n_rows},{n_cols}) nnz={int(data.numel())}") b_rows = n_rows if ast_ops._spmm_coo_op_transposes(op) else n_cols + stage_t0 = _start(f"build dense B shape=({b_rows},{dense_cols})") B = _materialize_dense_layout_for_test( _build_dense_matrix(b_rows, dense_cols, dtype, device), layout, ) b_stride = _stride_string(B) + _done("build dense B", stage_t0, f"stride={b_stride}") + stage_t0 = _start("prepare canonical reference") case = _prepare_canonical_case(data, row, col, shape, B, op=op, layout=layout) ref, pytorch_op, _torch_format, pytorch_reason = _build_pytorch_reference( data, @@ -809,16 +827,21 @@ def run_one_alg_case( op=op, layout=layout, ) + _done("prepare canonical reference", stage_t0) torch_ms = None try: + stage_t0 = _start("time PyTorch COO reference") _torch_out, torch_ms = _cuda_event_benchmark(pytorch_op, warmup, iters) + _done("time PyTorch COO reference", stage_t0, f"ms={_fmt_ms(torch_ms)}") except Exception as exc: pytorch_reason = str(exc) if pytorch_reason is None else f"{pytorch_reason}; timing: {exc}" + _done("time PyTorch COO reference", stage_t0, f"failed={exc}") cusparse_out = None cusparse_ms = None cusparse_reason = "" if run_cusparse: + stage_t0 = _start("time CuPy/cuSPARSE COO reference") cusparse_out, cusparse_ms, cusparse_reason = _time_cusparse_coo( case, ref, @@ -827,12 +850,16 @@ def run_one_alg_case( iters, layout=layout, ) + _done("time CuPy/cuSPARSE COO reference", stage_t0, f"ms={_fmt_ms(cusparse_ms)} reason={cusparse_reason or ''}") rows = [] diag_rows = [] try: + stage_t0 = _start("prepare COO route") prepared = ast.prepare_spmm_coo_route(data, row, col, shape, op=op, alg="auto") + _done("prepare COO route", stage_t0, f"row_runs={prepared.n_segs}") except Exception as exc: + _done("prepare COO route", stage_t0, f"failed={exc}") for alg in _expand_algs(alg_names, op, dtype): rows.append( _skip_alg_row( @@ -856,8 +883,10 @@ def run_one_alg_case( return rows, diag_rows for alg in _expand_algs(alg_names, op, dtype): + stage_t0 = None try: ast.resolve_spmm_coo_algorithm(alg, op, dtype) + stage_t0 = _start(f"run {alg}") result = _time_coo_algorithm( prepared, B, @@ -868,7 +897,10 @@ def run_one_alg_case( diagnose=diagnose, layout=layout, ) + _done(f"run {alg}", stage_t0, f"ms={_fmt_ms(result['ms'])}") except (ast.SpmmCooAlgorithmUnavailable, ValueError, TypeError) as exc: + if stage_t0 is not None: + _done(f"run {alg}", stage_t0, f"skip={exc}") rows.append( _skip_alg_row( path, @@ -2010,6 +2042,11 @@ def run_all_dtypes_export_csv( csv_path = _normalize_csv_path(csv_path) rows = [] diag_rows = [] + fieldnames = list(PERF_FIELDS) + if timing: + fieldnames += TIMING_FIELDS + best_path = csv_path[:-4] + ".best.csv" + diag_path = csv_path[:-4] + ".diagnose.csv" value_dtypes = CSV_VALUE_DTYPES if value_dtypes is None else value_dtypes index_dtypes = CSV_INDEX_DTYPES if index_dtypes is None else index_dtypes op_names = ["non"] if op_names is None else op_names @@ -2021,9 +2058,11 @@ def run_all_dtypes_export_csv( print("=" * 150) print( f"COO SpMM registry | dtype={_dtype_name(value_dtype)} index={_dtype_name(index_dtype)} " - f"op={op_name} layout={layout_name} alg={','.join(alg_names)}" + f"op={op_name} layout={layout_name} alg={','.join(alg_names)}", + flush=True, ) - for path in paths: + for matrix_idx, path in enumerate(paths, start=1): + print(f"[{matrix_idx}/{len(paths)}] [{os.path.basename(path)}] START", flush=True) case_rows, case_diag_rows = run_one_alg_case( path, value_dtype, @@ -2038,6 +2077,7 @@ def run_all_dtypes_export_csv( run_cusparse, timing, diagnose, + progress=True, ) rows.extend(case_rows) diag_rows.extend(case_diag_rows) @@ -2045,16 +2085,12 @@ def run_all_dtypes_export_csv( print( f"{row['matrix']:<32} {row['alg']:<16} {row['status']:<5} " f"ms={_fmt_ms(row['ms'])} torch={_fmt_ms(row['torch_ms'])} " - f"err={_fmt_err(row['err_vs_torch'])} reason={row.get('reason') or row.get('cusparse_reason') or ''}" + f"err={_fmt_err(row['err_vs_torch'])} reason={row.get('reason') or row.get('cusparse_reason') or ''}", + flush=True, ) - fieldnames = list(PERF_FIELDS) - if timing: - fieldnames += TIMING_FIELDS _write_csv(csv_path, rows, fieldnames) - best_path = csv_path[:-4] + ".best.csv" _write_csv(best_path, _best_rows(rows), BEST_FIELDS) if diagnose: - diag_path = csv_path[:-4] + ".diagnose.csv" _write_csv(diag_path, diag_rows, DIAG_FIELDS) print(f"Wrote {len(diag_rows)} diagnose rows to {diag_path}") print(f"Wrote best rows to {best_path}") From 030f737bba9753cb069539ce9d044ca89cf71cef Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Thu, 2 Jul 2026 18:56:42 +0800 Subject: [PATCH 05/13] spmm_coo_opt --- src/flagsparse/sparse_operations/spmm_coo.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/src/flagsparse/sparse_operations/spmm_coo.py b/src/flagsparse/sparse_operations/spmm_coo.py index 6c63602..a83b579 100644 --- a/src/flagsparse/sparse_operations/spmm_coo.py +++ b/src/flagsparse/sparse_operations/spmm_coo.py @@ -1196,7 +1196,9 @@ def _spmm_coo_alg1_build_bucket_descriptors(segs_flat, counts, offsets): "bucket_id": bucket_id, "rows": segs_flat.narrow(0, offset, count), "count": count, - "block_nnz": (32, 128, 512, 2048, 2048)[bucket_id], + # Keep large row-run buckets logically separate, but cap the + # static inner loop to avoid enormous Triton specializations. + "block_nnz": (32, 64, 128, 128, 256)[bucket_id], } ) process_cpu_ms = (time.perf_counter() - t0) * 1000.0 @@ -1320,7 +1322,7 @@ def _run_spmm_coo_alg1_route(prepared, B, *, timing=False, diagnostics=False, de "long_part_count": counts_cpu[-1], "launch_version": "spmm_coo_alg1_v1", "block_n": launch["block_n"], - "block_nnz": "32|128|512|2048|2048", + "block_nnz": "32|64|128|128|256", "warp_size": launch["heuristic_warp_size"], "factor": launch["heuristic_factor"], "grid_m": prepared.n_segs, From 900fc549bee910e1334e301abb06e0b6938dfab5 Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Thu, 9 Jul 2026 21:15:12 +0800 Subject: [PATCH 06/13] spmv_bsr --- conf/operators.yaml | 17 + ops_support.py | 6 + pytest.ini | 1 + run_flagsparse_pytest.py | 17 + src/flagsparse/__init__.py | 6 + src/flagsparse/sparse_operations/__init__.py | 4 + src/flagsparse/sparse_operations/spmv_bsr.py | 577 ++++++++++++++ tests/ci/test_cli_help.py | 2 + tests/ci/test_package_smoke.py | 4 + tests/ci/test_public_api.py | 4 + tests/ci/test_runtime_policies.py | 82 ++ tests/pytest/test_spmv_bsr_accuracy.py | 238 ++++++ tests/test_spmv_bsr.py | 797 +++++++++++++++++++ tools/ci/run_gpu_benchmark.py | 11 + 14 files changed, 1766 insertions(+) create mode 100644 src/flagsparse/sparse_operations/spmv_bsr.py create mode 100644 tests/pytest/test_spmv_bsr_accuracy.py create mode 100644 tests/test_spmv_bsr.py diff --git a/conf/operators.yaml b/conf/operators.yaml index cef3241..b366035 100644 --- a/conf/operators.yaml +++ b/conf/operators.yaml @@ -80,6 +80,23 @@ ops: stages: - beta: "1.0" + - id: spmv_bsr + description: | + Computes sparse matrix-vector multiplication for a BSR matrix. + for: + - flagsparse_spmv_bsr + - prepare_spmv_bsr + labels: + - flagsparse + - sparse + - bsr + - triton + - public-api + kind: + - SparseLinearAlg + stages: + - beta: "1.0" + - id: spmv_coo_tocsr description: | Computes COO SpMV through a COO-to-CSR preparation path. diff --git a/ops_support.py b/ops_support.py index 2fd01e3..8f4099a 100644 --- a/ops_support.py +++ b/ops_support.py @@ -233,6 +233,11 @@ def registry(modules: dict[str, SourceModule]) -> tuple[ApiSpec, ...]: if "spmv_csc" in modules else ("non", "trans", "conj") ) + spmv_bsr_ops = ( + op_names(modules["spmv_bsr"], "SPMV_BSR_SUPPORTED_OP_NAMES") + if "spmv_bsr" in modules + else ("non",) + ) spmm_values = ( normalize_dtype_values(modules["spmm_csr"].get("SUPPORTED_SPMM_VALUE_DTYPES")) if "spmm_csr" in modules @@ -270,6 +275,7 @@ def registry(modules: dict[str, SourceModule]) -> tuple[ApiSpec, ...]: ApiSpec("spmv", "flagsparse_spmv_csr", "spmv_csr", "CSR", "triton", value_const="SUPPORTED_SPMV_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmv_ops, notes="op supports non/trans/conj; conj on real dtypes is transpose-equivalent"), ApiSpec("spmv", "flagsparse_spmv_coo", "spmv_coo", "COO", "triton", value_const="SUPPORTED_SPMV_COO_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmv_coo_ops, notes="COO path stores canonical row/col tensors and supports non/trans/conj"), ApiSpec("spmv", "flagsparse_spmv_csc", "spmv_csc", "CSC", "triton", value_const="SUPPORTED_SPMV_CSC_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmv_csc_ops, notes="native CSC path supports non/trans/conj without CSR/COO conversion"), + ApiSpec("spmv", "flagsparse_spmv_bsr", "spmv_bsr", "BSR", "triton", value_const="SUPPORTED_SPMV_BSR_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmv_bsr_ops, notes="native BSR v1 supports non; trans/conj are reserved but unsupported"), ApiSpec("spmv", "flagsparse_spmv_coo_tocsr", "spmv_csr", "COO->CSR", "triton", value_const="SUPPORTED_SPMV_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=("non",), notes="COO input is converted to CSR before compute"), ApiSpec("spmm", "flagsparse_spmm_csr", "spmm_csr", "CSR", "triton", value_const="SUPPORTED_SPMM_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmm_ops, notes="op supports non/trans/conj; conj on real dtypes is transpose-equivalent"), ApiSpec("spmm", "flagsparse_spmm_csr_opt", "spmm_csr", "CSR", "triton_opt", values=("float32", "float64"), index_const="SUPPORTED_INDEX_DTYPES", ops=("non",), notes="bucketed opt path only supports float32/float64"), diff --git a/pytest.ini b/pytest.ini index 9c4910e..bb035fc 100644 --- a/pytest.ini +++ b/pytest.ini @@ -12,6 +12,7 @@ markers = spmv_csr: CSR SpMV accuracy (tests/pytest) spmv_coo: COO SpMV accuracy (tests/pytest) spmv_csc: CSC SpMV accuracy (tests/pytest) + spmv_bsr: BSR SpMV accuracy (tests/pytest) spmv_coo_tocsr: COO-to-CSR SpMV accuracy (tests/pytest) spmm_csr: CSR SpMM accuracy (tests/pytest) spmm_csr_opt: optimized CSR SpMM accuracy (tests/pytest) diff --git a/run_flagsparse_pytest.py b/run_flagsparse_pytest.py index 47b1346..e1a222d 100644 --- a/run_flagsparse_pytest.py +++ b/run_flagsparse_pytest.py @@ -87,6 +87,8 @@ "pytorch_ms", "cusparse_ms", "cupy_ms", + "csc_ms", + "bsr_ms", "base_ms", "alg1_ms", "alg2_ms", @@ -101,6 +103,10 @@ ("triton_speedup_vs_pytorch", "pytorch_ms", "triton_ms"), ("triton_speedup_vs_cusparse", "cusparse_ms", "triton_ms"), ("triton_speedup_vs_cupy", "cupy_ms", "triton_ms"), + ("csc_speedup_vs_pytorch", "pytorch_ms", "csc_ms"), + ("csc_speedup_vs_cusparse", "cusparse_ms", "csc_ms"), + ("bsr_speedup_vs_pytorch", "pytorch_ms", "bsr_ms"), + ("bsr_speedup_vs_cusparse", "cusparse_ms", "bsr_ms"), ("opt_speedup_vs_pytorch", "pytorch_ms", "opt_ms"), ("opt_speedup_vs_cusparse", "cusparse_ms", "opt_ms"), ("opt_vs_base", "base_ms", "opt_ms"), @@ -172,6 +178,16 @@ class OperatorTestConfig: "--iters", "{iters}", ), + "spmv_bsr": ( + "tests/test_spmv_bsr.py", + "{input}", + "--csv-bsr", + "{csv}", + "--warmup", + "{warmup}", + "--iters", + "{iters}", + ), "spmv_coo_tocsr": ( "tests/test_spmv_coo.py", "{input}", @@ -315,6 +331,7 @@ class OperatorTestConfig: "spmv_csr": OperatorTestConfig("spmv_csr", PERFORMANCE_COMMANDS["spmv_csr"]), "spmv_coo": OperatorTestConfig("spmv_coo", PERFORMANCE_COMMANDS["spmv_coo"]), "spmv_csc": OperatorTestConfig("spmv_csc", PERFORMANCE_COMMANDS["spmv_csc"]), + "spmv_bsr": OperatorTestConfig("spmv_bsr", PERFORMANCE_COMMANDS["spmv_bsr"]), "spmv_coo_tocsr": OperatorTestConfig( "spmv_coo_tocsr", PERFORMANCE_COMMANDS["spmv_coo_tocsr"] ), diff --git a/src/flagsparse/__init__.py b/src/flagsparse/__init__.py index bfcd42d..9714dcd 100644 --- a/src/flagsparse/__init__.py +++ b/src/flagsparse/__init__.py @@ -17,6 +17,7 @@ "comprehensive_scatter_test", "PreparedCoo", "PreparedCooSpmmRoute", + "PreparedBsrSpmv", "PreparedCscSpmv", "PreparedAlphaSpmmAlg1", "PreparedCsrSpmv", @@ -40,6 +41,7 @@ "prepare_spmm_csr_opt_alg2", "prepare_spmm_csr_opt_alg2_preprocess", "prepare_spmv_coo", + "prepare_spmv_bsr", "prepare_spmv_coo_tocsr", "prepare_spmv_csc", "flagsparse_spmv_csr", @@ -54,6 +56,7 @@ "alpha_spmm_alg1_tle_opt_unavailable_reason", "alpha_spmm_alg1_tle_opt2_unavailable_reason", "flagsparse_spmv_coo", + "flagsparse_spmv_bsr", "flagsparse_spmv_coo_tocsr", "flagsparse_spmv_csc", "flagsparse_spmm_csr_opt", @@ -148,6 +151,7 @@ "comprehensive_scatter_test", "PreparedCoo", "PreparedCooSpmmRoute", + "PreparedBsrSpmv", "PreparedCscSpmv", "PreparedAlphaSpmmAlg1", "PreparedCsrSpmv", @@ -171,6 +175,7 @@ "prepare_spmm_csr_opt_alg2", "prepare_spmm_csr_opt_alg2_preprocess", "prepare_spmv_coo", + "prepare_spmv_bsr", "prepare_spmv_coo_tocsr", "prepare_spmv_csc", "flagsparse_spmv_csr", @@ -185,6 +190,7 @@ "alpha_spmm_alg1_tle_opt_unavailable_reason", "alpha_spmm_alg1_tle_opt2_unavailable_reason", "flagsparse_spmv_coo", + "flagsparse_spmv_bsr", "flagsparse_spmv_coo_tocsr", "flagsparse_spmv_csc", "flagsparse_spmm_csr_opt", diff --git a/src/flagsparse/sparse_operations/__init__.py b/src/flagsparse/sparse_operations/__init__.py index d097a8b..bed39a2 100644 --- a/src/flagsparse/sparse_operations/__init__.py +++ b/src/flagsparse/sparse_operations/__init__.py @@ -74,6 +74,7 @@ prepare_spmm_csr_opt_alg2_preprocess, ) from .spmv_coo import PreparedCoo, flagsparse_spmv_coo, prepare_spmv_coo +from .spmv_bsr import PreparedBsrSpmv, flagsparse_spmv_bsr, prepare_spmv_bsr from .spmv_csc import PreparedCscSpmv, flagsparse_spmv_csc, prepare_spmv_csc from .spmv_csr import ( PreparedCsrSpmv, @@ -122,6 +123,7 @@ "PreparedCoo", "PreparedCooSpmmRoute", "PreparedAlphaSpmmAlg1", + "PreparedBsrSpmv", "PreparedCsrSpmv", "PreparedCscSpmv", "PreparedCsrSpmmOpt", @@ -174,6 +176,7 @@ "flagsparse_spmm_csr_opt_alg2", "flagsparse_spmm_csr_opt_alg2_preprocess", "flagsparse_spmv_coo", + "flagsparse_spmv_bsr", "flagsparse_spmv_coo_tocsr", "flagsparse_spmv_csc", "flagsparse_spmv_csr", @@ -220,6 +223,7 @@ "prepare_spmm_csr_opt_alg2", "prepare_spmm_csr_opt_alg2_preprocess", "prepare_spmv_coo", + "prepare_spmv_bsr", "prepare_spmv_coo_tocsr", "prepare_spmv_csc", "prepare_spmv_csr", diff --git a/src/flagsparse/sparse_operations/spmv_bsr.py b/src/flagsparse/sparse_operations/spmv_bsr.py new file mode 100644 index 0000000..3d3785d --- /dev/null +++ b/src/flagsparse/sparse_operations/spmv_bsr.py @@ -0,0 +1,577 @@ +"""Native BSR SpMV kernels and public helpers.""" + +from ._common import * + +import triton +import triton.language as tl + + +SUPPORTED_SPMV_BSR_VALUE_DTYPES = ( + torch.float32, + torch.float64, + torch.complex64, + torch.complex128, +) + +SPMV_BSR_OP_NON = 0 +SPMV_BSR_OP_TRANS = 1 +SPMV_BSR_OP_CONJ_TRANS = 2 +SPMV_BSR_OP_NAMES = { + SPMV_BSR_OP_NON: "non", + SPMV_BSR_OP_TRANS: "trans", + SPMV_BSR_OP_CONJ_TRANS: "conj", +} +SPMV_BSR_SUPPORTED_OP_NAMES = ("non",) +_SPMV_BSR_OP_NAME_TO_CODE = { + name: code for code, name in SPMV_BSR_OP_NAMES.items() +} + + +def _spmv_bsr_dtype_error_message(): + return "BSR SpMV supports float32, float64, complex64, and complex128" + + +def _normalize_spmv_bsr_op(op=None, transpose=False): + if op is None: + return SPMV_BSR_OP_TRANS if bool(transpose) else SPMV_BSR_OP_NON + if isinstance(op, str): + token = op.strip().lower() + if token not in _SPMV_BSR_OP_NAME_TO_CODE: + raise ValueError("op must be one of: 0=non, 1=trans, 2=conj") + return _SPMV_BSR_OP_NAME_TO_CODE[token] + try: + op_code = int(op) + except (TypeError, ValueError) as exc: + raise ValueError("op must be one of: 0=non, 1=trans, 2=conj") from exc + if op_code not in SPMV_BSR_OP_NAMES: + raise ValueError("op must be one of: 0=non, 1=trans, 2=conj") + return op_code + + +def _spmv_bsr_op_to_name(op): + return SPMV_BSR_OP_NAMES[_normalize_spmv_bsr_op(op)] + + +def _spmv_bsr_op_transposes(op): + return _normalize_spmv_bsr_op(op) in ( + SPMV_BSR_OP_TRANS, + SPMV_BSR_OP_CONJ_TRANS, + ) + + +def _ensure_spmv_bsr_supported_op(op_code): + if op_code != SPMV_BSR_OP_NON: + raise NotImplementedError( + "BSR SpMV v1 only supports op='non'; trans/conj are reserved for a future native BSR kernel" + ) + + +def _normalize_spmv_bsr_index_fallback_policy(index_fallback_policy): + policy = str(index_fallback_policy).lower() + if policy not in ("auto", "strict"): + raise ValueError("index_fallback_policy must be 'auto' or 'strict'") + return policy + + +class PreparedBsrSpmv: + """Prepared BSR metadata for repeated SpMV calls.""" + + __slots__ = ( + "data", + "kernel_indices", + "kernel_indptr", + "shape", + "n_rows", + "n_cols", + "block_dim", + "n_block_rows", + "n_block_cols", + "nnzb", + "stored_nnz", + "block_row_lengths", + "max_block_row_nnz", + "block_nnz", + "max_segments", + "op", + "transpose", + "index_fallback_policy", + "index_fallback_applied", + "index_fallback_reason", + ) + + def __init__( + self, + data, + kernel_indices, + kernel_indptr, + shape, + block_dim, + n_block_rows, + n_block_cols, + block_nnz, + max_segments, + max_block_row_nnz, + block_row_lengths=None, + op=None, + transpose=False, + index_fallback_policy="auto", + index_fallback_applied=False, + index_fallback_reason=None, + ): + self.data = data + self.kernel_indices = kernel_indices + self.kernel_indptr = kernel_indptr + self.shape = (int(shape[0]), int(shape[1])) + self.n_rows = int(shape[0]) + self.n_cols = int(shape[1]) + self.block_dim = int(block_dim) + self.n_block_rows = int(n_block_rows) + self.n_block_cols = int(n_block_cols) + self.nnzb = int(data.shape[0]) + self.stored_nnz = int(data.numel()) + if block_row_lengths is None: + block_row_lengths = kernel_indptr[1:] - kernel_indptr[:-1] + self.block_row_lengths = block_row_lengths + self.max_block_row_nnz = int(max_block_row_nnz) + self.block_nnz = int(block_nnz) + self.max_segments = int(max_segments) + self.op = _normalize_spmv_bsr_op(op, transpose=transpose) + self.transpose = _spmv_bsr_op_transposes(self.op) + self.index_fallback_policy = str(index_fallback_policy).lower() + self.index_fallback_applied = bool(index_fallback_applied) + self.index_fallback_reason = index_fallback_reason + + +@triton.jit +def _spmv_bsr_non_real_kernel( + data_ptr, + indices_ptr, + indptr_ptr, + x_ptr, + y_ptr, + n_rows, + n_cols, + n_block_rows, + BLOCK_DIM: tl.constexpr, + BLOCK_NNZ: tl.constexpr, + SEG: tl.constexpr, +): + brow = tl.program_id(0) + inner_row = tl.program_id(1) + if brow >= n_block_rows: + return + row = brow * BLOCK_DIM + inner_row + if row >= n_rows: + return + start = tl.load(indptr_ptr + brow) + end = tl.load(indptr_ptr + brow + 1) + offs = start + SEG * BLOCK_NNZ + tl.arange(0, BLOCK_NNZ) + mask = offs < end + bcols = tl.load(indices_ptr + offs, mask=mask, other=0) + acc = tl.load( + data_ptr + start * BLOCK_DIM * BLOCK_DIM + inner_row * BLOCK_DIM, + mask=start < end, + other=0.0, + ) * 0 + for inner_col in tl.static_range(0, BLOCK_DIM): + col = bcols * BLOCK_DIM + inner_col + valid = mask & (col < n_cols) + vals = tl.load( + data_ptr + offs * BLOCK_DIM * BLOCK_DIM + inner_row * BLOCK_DIM + inner_col, + mask=mask, + other=0.0, + ) + x_vals = tl.load(x_ptr + col, mask=valid, other=0.0) + acc += tl.sum(tl.where(valid, vals * x_vals, 0.0)) + tl.atomic_add(y_ptr + row, acc) + + +@triton.jit +def _spmv_bsr_non_complex_kernel( + data_ri_ptr, + indices_ptr, + indptr_ptr, + x_ri_ptr, + y_ri_ptr, + n_rows, + n_cols, + n_block_rows, + BLOCK_DIM: tl.constexpr, + BLOCK_NNZ: tl.constexpr, + SEG: tl.constexpr, +): + brow = tl.program_id(0) + inner_row = tl.program_id(1) + if brow >= n_block_rows: + return + row = brow * BLOCK_DIM + inner_row + if row >= n_rows: + return + start = tl.load(indptr_ptr + brow) + end = tl.load(indptr_ptr + brow + 1) + offs = start + SEG * BLOCK_NNZ + tl.arange(0, BLOCK_NNZ) + mask = offs < end + bcols = tl.load(indices_ptr + offs, mask=mask, other=0) + acc_re = tl.load( + data_ri_ptr + (start * BLOCK_DIM * BLOCK_DIM + inner_row * BLOCK_DIM) * 2, + mask=start < end, + other=0.0, + ) * 0 + acc_im = tl.load( + data_ri_ptr + (start * BLOCK_DIM * BLOCK_DIM + inner_row * BLOCK_DIM) * 2 + 1, + mask=start < end, + other=0.0, + ) * 0 + for inner_col in tl.static_range(0, BLOCK_DIM): + col = bcols * BLOCK_DIM + inner_col + valid = mask & (col < n_cols) + elem = offs * BLOCK_DIM * BLOCK_DIM + inner_row * BLOCK_DIM + inner_col + a_re = tl.load(data_ri_ptr + elem * 2, mask=mask, other=0.0) + a_im = tl.load(data_ri_ptr + elem * 2 + 1, mask=mask, other=0.0) + x_re = tl.load(x_ri_ptr + col * 2, mask=valid, other=0.0) + x_im = tl.load(x_ri_ptr + col * 2 + 1, mask=valid, other=0.0) + prod_re = a_re * x_re - a_im * x_im + prod_im = a_re * x_im + a_im * x_re + acc_re += tl.sum(tl.where(valid, prod_re, 0.0)) + acc_im += tl.sum(tl.where(valid, prod_im, 0.0)) + tl.atomic_add(y_ri_ptr + row * 2, acc_re) + tl.atomic_add(y_ri_ptr + row * 2 + 1, acc_im) + + +def _prepare_spmv_bsr_matrix(data, indices, indptr, shape, block_dim): + if not all(torch.is_tensor(t) for t in (data, indices, indptr)): + raise TypeError("data, indices, indptr must all be torch.Tensor") + if data.ndim != 3: + raise ValueError("data must have shape (nnzb, block_dim, block_dim)") + if indices.ndim != 1 or indptr.ndim != 1: + raise ValueError("indices and indptr must be 1D tensors") + n_rows, n_cols = int(shape[0]), int(shape[1]) + block_dim = int(block_dim) + if block_dim <= 1: + raise ValueError("block_dim must be greater than 1 for BSR SpMV") + if data.shape[1] != block_dim or data.shape[2] != block_dim: + raise ValueError("data block dimensions must match block_dim") + n_block_rows = (n_rows + block_dim - 1) // block_dim + n_block_cols = (n_cols + block_dim - 1) // block_dim + if indptr.numel() != n_block_rows + 1: + raise ValueError( + f"indptr length must be n_block_rows+1={n_block_rows + 1}, got {indptr.numel()}" + ) + if data.shape[0] != indices.numel(): + raise ValueError("data.shape[0] and indices length must both equal nnzb") + if not all(t.is_cuda for t in (data, indices, indptr)): + raise ValueError("data, indices, indptr must be CUDA tensors") + if not all(t.device == data.device for t in (indices, indptr)): + raise ValueError("data, indices, indptr must be on the same CUDA device") + if data.dtype not in SUPPORTED_SPMV_BSR_VALUE_DTYPES: + raise TypeError(_spmv_bsr_dtype_error_message()) + if indices.dtype not in SUPPORTED_INDEX_DTYPES: + raise TypeError("indices dtype must be torch.int32 or torch.int64") + if indptr.dtype not in SUPPORTED_INDEX_DTYPES: + raise TypeError("indptr dtype must be torch.int32 or torch.int64") + data = data.contiguous() + indices = indices.contiguous() + indptr = indptr.contiguous() + if int(indptr[0].item()) != 0: + raise ValueError("indptr must start at zero") + if int(indptr[-1].item()) != data.shape[0]: + raise ValueError("indptr[-1] must equal nnzb") + if indptr.numel() > 1 and torch.any(indptr[1:] < indptr[:-1]).item(): + raise ValueError("indptr must be non-decreasing") + if indices.numel() > 0: + min_index = int(indices.min().item()) + max_index = int(indices.max().item()) + if min_index < 0 or max_index >= n_block_cols: + raise IndexError("indices out of range for n_block_cols") + block_row_lengths = indptr[1:] - indptr[:-1] + max_block_row_nnz = ( + int(block_row_lengths.max().item()) if n_block_rows > 0 else 0 + ) + return ( + data, + indices, + indptr, + n_rows, + n_cols, + n_block_rows, + n_block_cols, + block_row_lengths, + max_block_row_nnz, + ) + + +def prepare_spmv_bsr( + data, + indices, + indptr, + shape, + block_dim, + block_nnz=128, + max_segments=None, + transpose=False, + op=None, + index_fallback_policy="auto", +): + index_fallback_policy = _normalize_spmv_bsr_index_fallback_policy( + index_fallback_policy + ) + op_code = _normalize_spmv_bsr_op(op, transpose=transpose) + _ensure_spmv_bsr_supported_op(op_code) + ( + data, + indices, + indptr, + n_rows, + n_cols, + n_block_rows, + n_block_cols, + block_row_lengths, + max_block_row_nnz, + ) = _prepare_spmv_bsr_matrix(data, indices, indptr, shape, block_dim) + block_nnz_use = int(block_nnz) + if block_nnz_use <= 0: + raise ValueError("block_nnz must be positive") + if max_segments is None: + max_segments_use = max((max_block_row_nnz + block_nnz_use - 1) // block_nnz_use, 1) + while max_segments_use > 2048 and block_nnz_use < 65536: + block_nnz_use *= 2 + max_segments_use = max( + (max_block_row_nnz + block_nnz_use - 1) // block_nnz_use, + 1, + ) + else: + max_segments_use = max(1, int(max_segments)) + return PreparedBsrSpmv( + data=data, + kernel_indices=indices, + kernel_indptr=indptr, + shape=shape, + block_dim=block_dim, + n_block_rows=n_block_rows, + n_block_cols=n_block_cols, + block_nnz=block_nnz_use, + max_segments=max_segments_use, + max_block_row_nnz=max_block_row_nnz, + block_row_lengths=block_row_lengths, + op=op_code, + index_fallback_policy=index_fallback_policy, + ) + + +def _validate_spmv_bsr_x(x, prepared, op_code): + if x is None or not torch.is_tensor(x): + raise TypeError("x must be a torch.Tensor") + if x.ndim != 1: + raise ValueError("x must be a 1D tensor") + if not x.is_cuda: + raise ValueError("x must be a CUDA tensor") + if x.dtype != prepared.data.dtype: + raise TypeError("x dtype must match sparse matrix dtype") + expected = prepared.n_rows if _spmv_bsr_op_transposes(op_code) else prepared.n_cols + if x.numel() != expected: + raise ValueError(f"x length must be {expected}, got {x.numel()}") + if x.device != prepared.data.device: + raise ValueError("x must be on the same device as sparse matrix data") + return x.contiguous() + + +def _triton_spmv_bsr_kernel(prepared, x, op_code): + _ensure_spmv_bsr_supported_op(op_code) + dtype = prepared.data.dtype + y = torch.zeros(prepared.n_rows, dtype=dtype, device=prepared.data.device) + if prepared.nnzb == 0: + return y + for seg in range(prepared.max_segments): + grid = (prepared.n_block_rows, prepared.block_dim) + if _is_complex_dtype(dtype): + data_ri = torch.view_as_real(prepared.data).reshape(-1) + x_ri = torch.view_as_real(x).reshape(-1) + y_ri = torch.view_as_real(y).reshape(-1) + _spmv_bsr_non_complex_kernel[grid]( + data_ri, + prepared.kernel_indices, + prepared.kernel_indptr, + x_ri, + y_ri, + prepared.n_rows, + prepared.n_cols, + prepared.n_block_rows, + BLOCK_DIM=prepared.block_dim, + BLOCK_NNZ=prepared.block_nnz, + SEG=seg, + ) + else: + _spmv_bsr_non_real_kernel[grid]( + prepared.data, + prepared.kernel_indices, + prepared.kernel_indptr, + x, + y, + prepared.n_rows, + prepared.n_cols, + prepared.n_block_rows, + BLOCK_DIM=prepared.block_dim, + BLOCK_NNZ=prepared.block_nnz, + SEG=seg, + ) + return y + + +def _spmv_bsr_uses_int64_indices(prepared): + return ( + prepared.kernel_indices.dtype == torch.int64 + or prepared.kernel_indptr.dtype == torch.int64 + ) + + +def _spmv_bsr_int32_fallback_blocker(prepared): + if prepared.nnzb > _INDEX_LIMIT_INT32: + return f"nnzb {prepared.nnzb} cannot fit int32" + if prepared.kernel_indices.numel() > 0: + max_col = int(prepared.kernel_indices.max().item()) + if max_col > _INDEX_LIMIT_INT32: + return f"block column index {max_col} cannot fit int32" + if prepared.kernel_indptr.numel() > 0: + max_ptr = int(prepared.kernel_indptr[-1].item()) + if max_ptr > _INDEX_LIMIT_INT32: + return f"indptr offset {max_ptr} cannot fit int32" + return None + + +def _spmv_bsr_prepared_with_int32_indices(prepared, reason): + blocker = _spmv_bsr_int32_fallback_blocker(prepared) + if blocker is not None: + raise RuntimeError(f"int32 fallback is unsafe: {blocker}") from reason + return PreparedBsrSpmv( + data=prepared.data, + kernel_indices=prepared.kernel_indices.to(torch.int32).contiguous(), + kernel_indptr=prepared.kernel_indptr.to(torch.int32).contiguous(), + shape=prepared.shape, + block_dim=prepared.block_dim, + n_block_rows=prepared.n_block_rows, + n_block_cols=prepared.n_block_cols, + block_nnz=prepared.block_nnz, + max_segments=prepared.max_segments, + max_block_row_nnz=prepared.max_block_row_nnz, + block_row_lengths=prepared.block_row_lengths, + op=prepared.op, + index_fallback_policy=prepared.index_fallback_policy, + index_fallback_applied=True, + index_fallback_reason=str(reason), + ) + + +def _run_spmv_bsr_prepared_with_fallback(prepared, x, op_code): + try: + return _triton_spmv_bsr_kernel(prepared, x, op_code) + except RuntimeError as exc: + if ( + prepared.index_fallback_policy != "auto" + or not _spmv_bsr_uses_int64_indices(prepared) + ): + raise + fallback_prepared = _spmv_bsr_prepared_with_int32_indices(prepared, exc) + return _triton_spmv_bsr_kernel(fallback_prepared, x, op_code) + + +def flagsparse_spmv_bsr( + data=None, + indices=None, + indptr=None, + x=None, + shape=None, + block_dim=None, + block_nnz=128, + max_segments=None, + out=None, + return_time=False, + return_meta=False, + prepared=None, + transpose=None, + op=None, + index_fallback_policy="auto", +): + """BSR SpMV using a native Triton BSR kernel.""" + op_explicit = op is not None + op_code = _normalize_spmv_bsr_op( + op, + transpose=False if transpose is None else bool(transpose), + ) + if ( + op_explicit + and transpose is not None + and bool(transpose) != _spmv_bsr_op_transposes(op_code) + ): + raise ValueError("transpose conflicts with op") + _ensure_spmv_bsr_supported_op(op_code) + if prepared is None: + if any(arg is None for arg in (data, indices, indptr, shape, block_dim)): + raise ValueError( + "data, indices, indptr, shape, and block_dim are required when prepared is not provided" + ) + prepared = prepare_spmv_bsr( + data, + indices, + indptr, + shape, + block_dim, + block_nnz=block_nnz, + max_segments=max_segments, + op=op_code, + index_fallback_policy=index_fallback_policy, + ) + else: + if op_explicit and op_code != prepared.op: + raise ValueError( + f"op={_spmv_bsr_op_to_name(op_code)} does not match prepared.op={_spmv_bsr_op_to_name(prepared.op)}" + ) + if ( + not op_explicit + and transpose is not None + and bool(transpose) != prepared.transpose + ): + raise ValueError( + f"transpose={bool(transpose)} does not match prepared.transpose={prepared.transpose}" + ) + if not op_explicit: + op_code = prepared.op + x = _validate_spmv_bsr_x(x, prepared, op_code) + do_timing = bool(return_time or return_meta) + if do_timing: + torch.cuda.synchronize() + t0 = time.perf_counter() + y = _run_spmv_bsr_prepared_with_fallback(prepared, x, op_code) + if do_timing: + torch.cuda.synchronize() + compute_ms = (time.perf_counter() - t0) * 1000.0 + op_total_ms = compute_ms + else: + compute_ms = None + op_total_ms = None + if out is not None: + if not out.is_cuda: + raise ValueError("out must be a CUDA tensor") + if out.device != y.device: + raise ValueError("out must be on the same CUDA device as the result") + if out.shape != y.shape or out.dtype != y.dtype: + raise ValueError("out shape/dtype must match result") + out.copy_(y) + y = out + if return_meta: + meta = { + "op": _spmv_bsr_op_to_name(op_code), + "block_dim": prepared.block_dim, + "nnzb": prepared.nnzb, + "stored_nnz": prepared.stored_nnz, + "symbolic_ms": 0.0 if do_timing else None, + "compute_ms": compute_ms, + "op_total_ms": op_total_ms, + "index_fallback_applied": prepared.index_fallback_applied, + "index_fallback_reason": prepared.index_fallback_reason, + } + if return_time: + return y, op_total_ms, meta + return y, meta + if return_time: + return y, op_total_ms + return y diff --git a/tests/ci/test_cli_help.py b/tests/ci/test_cli_help.py index 0e7c15c..601c568 100644 --- a/tests/ci/test_cli_help.py +++ b/tests/ci/test_cli_help.py @@ -20,6 +20,8 @@ "run_flagsparse_pytest.py", "tests/test_spmv.py", "tests/test_spmv_coo.py", + "tests/test_spmv_csc.py", + "tests/test_spmv_bsr.py", "tests/test_spmm.py", "tests/test_spgemm.py", "tests/test_spsv.py", diff --git a/tests/ci/test_package_smoke.py b/tests/ci/test_package_smoke.py index 972b571..f05d13f 100644 --- a/tests/ci/test_package_smoke.py +++ b/tests/ci/test_package_smoke.py @@ -10,5 +10,9 @@ def test_package_version_is_exposed(): def test_public_exports_are_listed(): exported = set(dir(flagsparse)) assert "flagsparse_spmv_csr" in exported + assert "flagsparse_spmv_csc" in exported + assert "flagsparse_spmv_bsr" in exported + assert "prepare_spmv_csc" in exported + assert "prepare_spmv_bsr" in exported assert "create_csr_matrix" in exported assert "read_mtx_file" in exported diff --git a/tests/ci/test_public_api.py b/tests/ci/test_public_api.py index 3168d9c..6843dbd 100644 --- a/tests/ci/test_public_api.py +++ b/tests/ci/test_public_api.py @@ -7,6 +7,10 @@ "flagsparse_scatter", "flagsparse_spmv_csr", "flagsparse_spmv_coo", + "flagsparse_spmv_csc", + "flagsparse_spmv_bsr", + "prepare_spmv_csc", + "prepare_spmv_bsr", "flagsparse_spmm_csr", "flagsparse_spmm_coo", "flagsparse_spgemm_csr", diff --git a/tests/ci/test_runtime_policies.py b/tests/ci/test_runtime_policies.py index 01e94e6..5d3c688 100644 --- a/tests/ci/test_runtime_policies.py +++ b/tests/ci/test_runtime_policies.py @@ -18,6 +18,8 @@ gather_scatter as gather_scatter_ops, ) from flagsparse.sparse_operations import spmv_coo as spmv_coo_ops # noqa: E402 +from flagsparse.sparse_operations import spmv_bsr as spmv_bsr_ops # noqa: E402 +from flagsparse.sparse_operations import spmv_csc as spmv_csc_ops # noqa: E402 from flagsparse.sparse_operations import spmv_csr as spmv_csr_ops # noqa: E402 # isort: on # fmt: on @@ -51,6 +53,16 @@ def test_spmv_coo_index_fallback_policy_normalization(policy): assert spmv_coo_ops._normalize_spmv_coo_index_fallback_policy(policy) == policy +@pytest.mark.parametrize("policy", ["auto", "strict"]) +def test_spmv_csc_index_fallback_policy_normalization(policy): + assert spmv_csc_ops._normalize_spmv_csc_index_fallback_policy(policy) == policy + + +@pytest.mark.parametrize("policy", ["auto", "strict"]) +def test_spmv_bsr_index_fallback_policy_normalization(policy): + assert spmv_bsr_ops._normalize_spmv_bsr_index_fallback_policy(policy) == policy + + @pytest.mark.parametrize( ("op", "expected"), [ @@ -64,6 +76,32 @@ def test_spmv_coo_op_normalization(op, expected): assert spmv_coo_ops._normalize_spmv_coo_op(op) == expected +@pytest.mark.parametrize( + ("op", "expected"), + [ + (None, 0), + ("non", 0), + ("trans", 1), + ("conj", 2), + ], +) +def test_spmv_csc_op_normalization(op, expected): + assert spmv_csc_ops._normalize_spmv_csc_op(op) == expected + + +@pytest.mark.parametrize( + ("op", "expected"), + [ + (None, 0), + ("non", 0), + ("trans", 1), + ("conj", 2), + ], +) +def test_spmv_bsr_op_normalization(op, expected): + assert spmv_bsr_ops._normalize_spmv_bsr_op(op) == expected + + @pytest.mark.parametrize("op", ["non", "trans", "conj"]) def test_spmv_csr_op_transpose_contract(op): if op == "non": @@ -78,6 +116,50 @@ def test_spmv_csr_op_transpose_contract(op): ) +@pytest.mark.parametrize("op", ["non", "trans", "conj"]) +def test_spmv_csc_op_transpose_contract(op): + if op == "non": + assert ( + spmv_csc_ops._spmv_csc_op_transposes( + spmv_csc_ops._normalize_spmv_csc_op(op) + ) + is False + ) + else: + assert ( + spmv_csc_ops._spmv_csc_op_transposes( + spmv_csc_ops._normalize_spmv_csc_op(op) + ) + is True + ) + + +@pytest.mark.parametrize("op", ["non", "trans", "conj"]) +def test_spmv_bsr_op_transpose_contract(op): + if op == "non": + assert ( + spmv_bsr_ops._spmv_bsr_op_transposes( + spmv_bsr_ops._normalize_spmv_bsr_op(op) + ) + is False + ) + else: + assert ( + spmv_bsr_ops._spmv_bsr_op_transposes( + spmv_bsr_ops._normalize_spmv_bsr_op(op) + ) + is True + ) + + +@pytest.mark.parametrize("op", ["trans", "conj"]) +def test_spmv_bsr_unsupported_ops_rejected_by_policy(op): + with pytest.raises(NotImplementedError, match="only supports op='non'"): + spmv_bsr_ops._ensure_spmv_bsr_supported_op( + spmv_bsr_ops._normalize_spmv_bsr_op(op) + ) + + def test_scatter_policy_validator_rejects_unknown_policy(): with pytest.raises( ValueError, match="index_fallback_policy must be 'auto' or 'strict'" diff --git a/tests/pytest/test_spmv_bsr_accuracy.py b/tests/pytest/test_spmv_bsr_accuracy.py new file mode 100644 index 0000000..ff027e9 --- /dev/null +++ b/tests/pytest/test_spmv_bsr_accuracy.py @@ -0,0 +1,238 @@ +import importlib + +import pytest +import torch + +from flagsparse import flagsparse_spmv_bsr, prepare_spmv_bsr +from tests.pytest.accuracy_utils import close_tolerances +from tests.pytest.param_shapes import SPMV_MN_SHAPES + + +spmv_bsr_mod = importlib.import_module("flagsparse.sparse_operations.spmv_bsr") +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") + + +def _value_dtype_cases(): + cases = [ + ("float32", torch.float32), + ("float64", torch.float64), + ("complex64", torch.complex64), + ("complex128", torch.complex128), + ] + return [(name, dtype) for name, dtype in cases if dtype is not None] + + +def _random_values(shape, dtype, device): + if dtype in (torch.float32, torch.float64): + return torch.randn(shape, dtype=dtype, device=device) + if dtype == torch.complex64: + return torch.complex( + torch.randn(shape, dtype=torch.float32, device=device), + torch.randn(shape, dtype=torch.float32, device=device), + ) + if dtype == torch.complex128: + return torch.complex( + torch.randn(shape, dtype=torch.float64, device=device), + torch.randn(shape, dtype=torch.float64, device=device), + ) + raise TypeError(f"unsupported dtype: {dtype}") + + +def _reference_dtype(dtype): + if dtype == torch.float32: + return torch.float64 + if dtype == torch.complex64: + return torch.complex128 + return dtype + + +def _dense_to_bsr(dense, index_dtype, block_dim): + device = dense.device + M, N = dense.shape + n_block_rows = (M + block_dim - 1) // block_dim + rows, cols = torch.nonzero(dense != 0, as_tuple=True) + blocks = {} + for row, col in zip(rows.tolist(), cols.tolist()): + brow = int(row) // block_dim + bcol = int(col) // block_dim + inner_row = int(row) % block_dim + inner_col = int(col) % block_dim + block = blocks.setdefault( + (brow, bcol), + torch.zeros((block_dim, block_dim), dtype=dense.dtype, device=device), + ) + block[inner_row, inner_col] = dense[row, col] + row_blocks = [[] for _ in range(n_block_rows)] + for key in sorted(blocks): + row_blocks[key[0]].append(key) + data = [] + indices = [] + indptr = [0] + for keys in row_blocks: + for key in keys: + indices.append(key[1]) + data.append(blocks[key]) + indptr.append(len(indices)) + if data: + data_tensor = torch.stack(data).contiguous() + else: + data_tensor = torch.empty((0, block_dim, block_dim), dtype=dense.dtype, device=device) + return ( + data_tensor, + torch.tensor(indices, dtype=index_dtype, device=device), + torch.tensor(indptr, dtype=index_dtype, device=device), + ) + + +def _random_bsr_mn(M, N, dtype, index_dtype, block_dim, device): + denom = max(M * N, 1) + p = min(0.25, max(0.06, 32.0 / denom)) + mask = torch.rand(M, N, device=device) < p + if int(mask.sum().item()) == 0: + mask[0, 0] = True + dense = torch.where( + mask, + _random_values((M, N), dtype, device), + torch.zeros((), dtype=dtype, device=device), + ) + data, indices, indptr = _dense_to_bsr(dense, index_dtype, block_dim) + return data, indices, indptr, dense + + +def _make_x(length, dtype, device): + return _random_values((length,), dtype, device) + + +def _assert_close(actual, expected, dtype): + rtol, atol = close_tolerances(dtype) + ref_dtype = _reference_dtype(dtype) + assert torch.allclose( + actual.to(ref_dtype), expected.to(ref_dtype), rtol=rtol, atol=atol + ) + + +@pytest.mark.spmv_bsr +@pytest.mark.parametrize("M, N", SPMV_MN_SHAPES) +@pytest.mark.parametrize( + "name,dtype", _value_dtype_cases(), ids=[c[0] for c in _value_dtype_cases()] +) +@pytest.mark.parametrize( + "index_dtype", [torch.int32, torch.int64], ids=["int32", "int64"] +) +@pytest.mark.parametrize("block_dim", [2, 4], ids=["block2", "block4"]) +def test_spmv_bsr_matches_dense_reference(M, N, name, dtype, index_dtype, block_dim): + device = torch.device("cuda") + data, indices, indptr, dense = _random_bsr_mn( + M, N, dtype, index_dtype, block_dim, device + ) + x = _make_x(N, dtype, device) + ref_dtype = _reference_dtype(dtype) + ref = (dense.to(ref_dtype) @ x.to(ref_dtype)).to(dtype) + out = flagsparse_spmv_bsr( + data, + indices, + indptr, + x, + shape=(M, N), + block_dim=block_dim, + index_fallback_policy="auto", + ) + _assert_close(out, ref, dtype) + + +@pytest.mark.spmv_bsr +def test_spmv_bsr_prepared_path_matches_dense_reference(): + device = torch.device("cuda") + M, N = 8, 10 + dtype = torch.complex64 + block_dim = 4 + data, indices, indptr, dense = _random_bsr_mn( + M, N, dtype, torch.int32, block_dim, device + ) + prepared = prepare_spmv_bsr(data, indices, indptr, (M, N), block_dim, op="non") + x = _make_x(N, dtype, device) + ref = (dense.to(torch.complex128) @ x.to(torch.complex128)).to(dtype) + out = flagsparse_spmv_bsr(x=x, prepared=prepared) + _assert_close(out, ref, dtype) + + +@pytest.mark.spmv_bsr +@pytest.mark.parametrize("op", ["trans", "conj"], ids=["trans", "conj"]) +def test_spmv_bsr_unsupported_ops_are_rejected(op): + device = torch.device("cuda") + data, indices, indptr, _dense = _random_bsr_mn( + 8, 10, torch.float32, torch.int32, 2, device + ) + x = torch.randn(8, dtype=torch.float32, device=device) + with pytest.raises(NotImplementedError, match="only supports op='non'"): + flagsparse_spmv_bsr(data, indices, indptr, x, shape=(8, 10), block_dim=2, op=op) + + +@pytest.mark.spmv_bsr +def test_spmv_bsr_x_length_mismatch_rejected(): + device = torch.device("cuda") + data, indices, indptr, _dense = _random_bsr_mn( + 8, 10, torch.float32, torch.int32, 2, device + ) + prepared = prepare_spmv_bsr(data, indices, indptr, (8, 10), 2) + x = torch.randn(8, dtype=torch.float32, device=device) + with pytest.raises(ValueError, match="x length must be 10"): + flagsparse_spmv_bsr(x=x, prepared=prepared) + + +@pytest.mark.spmv_bsr +def test_spmv_bsr_int64_auto_fallback_to_int32(monkeypatch): + device = torch.device("cuda") + data, indices, indptr, dense = _random_bsr_mn( + 12, 9, torch.float32, torch.int64, 2, device + ) + x = torch.randn(9, dtype=torch.float32, device=device) + ref = dense.to(torch.float64) @ x.to(torch.float64) + state = {"forced_once": False} + original = spmv_bsr_mod._triton_spmv_bsr_kernel + + def fail_int64_once(prepared, x_in, op_code): + if prepared.kernel_indices.dtype == torch.int64 and not state["forced_once"]: + state["forced_once"] = True + raise RuntimeError("forced int64 launch failure") + return original(prepared, x_in, op_code) + + monkeypatch.setattr(spmv_bsr_mod, "_triton_spmv_bsr_kernel", fail_int64_once) + out = flagsparse_spmv_bsr( + data, + indices, + indptr, + x, + shape=(12, 9), + block_dim=2, + index_fallback_policy="auto", + ) + assert state["forced_once"] + _assert_close(out, ref.to(torch.float32), torch.float32) + + +@pytest.mark.spmv_bsr +def test_spmv_bsr_int64_strict_no_fallback(monkeypatch): + device = torch.device("cuda") + data, indices, indptr, _dense = _random_bsr_mn( + 12, 9, torch.float32, torch.int64, 2, device + ) + x = torch.randn(9, dtype=torch.float32, device=device) + original = spmv_bsr_mod._triton_spmv_bsr_kernel + + def fail_int64(prepared, x_in, op_code): + if prepared.kernel_indices.dtype == torch.int64: + raise RuntimeError("forced int64 launch failure") + return original(prepared, x_in, op_code) + + monkeypatch.setattr(spmv_bsr_mod, "_triton_spmv_bsr_kernel", fail_int64) + with pytest.raises(RuntimeError, match="forced int64 launch failure"): + flagsparse_spmv_bsr( + data, + indices, + indptr, + x, + shape=(12, 9), + block_dim=2, + index_fallback_policy="strict", + ) diff --git a/tests/test_spmv_bsr.py b/tests/test_spmv_bsr.py new file mode 100644 index 0000000..edde03d --- /dev/null +++ b/tests/test_spmv_bsr.py @@ -0,0 +1,797 @@ +"""Native BSR SpMV benchmark and correctness script.""" + +import argparse +import csv +import glob +import math +import os +import sys +from pathlib import Path + +import torch + +_PROJECT_ROOT = Path(__file__).resolve().parents[1] +_SRC_ROOT = _PROJECT_ROOT / "src" +if str(_SRC_ROOT) not in sys.path: + sys.path.insert(0, str(_SRC_ROOT)) + +import flagsparse as fs + +try: + import cupy as cp + import cupyx.scipy.sparse as cpx_sparse +except ImportError: + cp = None + cpx_sparse = None + + +VALUE_DTYPES = (torch.float32, torch.float64, torch.complex64, torch.complex128) +INDEX_DTYPES = (torch.int32, torch.int64) +OPS = ("non", "trans", "conj") +SUPPORTED_OPS = ("non",) +TEST_SIZES = ((64, 96), (160, 1024), (128, 256)) +DEFAULT_BLOCK_DIMS = (4,) +WARMUP = 10 +ITERS = 50 + + +def _dtype_name(dtype): + return str(dtype).replace("torch.", "") + + +DTYPE_MAP = { + "float32": torch.float32, + "float64": torch.float64, + "complex64": torch.complex64, + "complex128": torch.complex128, +} +INDEX_DTYPE_MAP = {"int32": torch.int32, "int64": torch.int64} + + +def _parse_csv_tokens(value, mapping, option_name): + tokens = [token.strip().lower() for token in str(value).split(",") if token.strip()] + if not tokens: + raise ValueError(f"{option_name} must not be empty") + invalid = [token for token in tokens if token not in mapping] + if invalid: + raise ValueError( + f"unsupported {option_name}: {', '.join(invalid)}; allowed: {', '.join(mapping)}" + ) + return [mapping[token] for token in tokens] + + +def _parse_ops(value): + token = "non" if value is None else str(value).strip().lower() + if token == "all": + return list(OPS) + ops = [item.strip().lower() for item in token.split(",") if item.strip()] + invalid = [op for op in ops if op not in OPS] + if not ops or invalid: + raise ValueError(f"unsupported --ops: {', '.join(invalid or ops)}") + return ops + + +def _parse_block_dims(value): + token = str(value or "4").strip().lower() + if token == "auto": + return ["auto"] + dims = [] + for item in token.split(","): + item = item.strip() + if not item: + continue + dim = int(item) + if dim <= 1: + raise ValueError("--block-dims values must be greater than 1") + dims.append(dim) + if not dims: + raise ValueError("--block-dims must not be empty") + return dims + + +def _random_values(shape, dtype, device): + if dtype in (torch.float32, torch.float64): + return torch.randn(shape, dtype=dtype, device=device) + if dtype == torch.complex64: + return torch.complex( + torch.randn(shape, dtype=torch.float32, device=device), + torch.randn(shape, dtype=torch.float32, device=device), + ) + if dtype == torch.complex128: + return torch.complex( + torch.randn(shape, dtype=torch.float64, device=device), + torch.randn(shape, dtype=torch.float64, device=device), + ) + raise TypeError(f"unsupported dtype: {dtype}") + + +def _reference_dtype(dtype): + if dtype == torch.float32: + return torch.float64 + if dtype == torch.complex64: + return torch.complex128 + return dtype + + +def _reference_tolerance(dtype): + if dtype in (torch.float32, torch.complex64): + return 1.3e-6, 1e-3 + if dtype in (torch.float64, torch.complex128): + return 1e-7, 1e-5 + return 1e-6, 1e-5 + + +def _mtx_value_for_dtype(raw_value, dtype): + if dtype in (torch.complex64, torch.complex128): + return complex(raw_value) + return float(raw_value.real if isinstance(raw_value, complex) else raw_value) + + +def _zero_value(dtype): + return 0j if dtype in (torch.complex64, torch.complex128) else 0.0 + + +def _choose_auto_block_dim(entries, shape): + n_rows, n_cols = shape + nnz = max(1, len(entries)) + for block_dim in (16, 8, 4, 2): + blocks = { + (int(row) // block_dim, int(col) // block_dim) + for row, col in entries.keys() + } + stored = len(blocks) * block_dim * block_dim + if stored <= 2.0 * nnz: + return block_dim + return 4 if max(n_rows, n_cols) >= 4 else 2 + + +def _entries_to_bsr_torch(entries, shape, dtype, index_dtype, block_dim, device): + n_rows, n_cols = int(shape[0]), int(shape[1]) + block_dim = int(block_dim) + blocks = {} + for (row, col), value in entries.items(): + brow = int(row) // block_dim + bcol = int(col) // block_dim + inner_row = int(row) % block_dim + inner_col = int(col) % block_dim + block = blocks.setdefault( + (brow, bcol), + [_zero_value(dtype) for _ in range(block_dim * block_dim)], + ) + block[inner_row * block_dim + inner_col] += _mtx_value_for_dtype(value, dtype) + n_block_rows = (n_rows + block_dim - 1) // block_dim + rows = [[] for _ in range(n_block_rows)] + for key in sorted(blocks): + rows[key[0]].append(key) + data_values = [] + indices_values = [] + indptr_values = [0] + for row_blocks in rows: + for key in row_blocks: + indices_values.append(key[1]) + data_values.extend(blocks[key]) + indptr_values.append(len(indices_values)) + data = torch.tensor(data_values, dtype=dtype, device=device) + data = data.reshape(-1, block_dim, block_dim).contiguous() + indices = torch.tensor(indices_values, dtype=index_dtype, device=device) + indptr = torch.tensor(indptr_values, dtype=index_dtype, device=device) + return data, indices.contiguous(), indptr.contiguous() + + +def _dense_to_bsr(dense, index_dtype, block_dim): + rows, cols = dense.nonzero(as_tuple=True) + entries = { + (int(row.item()), int(col.item())): dense[row, col].item() + for row, col in zip(rows, cols) + } + return _entries_to_bsr_torch( + entries, + tuple(dense.shape), + dense.dtype, + index_dtype, + block_dim, + dense.device, + ) + + +def _bsr_block_rows(indptr): + counts = indptr[1:].to(torch.int64) - indptr[:-1].to(torch.int64) + return torch.repeat_interleave( + torch.arange(indptr.numel() - 1, dtype=torch.int64, device=indptr.device), + counts, + ) + + +def _bsr_to_torch_coo(data, indices, indptr, shape, block_dim): + block_rows = _bsr_block_rows(indptr) + nnzb = int(data.shape[0]) + if nnzb == 0: + empty = torch.empty(0, dtype=torch.int64, device=data.device) + return torch.sparse_coo_tensor( + torch.stack([empty, empty]), + data.reshape(-1), + size=shape, + device=data.device, + dtype=data.dtype, + ).coalesce() + local = torch.arange(block_dim * block_dim, dtype=torch.int64, device=data.device) + inner_rows = local // block_dim + inner_cols = local % block_dim + rows = block_rows[:, None] * block_dim + inner_rows[None, :] + cols = indices.to(torch.int64)[:, None] * block_dim + inner_cols[None, :] + values = data.reshape(nnzb, block_dim * block_dim) + mask = (rows < int(shape[0])) & (cols < int(shape[1])) & (values != 0) + rows = rows[mask] + cols = cols[mask] + values = values[mask] + return torch.sparse_coo_tensor( + torch.stack([rows, cols]), + values, + size=shape, + device=data.device, + dtype=data.dtype, + ).coalesce() + + +def _pytorch_reference(data, indices, indptr, x, shape, dtype, block_dim): + ref_dtype = _reference_dtype(dtype) + A = _bsr_to_torch_coo( + data.to(ref_dtype), + indices, + indptr, + shape, + block_dim, + ) + return torch.sparse.mm(A, x.to(ref_dtype).unsqueeze(1)).squeeze(1).to(dtype) + + +def _allclose_error_ratio(actual, expected, atol, rtol): + if expected.numel() == 0: + return 0.0 + diff = torch.abs(actual - expected).to(torch.float64) + denom = atol + rtol * torch.abs(expected).to(torch.float64) + return float(torch.max(diff / denom).item()) + + +def _cuda_event_benchmark(op, warmup, iters): + out = None + count = max(1, int(iters)) + for _ in range(max(0, int(warmup))): + out = op() + torch.cuda.synchronize() + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(count): + out = op() + end.record() + torch.cuda.synchronize() + return out, start.elapsed_time(end) / count + + +def _time_flagsparse_bsr(data, indices, indptr, x, shape, block_dim, warmup, iters, timing=False): + prepared = fs.prepare_spmv_bsr(data, indices, indptr, shape, block_dim, op="non") + out, gpu_ms = _cuda_event_benchmark( + lambda: fs.flagsparse_spmv_bsr(x=x, prepared=prepared), + warmup, + iters, + ) + return { + "out": out, + "ms": gpu_ms, + "gpu_ms": gpu_ms, + "process_cpu_ms": 0.0, + "process_gpu_ms": 0.0 if timing else None, + "compute_ms": gpu_ms if timing else None, + } + + +def _time_pytorch(data, indices, indptr, x, shape, block_dim, warmup, iters): + A = _bsr_to_torch_coo(data, indices, indptr, shape, block_dim) + fn = lambda: torch.sparse.mm(A, x.unsqueeze(1)).squeeze(1) + _, ms = _cuda_event_benchmark(fn, warmup, iters) + return ms + + +def _time_cusparse(data, indices, indptr, x, shape, block_dim, warmup, iters): + if cp is None or cpx_sparse is None: + return None + if data.dtype not in (torch.float32, torch.float64, torch.complex64, torch.complex128): + return None + data_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(data)) + ind_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(indices.to(torch.int64))) + ptr_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(indptr.to(torch.int64))) + x_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(x)) + A = cpx_sparse.bsr_matrix((data_cp, ind_cp, ptr_cp), shape=shape) + fn = lambda: A @ x_cp + for _ in range(max(0, int(warmup))): + _ = fn() + cp.cuda.runtime.deviceSynchronize() + start = cp.cuda.Event() + end = cp.cuda.Event() + count = max(1, int(iters)) + start.record() + for _ in range(count): + _ = fn() + end.record() + end.synchronize() + return cp.cuda.get_elapsed_time(start, end) / count + + +def _fmt(v): + return "N/A" if v is None else f"{v:.4f}" + + +def _fmt_err(v): + return "N/A" if v is None else f"{v:.2e}" + + +def _spd(base, other): + if base is None or other is None or other <= 0: + return "N/A" + return f"{base / other:.2f}x" + + +def _status(ok): + return "PASS" if ok else "FAIL" + + +def _header(timing=False): + split = f" {'ProcGPU':>9} {'Compute':>9}" if timing else "" + return ( + f"{'Matrix':<28} {'Op':>5} {'BDim':>5} {'Out':>7} {'Rows':>7} {'Cols':>7} {'NNZB':>9} {'Pad':>7} " + f"{'BSR(ms)':>9} {'BSRGPU':>9} {'CPUProc':>9}{split} " + f"{'PT(ms)':>9} {'CU(ms)':>9} {'BSR/PT':>8} {'BSR/CU':>8} " + f"{'Err':>10} {'Status':>6}" + ) + + +def _sep(timing=False): + return "-" * (158 if timing else 138) + + +def _print_row(row, timing=False): + name = str(row["matrix"])[:27] + if len(str(row["matrix"])) > 27: + name += "..." + split = ( + f" {_fmt(row.get('process_gpu_ms')):>9} {_fmt(row.get('compute_ms')):>9}" + if timing + else "" + ) + print( + f"{name:<28} {row['op']:>5} {row['block_dim']:>5} {row['out_size']:>7} {row['n_rows']:>7} {row['n_cols']:>7} {row['nnzb']:>9} {row['padding_ratio']:>7} " + f"{_fmt(row['bsr_ms']):>9} {_fmt(row['bsr_gpu_ms']):>9} {_fmt(row['process_cpu_ms']):>9}{split} " + f"{_fmt(row['pytorch_ms']):>9} {_fmt(row['cusparse_ms']):>9} " + f"{_spd(row['pytorch_ms'], row['bsr_ms']):>8} {_spd(row['cusparse_ms'], row['bsr_ms']):>8} " + f"{_fmt_err(row['err']):>10} {row['status']:>6}" + ) + error = row.get("error") + if error: + print(f" error: {str(error)[:240]}") + + +def _base_row( + matrix_name, + dtype, + index_dtype, + op, + shape, + data, + block_dim, + logical_nnz=None, + status="ERROR", +): + nnzb = int(data.shape[0]) + stored_nnz = int(data.numel()) + logical_nnz = max(1, int(logical_nnz if logical_nnz is not None else stored_nnz)) + return { + "matrix": matrix_name, + "value_dtype": _dtype_name(dtype), + "index_dtype": _dtype_name(index_dtype), + "op": op, + "block_dim": int(block_dim), + "out_size": int(shape[0]) if op == "non" else "UNSUP", + "n_rows": int(shape[0]), + "n_cols": int(shape[1]), + "nnzb": nnzb, + "logical_nnz": logical_nnz, + "stored_nnz": stored_nnz, + "padding_ratio": f"{stored_nnz / logical_nnz:.2f}", + "bsr_ms": None, + "bsr_gpu_ms": None, + "process_cpu_ms": 0.0, + "process_gpu_ms": None, + "compute_ms": None, + "pytorch_ms": None, + "cusparse_ms": None, + "err": None, + "status": status, + "error": None, + } + + +def _run_one_case( + data, + indices, + indptr, + shape, + dtype, + index_dtype, + op, + matrix_name, + block_dim, + warmup, + iters, + timing=False, + run_cusparse=True, + logical_nnz=None, +): + data = data.contiguous() + indices = indices.to(index_dtype).contiguous() + indptr = indptr.to(index_dtype).contiguous() + row = _base_row( + matrix_name, + dtype, + index_dtype, + op, + shape, + data, + block_dim, + logical_nnz=logical_nnz, + ) + row["process_gpu_ms"] = 0.0 if timing else None + if op not in SUPPORTED_OPS: + row["status"] = "SKIP" + row["error"] = "BSR SpMV v1 only supports op=non" + return row + x = _random_values((int(shape[1]),), dtype, data.device) + atol, rtol = _reference_tolerance(dtype) + try: + bsr = _time_flagsparse_bsr(data, indices, indptr, x, shape, block_dim, warmup, iters, timing=timing) + except Exception as exc: + row["error"] = f"flagsparse_spmv_bsr failed: {exc}" + return row + row.update( + { + "bsr_ms": bsr["ms"], + "bsr_gpu_ms": bsr["gpu_ms"], + "process_cpu_ms": bsr["process_cpu_ms"], + "process_gpu_ms": bsr["process_gpu_ms"], + "compute_ms": bsr["compute_ms"], + } + ) + try: + y_ref = _pytorch_reference(data, indices, indptr, x, shape, dtype, block_dim) + err = _allclose_error_ratio(bsr["out"], y_ref, atol, rtol) + except Exception as exc: + row["error"] = f"reference failed after BSR run: {exc}" + return row + try: + row["pytorch_ms"] = _time_pytorch(data, indices, indptr, x, shape, block_dim, warmup, iters) + except Exception: + pass + if run_cusparse: + try: + row["cusparse_ms"] = _time_cusparse(data, indices, indptr, x, shape, block_dim, warmup, iters) + except Exception: + pass + ok = (not math.isnan(err)) and err <= 1.0 + row["err"] = err + row["status"] = _status(ok) + row["error"] = None if ok else "correctness check failed" + return row + + +def load_mtx_entries(path): + with open(path, "r", encoding="utf-8") as handle: + lines = handle.readlines() + mm_field = "real" + mm_symmetry = "general" + header = None + data_lines = [] + for line in lines: + stripped = line.strip() + if stripped.startswith("%%MatrixMarket"): + parts = stripped.split() + if len(parts) >= 5: + mm_field = parts[3].lower() + mm_symmetry = parts[4].lower() + continue + if stripped.startswith("%"): + continue + if header is None and stripped: + parts = stripped.split() + header = (int(parts[0]), int(parts[1]), int(parts[2]) if len(parts) > 2 else 0) + continue + if stripped: + data_lines.append(stripped) + if header is None: + raise ValueError(f"Cannot parse .mtx header: {path}") + n_rows, n_cols, nnz = header + entries = {} + + def add_entry(r, c, value): + key = (int(r), int(c)) + entries[key] = entries.get(key, 0.0) + value + + is_pattern = mm_field == "pattern" + is_complex = mm_field == "complex" + is_symmetric = mm_symmetry == "symmetric" + is_skew = mm_symmetry == "skew-symmetric" + is_hermitian = mm_symmetry == "hermitian" + for line in data_lines[:nnz]: + parts = line.split() + if len(parts) < 2: + continue + r = int(parts[0]) - 1 + c = int(parts[1]) - 1 + if not (0 <= r < n_rows and 0 <= c < n_cols): + continue + if is_pattern: + value = 1.0 + elif is_complex: + value = complex(float(parts[2]), float(parts[3])) + else: + value = float(parts[2]) + add_entry(r, c, value) + if r != c: + if is_symmetric and 0 <= c < n_rows and 0 <= r < n_cols: + add_entry(c, r, value) + elif is_skew and 0 <= c < n_rows and 0 <= r < n_cols: + add_entry(c, r, -value) + elif is_hermitian and 0 <= c < n_rows and 0 <= r < n_cols: + add_entry(c, r, value.conjugate() if isinstance(value, complex) else value) + return entries, (n_rows, n_cols) + + +def _resolve_block_dims(block_dims, entries, shape): + if block_dims == ["auto"]: + return [_choose_auto_block_dim(entries, shape)] + return block_dims + + +def run_synthetic(value_dtypes=None, index_dtypes=None, block_dims=None, ops=None, warmup=WARMUP, iters=ITERS, timing=False, run_cusparse=True): + if not torch.cuda.is_available(): + print("CUDA is not available.") + return + device = torch.device("cuda") + value_dtypes = VALUE_DTYPES if value_dtypes is None else value_dtypes + index_dtypes = INDEX_DTYPES if index_dtypes is None else index_dtypes + block_dims = list(DEFAULT_BLOCK_DIMS) if block_dims is None else block_dims + ops = SUPPORTED_OPS if ops is None else ops + print("=" * 140) + print("FLAGSPARSE SpMV BSR BENCHMARK (native BSR Triton)") + print("=" * 140) + print("Timing policy: bsr_ms = process_cpu_ms + bsr_gpu_ms; BSR construction is setup.") + for dtype in value_dtypes: + for index_dtype in index_dtypes: + for block_dim in block_dims: + for op in ops: + print(_sep(timing)) + print(f"dtype: {_dtype_name(dtype)} | index_dtype: {_dtype_name(index_dtype)} | block_dim: {block_dim} | op: {op}") + print(_sep(timing)) + print(_header(timing)) + print(_sep(timing)) + for m, n in TEST_SIZES: + dense = _random_values((m, n), dtype, device) + dense *= (torch.rand(m, n, device=device) < 0.1).to(dtype=dtype) + logical_nnz = int(torch.count_nonzero(dense).item()) + data, indices, indptr = _dense_to_bsr(dense, index_dtype, int(block_dim)) + row = _run_one_case( + data, + indices, + indptr, + (m, n), + dtype, + index_dtype, + op, + f"{m}x{n}", + int(block_dim), + warmup, + iters, + timing=timing, + run_cusparse=run_cusparse, + logical_nnz=logical_nnz, + ) + _print_row(row, timing=timing) + print(_sep(timing)) + print() + + +def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dims=None, ops=None, warmup=WARMUP, iters=ITERS, timing=False, run_cusparse=True, fail_fast=False): + if not torch.cuda.is_available(): + print("CUDA is not available.") + return + device = torch.device("cuda") + value_dtypes = VALUE_DTYPES if value_dtypes is None else value_dtypes + index_dtypes = INDEX_DTYPES if index_dtypes is None else index_dtypes + block_dims = list(DEFAULT_BLOCK_DIMS) if block_dims is None else block_dims + ops = SUPPORTED_OPS if ops is None else ops + rows = [] + for dtype in value_dtypes: + for index_dtype in index_dtypes: + for op in ops: + print(_sep(timing)) + print(f"Value dtype: {_dtype_name(dtype)} | Index dtype: {_dtype_name(index_dtype)} | op: {op}") + print(_sep(timing)) + print(_header(timing)) + print(_sep(timing)) + for path in mtx_paths: + try: + entries, shape = load_mtx_entries(path) + for block_dim in _resolve_block_dims(block_dims, entries, shape): + data, indices, indptr = _entries_to_bsr_torch( + entries, shape, dtype, index_dtype, int(block_dim), device + ) + row = _run_one_case( + data, + indices, + indptr, + shape, + dtype, + index_dtype, + op, + os.path.basename(path), + int(block_dim), + warmup, + iters, + timing=timing, + run_cusparse=run_cusparse, + logical_nnz=len(entries), + ) + if fail_fast and row.get("status") == "ERROR": + raise RuntimeError(row.get("error") or "BSR SpMV case failed") + rows.append(row) + _print_row(row, timing=timing) + except Exception as exc: + if fail_fast: + raise + row = { + "matrix": os.path.basename(path), + "value_dtype": _dtype_name(dtype), + "index_dtype": _dtype_name(index_dtype), + "op": op, + "block_dim": "ERR", + "out_size": "ERR", + "n_rows": "ERR", + "n_cols": "ERR", + "nnzb": "ERR", + "logical_nnz": "ERR", + "stored_nnz": "ERR", + "padding_ratio": "ERR", + "bsr_ms": None, + "bsr_gpu_ms": None, + "process_cpu_ms": None, + "process_gpu_ms": None, + "compute_ms": None, + "pytorch_ms": None, + "cusparse_ms": None, + "err": None, + "status": "ERROR", + "error": str(exc), + } + rows.append(row) + _print_row(row, timing=timing) + print(_sep(timing)) + fieldnames = [ + "matrix", + "value_dtype", + "index_dtype", + "op", + "block_dim", + "out_size", + "n_rows", + "n_cols", + "nnzb", + "logical_nnz", + "stored_nnz", + "padding_ratio", + "bsr_ms", + "bsr_gpu_ms", + "process_cpu_ms", + "process_gpu_ms", + "compute_ms", + "pytorch_ms", + "cusparse_ms", + "err", + "status", + "error", + ] + if not timing: + fieldnames = [ + field + for field in fieldnames + if field not in ("process_gpu_ms", "compute_ms") + ] + csv_parent = Path(csv_path).parent + if str(csv_parent) not in ("", "."): + csv_parent.mkdir(parents=True, exist_ok=True) + with open(csv_path, "w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fieldnames, extrasaction="ignore") + writer.writeheader() + for row in rows: + writer.writerow({key: ("" if value is None else value) for key, value in row.items()}) + print(f"Wrote {len(rows)} rows to {csv_path}") + + +def main(): + parser = argparse.ArgumentParser(description="Native BSR SpMV benchmark/test.") + parser.add_argument("mtx", nargs="*", help=".mtx files or directories") + parser.add_argument("--synthetic", action="store_true") + parser.add_argument("--csv-bsr", type=str, default=None, metavar="FILE") + parser.add_argument("--dtypes", default="float32,float64,complex64,complex128") + parser.add_argument("--index-dtypes", default="int32,int64") + parser.add_argument("--block-dims", default="4") + parser.add_argument("--ops", default="non") + parser.add_argument("--warmup", type=int, default=WARMUP) + parser.add_argument("--iters", type=int, default=ITERS) + parser.add_argument("--timing", action="store_true") + parser.add_argument("--no-cusparse", action="store_true") + parser.add_argument("--fail-fast", action="store_true") + args = parser.parse_args() + try: + value_dtypes = _parse_csv_tokens(args.dtypes, DTYPE_MAP, "--dtypes") + index_dtypes = _parse_csv_tokens(args.index_dtypes, INDEX_DTYPE_MAP, "--index-dtypes") + block_dims = _parse_block_dims(args.block_dims) + ops = _parse_ops(args.ops) + except ValueError as exc: + parser.error(str(exc)) + if args.synthetic: + run_synthetic( + value_dtypes=value_dtypes, + index_dtypes=index_dtypes, + block_dims=block_dims, + ops=ops, + warmup=args.warmup, + iters=args.iters, + timing=args.timing, + run_cusparse=not args.no_cusparse, + ) + return + paths = [] + for path in args.mtx: + if os.path.isfile(path) and path.endswith(".mtx"): + paths.append(path) + elif os.path.isdir(path): + paths.extend(sorted(glob.glob(os.path.join(path, "*.mtx")))) + if args.csv_bsr: + if not paths: + paths = sorted(glob.glob("*.mtx")) + if not paths: + print("No .mtx files found for --csv-bsr") + return + run_csv( + paths, + args.csv_bsr, + value_dtypes=value_dtypes, + index_dtypes=index_dtypes, + block_dims=block_dims, + ops=ops, + warmup=args.warmup, + iters=args.iters, + timing=args.timing, + run_cusparse=not args.no_cusparse, + fail_fast=args.fail_fast, + ) + return + if not paths: + print("No .mtx files. Use --synthetic or --csv-bsr with inputs.") + return + run_csv( + paths, + "spmv_bsr_results.csv", + value_dtypes=value_dtypes, + index_dtypes=index_dtypes, + block_dims=block_dims, + ops=ops, + warmup=args.warmup, + iters=args.iters, + timing=args.timing, + run_cusparse=not args.no_cusparse, + fail_fast=args.fail_fast, + ) + + +if __name__ == "__main__": + main() diff --git a/tools/ci/run_gpu_benchmark.py b/tools/ci/run_gpu_benchmark.py index 27b40b8..7bf4277 100644 --- a/tools/ci/run_gpu_benchmark.py +++ b/tools/ci/run_gpu_benchmark.py @@ -28,6 +28,7 @@ def _parse_args() -> argparse.Namespace: "spmv", "spmv-coo", "spmv-csc", + "spmv-bsr", "spmm", "spmm-coo", "spsv", @@ -103,6 +104,15 @@ def _command_specs( str(args.iters), *no_cusparse, ], + "spmv-bsr": [ + "tests/test_spmv_bsr.py", + "--synthetic", + "--warmup", + str(args.warmup), + "--iters", + str(args.iters), + *no_cusparse, + ], "spmm": [ "tests/test_spmm.py", "--synthetic", @@ -146,6 +156,7 @@ def _command_specs( "spmv", "spmv-coo", "spmv-csc", + "spmv-bsr", "spmm", "spmm-coo", "spsv", From 1b4e2b120d90bda880b9120b6c146ec179e7082f Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Sat, 11 Jul 2026 19:38:36 +0800 Subject: [PATCH 07/13] spmv_bsr --- tests/test_spmv_bsr.py | 76 +++++++++++++++++++++++++++++------------- 1 file changed, 53 insertions(+), 23 deletions(-) diff --git a/tests/test_spmv_bsr.py b/tests/test_spmv_bsr.py index edde03d..c8f543b 100644 --- a/tests/test_spmv_bsr.py +++ b/tests/test_spmv_bsr.py @@ -295,27 +295,42 @@ def _time_pytorch(data, indices, indptr, x, shape, block_dim, warmup, iters): def _time_cusparse(data, indices, indptr, x, shape, block_dim, warmup, iters): if cp is None or cpx_sparse is None: - return None + return None, "CuPy/cupyx.scipy.sparse is not available" if data.dtype not in (torch.float32, torch.float64, torch.complex64, torch.complex128): - return None - data_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(data)) - ind_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(indices.to(torch.int64))) - ptr_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(indptr.to(torch.int64))) - x_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(x)) - A = cpx_sparse.bsr_matrix((data_cp, ind_cp, ptr_cp), shape=shape) - fn = lambda: A @ x_cp - for _ in range(max(0, int(warmup))): - _ = fn() - cp.cuda.runtime.deviceSynchronize() - start = cp.cuda.Event() - end = cp.cuda.Event() - count = max(1, int(iters)) - start.record() - for _ in range(count): - _ = fn() - end.record() - end.synchronize() - return cp.cuda.get_elapsed_time(start, end) / count + return None, f"unsupported cuSPARSE dtype: {_dtype_name(data.dtype)}" + + def run_with_index_dtype(index_dtype, fallback_note=None): + data_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(data)) + ind_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(indices.to(index_dtype))) + ptr_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(indptr.to(index_dtype))) + x_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(x)) + A = cpx_sparse.bsr_matrix((data_cp, ind_cp, ptr_cp), shape=shape) + fn = lambda: A @ x_cp + for _ in range(max(0, int(warmup))): + _ = fn() + cp.cuda.runtime.deviceSynchronize() + start = cp.cuda.Event() + end = cp.cuda.Event() + count = max(1, int(iters)) + start.record() + for _ in range(count): + _ = fn() + end.record() + end.synchronize() + return cp.cuda.get_elapsed_time(start, end) / count, fallback_note + + try: + return run_with_index_dtype(indices.dtype) + except Exception as exc: + if indices.dtype == torch.int32 and indptr.dtype == torch.int32: + raise + try: + return run_with_index_dtype( + torch.int32, + f"CuPy BSR failed with {_dtype_name(indices.dtype)} indices; used int32 baseline fallback: {exc}", + ) + except Exception: + raise exc def _fmt(v): @@ -369,6 +384,9 @@ def _print_row(row, timing=False): error = row.get("error") if error: print(f" error: {str(error)[:240]}") + cusparse_error = row.get("cusparse_error") + if cusparse_error and row.get("cusparse_ms") is None: + print(f" cusparse: {str(cusparse_error)[:240]}") def _base_row( @@ -405,6 +423,7 @@ def _base_row( "compute_ms": None, "pytorch_ms": None, "cusparse_ms": None, + "cusparse_error": None, "err": None, "status": status, "error": None, @@ -473,9 +492,18 @@ def _run_one_case( pass if run_cusparse: try: - row["cusparse_ms"] = _time_cusparse(data, indices, indptr, x, shape, block_dim, warmup, iters) - except Exception: - pass + row["cusparse_ms"], row["cusparse_error"] = _time_cusparse( + data, + indices, + indptr, + x, + shape, + block_dim, + warmup, + iters, + ) + except Exception as exc: + row["cusparse_error"] = str(exc) ok = (not math.isnan(err)) and err <= 1.0 row["err"] = err row["status"] = _status(ok) @@ -667,6 +695,7 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "compute_ms": None, "pytorch_ms": None, "cusparse_ms": None, + "cusparse_error": None, "err": None, "status": "ERROR", "error": str(exc), @@ -694,6 +723,7 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "compute_ms", "pytorch_ms", "cusparse_ms", + "cusparse_error", "err", "status", "error", From 8e42723210bb7bce35f82db30b85bced0619b57a Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Sat, 11 Jul 2026 21:39:23 +0800 Subject: [PATCH 08/13] spmv_bsr --- src/flagsparse/sparse_operations/spmv_bsr.py | 8 +- tests/pytest/test_spmv_bsr_accuracy.py | 45 ++++-- tests/test_spmv_bsr.py | 151 +++++++++++++++++-- 3 files changed, 175 insertions(+), 29 deletions(-) diff --git a/src/flagsparse/sparse_operations/spmv_bsr.py b/src/flagsparse/sparse_operations/spmv_bsr.py index 3d3785d..a78bf82 100644 --- a/src/flagsparse/sparse_operations/spmv_bsr.py +++ b/src/flagsparse/sparse_operations/spmv_bsr.py @@ -251,8 +251,10 @@ def _prepare_spmv_bsr_matrix(data, indices, indptr, shape, block_dim): raise ValueError("block_dim must be greater than 1 for BSR SpMV") if data.shape[1] != block_dim or data.shape[2] != block_dim: raise ValueError("data block dimensions must match block_dim") - n_block_rows = (n_rows + block_dim - 1) // block_dim - n_block_cols = (n_cols + block_dim - 1) // block_dim + if n_rows % block_dim != 0 or n_cols % block_dim != 0: + raise ValueError("shape is not divisible by block_dim for standard BSR") + n_block_rows = n_rows // block_dim + n_block_cols = n_cols // block_dim if indptr.numel() != n_block_rows + 1: raise ValueError( f"indptr length must be n_block_rows+1={n_block_rows + 1}, got {indptr.numel()}" @@ -561,6 +563,8 @@ def flagsparse_spmv_bsr( meta = { "op": _spmv_bsr_op_to_name(op_code), "block_dim": prepared.block_dim, + "n_block_rows": prepared.n_block_rows, + "n_block_cols": prepared.n_block_cols, "nnzb": prepared.nnzb, "stored_nnz": prepared.stored_nnz, "symbolic_ms": 0.0 if do_timing else None, diff --git a/tests/pytest/test_spmv_bsr_accuracy.py b/tests/pytest/test_spmv_bsr_accuracy.py index ff027e9..504097c 100644 --- a/tests/pytest/test_spmv_bsr_accuracy.py +++ b/tests/pytest/test_spmv_bsr_accuracy.py @@ -5,11 +5,11 @@ from flagsparse import flagsparse_spmv_bsr, prepare_spmv_bsr from tests.pytest.accuracy_utils import close_tolerances -from tests.pytest.param_shapes import SPMV_MN_SHAPES spmv_bsr_mod = importlib.import_module("flagsparse.sparse_operations.spmv_bsr") pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +BSR_MN_SHAPES = ((8, 8), (16, 32), (64, 96)) def _value_dtype_cases(): @@ -49,7 +49,9 @@ def _reference_dtype(dtype): def _dense_to_bsr(dense, index_dtype, block_dim): device = dense.device M, N = dense.shape - n_block_rows = (M + block_dim - 1) // block_dim + if M % block_dim != 0 or N % block_dim != 0: + raise ValueError("shape is not divisible by block_dim for standard BSR") + n_block_rows = M // block_dim rows, cols = torch.nonzero(dense != 0, as_tuple=True) blocks = {} for row, col in zip(rows.tolist(), cols.tolist()): @@ -112,7 +114,7 @@ def _assert_close(actual, expected, dtype): @pytest.mark.spmv_bsr -@pytest.mark.parametrize("M, N", SPMV_MN_SHAPES) +@pytest.mark.parametrize("M, N", BSR_MN_SHAPES) @pytest.mark.parametrize( "name,dtype", _value_dtype_cases(), ids=[c[0] for c in _value_dtype_cases()] ) @@ -143,7 +145,7 @@ def test_spmv_bsr_matches_dense_reference(M, N, name, dtype, index_dtype, block_ @pytest.mark.spmv_bsr def test_spmv_bsr_prepared_path_matches_dense_reference(): device = torch.device("cuda") - M, N = 8, 10 + M, N = 8, 12 dtype = torch.complex64 block_dim = 4 data, indices, indptr, dense = _random_bsr_mn( @@ -161,22 +163,22 @@ def test_spmv_bsr_prepared_path_matches_dense_reference(): def test_spmv_bsr_unsupported_ops_are_rejected(op): device = torch.device("cuda") data, indices, indptr, _dense = _random_bsr_mn( - 8, 10, torch.float32, torch.int32, 2, device + 8, 12, torch.float32, torch.int32, 2, device ) x = torch.randn(8, dtype=torch.float32, device=device) with pytest.raises(NotImplementedError, match="only supports op='non'"): - flagsparse_spmv_bsr(data, indices, indptr, x, shape=(8, 10), block_dim=2, op=op) + flagsparse_spmv_bsr(data, indices, indptr, x, shape=(8, 12), block_dim=2, op=op) @pytest.mark.spmv_bsr def test_spmv_bsr_x_length_mismatch_rejected(): device = torch.device("cuda") data, indices, indptr, _dense = _random_bsr_mn( - 8, 10, torch.float32, torch.int32, 2, device + 8, 12, torch.float32, torch.int32, 2, device ) - prepared = prepare_spmv_bsr(data, indices, indptr, (8, 10), 2) + prepared = prepare_spmv_bsr(data, indices, indptr, (8, 12), 2) x = torch.randn(8, dtype=torch.float32, device=device) - with pytest.raises(ValueError, match="x length must be 10"): + with pytest.raises(ValueError, match="x length must be 12"): flagsparse_spmv_bsr(x=x, prepared=prepared) @@ -184,9 +186,9 @@ def test_spmv_bsr_x_length_mismatch_rejected(): def test_spmv_bsr_int64_auto_fallback_to_int32(monkeypatch): device = torch.device("cuda") data, indices, indptr, dense = _random_bsr_mn( - 12, 9, torch.float32, torch.int64, 2, device + 12, 10, torch.float32, torch.int64, 2, device ) - x = torch.randn(9, dtype=torch.float32, device=device) + x = torch.randn(10, dtype=torch.float32, device=device) ref = dense.to(torch.float64) @ x.to(torch.float64) state = {"forced_once": False} original = spmv_bsr_mod._triton_spmv_bsr_kernel @@ -203,7 +205,7 @@ def fail_int64_once(prepared, x_in, op_code): indices, indptr, x, - shape=(12, 9), + shape=(12, 10), block_dim=2, index_fallback_policy="auto", ) @@ -215,9 +217,9 @@ def fail_int64_once(prepared, x_in, op_code): def test_spmv_bsr_int64_strict_no_fallback(monkeypatch): device = torch.device("cuda") data, indices, indptr, _dense = _random_bsr_mn( - 12, 9, torch.float32, torch.int64, 2, device + 12, 10, torch.float32, torch.int64, 2, device ) - x = torch.randn(9, dtype=torch.float32, device=device) + x = torch.randn(10, dtype=torch.float32, device=device) original = spmv_bsr_mod._triton_spmv_bsr_kernel def fail_int64(prepared, x_in, op_code): @@ -232,7 +234,20 @@ def fail_int64(prepared, x_in, op_code): indices, indptr, x, - shape=(12, 9), + shape=(12, 10), block_dim=2, index_fallback_policy="strict", ) + + +@pytest.mark.spmv_bsr +def test_spmv_bsr_rejects_non_divisible_shape(): + device = torch.device("cuda") + data = torch.ones((1, 4, 4), dtype=torch.float32, device=device) + indices = torch.tensor([0], dtype=torch.int32, device=device) + indptr = torch.tensor([0, 1], dtype=torch.int32, device=device) + x = torch.randn(8, dtype=torch.float32, device=device) + with pytest.raises(ValueError, match="shape is not divisible by block_dim"): + prepare_spmv_bsr(data, indices, indptr, (7, 8), 4) + with pytest.raises(ValueError, match="shape is not divisible by block_dim"): + flagsparse_spmv_bsr(data, indices, indptr, x, shape=(7, 8), block_dim=4) diff --git a/tests/test_spmv_bsr.py b/tests/test_spmv_bsr.py index c8f543b..2c4d16f 100644 --- a/tests/test_spmv_bsr.py +++ b/tests/test_spmv_bsr.py @@ -6,6 +6,7 @@ import math import os import sys +import warnings from pathlib import Path import torch @@ -35,6 +36,25 @@ ITERS = 50 +def _cupy_bsr_unavailable_reason(): + if cp is None or cpx_sparse is None: + return "CuPy/cupyx.scipy.sparse is not available" + if not hasattr(cpx_sparse, "bsr_matrix"): + return "CuPy cupyx.scipy.sparse has no bsr_matrix baseline" + return None + + +def _print_baseline_notes(run_cusparse=True): + print( + "PyTorch baseline: PT(ms) uses torch.sparse_bsr_tensor when shape is divisible by block_dim; " + "unsupported shapes are recorded in pytorch_error CSV." + ) + if run_cusparse: + reason = _cupy_bsr_unavailable_reason() + if reason: + print(f"CuPy baseline: unavailable for BSR ({reason}); CU(ms)=N/A.") + + def _dtype_name(dtype): return str(dtype).replace("torch.", "") @@ -134,7 +154,12 @@ def _zero_value(dtype): def _choose_auto_block_dim(entries, shape): n_rows, n_cols = shape nnz = max(1, len(entries)) + first_divisible = None for block_dim in (16, 8, 4, 2): + if n_rows % block_dim != 0 or n_cols % block_dim != 0: + continue + if first_divisible is None: + first_divisible = block_dim blocks = { (int(row) // block_dim, int(col) // block_dim) for row, col in entries.keys() @@ -142,12 +167,18 @@ def _choose_auto_block_dim(entries, shape): stored = len(blocks) * block_dim * block_dim if stored <= 2.0 * nnz: return block_dim - return 4 if max(n_rows, n_cols) >= 4 else 2 + return first_divisible + + +def _shape_divisible_by_block_dim(shape, block_dim): + return int(shape[0]) % int(block_dim) == 0 and int(shape[1]) % int(block_dim) == 0 def _entries_to_bsr_torch(entries, shape, dtype, index_dtype, block_dim, device): n_rows, n_cols = int(shape[0]), int(shape[1]) block_dim = int(block_dim) + if not _shape_divisible_by_block_dim((n_rows, n_cols), block_dim): + raise ValueError("shape is not divisible by block_dim for standard BSR") blocks = {} for (row, col), value in entries.items(): brow = int(row) // block_dim @@ -159,7 +190,7 @@ def _entries_to_bsr_torch(entries, shape, dtype, index_dtype, block_dim, device) [_zero_value(dtype) for _ in range(block_dim * block_dim)], ) block[inner_row * block_dim + inner_col] += _mtx_value_for_dtype(value, dtype) - n_block_rows = (n_rows + block_dim - 1) // block_dim + n_block_rows = n_rows // block_dim rows = [[] for _ in range(n_block_rows)] for key in sorted(blocks): rows[key[0]].append(key) @@ -287,15 +318,35 @@ def _time_flagsparse_bsr(data, indices, indptr, x, shape, block_dim, warmup, ite def _time_pytorch(data, indices, indptr, x, shape, block_dim, warmup, iters): - A = _bsr_to_torch_coo(data, indices, indptr, shape, block_dim) + if int(shape[0]) % int(block_dim) != 0 or int(shape[1]) % int(block_dim) != 0: + return ( + None, + "PyTorch BSR baseline requires both matrix dimensions to be divisible by block_dim", + ) + with warnings.catch_warnings(): + warnings.filterwarnings( + "ignore", + message="Sparse BSR tensor support is in beta state.*", + category=UserWarning, + ) + A = torch.sparse_bsr_tensor( + indptr, + indices, + data, + size=shape, + device=data.device, + dtype=data.dtype, + ) fn = lambda: torch.sparse.mm(A, x.unsqueeze(1)).squeeze(1) _, ms = _cuda_event_benchmark(fn, warmup, iters) - return ms + return ms, None def _time_cusparse(data, indices, indptr, x, shape, block_dim, warmup, iters): if cp is None or cpx_sparse is None: return None, "CuPy/cupyx.scipy.sparse is not available" + if not hasattr(cpx_sparse, "bsr_matrix"): + return None, "CuPy cupyx.scipy.sparse has no bsr_matrix baseline" if data.dtype not in (torch.float32, torch.float64, torch.complex64, torch.complex128): return None, f"unsupported cuSPARSE dtype: {_dtype_name(data.dtype)}" @@ -384,9 +435,6 @@ def _print_row(row, timing=False): error = row.get("error") if error: print(f" error: {str(error)[:240]}") - cusparse_error = row.get("cusparse_error") - if cusparse_error and row.get("cusparse_ms") is None: - print(f" cusparse: {str(cusparse_error)[:240]}") def _base_row( @@ -422,6 +470,7 @@ def _base_row( "process_gpu_ms": None, "compute_ms": None, "pytorch_ms": None, + "pytorch_error": None, "cusparse_ms": None, "cusparse_error": None, "err": None, @@ -430,6 +479,39 @@ def _base_row( } +def _skip_row(matrix_name, dtype, index_dtype, op, shape, block_dim, logical_nnz, error): + try: + block_dim_value = int(block_dim) + except (TypeError, ValueError): + block_dim_value = block_dim + return { + "matrix": matrix_name, + "value_dtype": _dtype_name(dtype), + "index_dtype": _dtype_name(index_dtype), + "op": op, + "block_dim": block_dim_value, + "out_size": int(shape[0]) if op == "non" else "UNSUP", + "n_rows": int(shape[0]), + "n_cols": int(shape[1]), + "nnzb": "SKIP", + "logical_nnz": max(1, int(logical_nnz)), + "stored_nnz": "SKIP", + "padding_ratio": "SKIP", + "bsr_ms": None, + "bsr_gpu_ms": None, + "process_cpu_ms": 0.0, + "process_gpu_ms": None, + "compute_ms": None, + "pytorch_ms": None, + "pytorch_error": None, + "cusparse_ms": None, + "cusparse_error": None, + "err": None, + "status": "SKIP", + "error": error, + } + + def _run_one_case( data, indices, @@ -487,9 +569,18 @@ def _run_one_case( row["error"] = f"reference failed after BSR run: {exc}" return row try: - row["pytorch_ms"] = _time_pytorch(data, indices, indptr, x, shape, block_dim, warmup, iters) - except Exception: - pass + row["pytorch_ms"], row["pytorch_error"] = _time_pytorch( + data, + indices, + indptr, + x, + shape, + block_dim, + warmup, + iters, + ) + except Exception as exc: + row["pytorch_error"] = str(exc) if run_cusparse: try: row["cusparse_ms"], row["cusparse_error"] = _time_cusparse( @@ -575,7 +666,8 @@ def add_entry(r, c, value): def _resolve_block_dims(block_dims, entries, shape): if block_dims == ["auto"]: - return [_choose_auto_block_dim(entries, shape)] + block_dim = _choose_auto_block_dim(entries, shape) + return [] if block_dim is None else [block_dim] return block_dims @@ -592,9 +684,12 @@ def run_synthetic(value_dtypes=None, index_dtypes=None, block_dims=None, ops=Non print("FLAGSPARSE SpMV BSR BENCHMARK (native BSR Triton)") print("=" * 140) print("Timing policy: bsr_ms = process_cpu_ms + bsr_gpu_ms; BSR construction is setup.") + _print_baseline_notes(run_cusparse=run_cusparse) for dtype in value_dtypes: for index_dtype in index_dtypes: for block_dim in block_dims: + if block_dim == "auto": + block_dim = 4 for op in ops: print(_sep(timing)) print(f"dtype: {_dtype_name(dtype)} | index_dtype: {_dtype_name(index_dtype)} | block_dim: {block_dim} | op: {op}") @@ -637,6 +732,7 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim block_dims = list(DEFAULT_BLOCK_DIMS) if block_dims is None else block_dims ops = SUPPORTED_OPS if ops is None else ops rows = [] + _print_baseline_notes(run_cusparse=run_cusparse) for dtype in value_dtypes: for index_dtype in index_dtypes: for op in ops: @@ -648,7 +744,36 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim for path in mtx_paths: try: entries, shape = load_mtx_entries(path) - for block_dim in _resolve_block_dims(block_dims, entries, shape): + resolved_block_dims = _resolve_block_dims(block_dims, entries, shape) + if not resolved_block_dims: + row = _skip_row( + os.path.basename(path), + dtype, + index_dtype, + op, + shape, + "auto", + len(entries), + "shape is not divisible by any supported standard BSR block_dim", + ) + rows.append(row) + _print_row(row, timing=timing) + continue + for block_dim in resolved_block_dims: + if not _shape_divisible_by_block_dim(shape, int(block_dim)): + row = _skip_row( + os.path.basename(path), + dtype, + index_dtype, + op, + shape, + int(block_dim), + len(entries), + "shape is not divisible by block_dim for standard BSR", + ) + rows.append(row) + _print_row(row, timing=timing) + continue data, indices, indptr = _entries_to_bsr_torch( entries, shape, dtype, index_dtype, int(block_dim), device ) @@ -694,6 +819,7 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "process_gpu_ms": None, "compute_ms": None, "pytorch_ms": None, + "pytorch_error": None, "cusparse_ms": None, "cusparse_error": None, "err": None, @@ -722,6 +848,7 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "process_gpu_ms", "compute_ms", "pytorch_ms", + "pytorch_error", "cusparse_ms", "cusparse_error", "err", From 9ec861c55588014b532519b5cb6fe2c8018a8862 Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Sat, 11 Jul 2026 22:28:55 +0800 Subject: [PATCH 09/13] spmv_bsr --- src/flagsparse/sparse_operations/spmv_bsr.py | 6 +-- tests/pytest/test_spmv_bsr_accuracy.py | 31 ++++++----- tests/test_spmv_bsr.py | 57 ++++---------------- 3 files changed, 28 insertions(+), 66 deletions(-) diff --git a/src/flagsparse/sparse_operations/spmv_bsr.py b/src/flagsparse/sparse_operations/spmv_bsr.py index a78bf82..6e88fd7 100644 --- a/src/flagsparse/sparse_operations/spmv_bsr.py +++ b/src/flagsparse/sparse_operations/spmv_bsr.py @@ -251,10 +251,8 @@ def _prepare_spmv_bsr_matrix(data, indices, indptr, shape, block_dim): raise ValueError("block_dim must be greater than 1 for BSR SpMV") if data.shape[1] != block_dim or data.shape[2] != block_dim: raise ValueError("data block dimensions must match block_dim") - if n_rows % block_dim != 0 or n_cols % block_dim != 0: - raise ValueError("shape is not divisible by block_dim for standard BSR") - n_block_rows = n_rows // block_dim - n_block_cols = n_cols // block_dim + n_block_rows = (n_rows + block_dim - 1) // block_dim + n_block_cols = (n_cols + block_dim - 1) // block_dim if indptr.numel() != n_block_rows + 1: raise ValueError( f"indptr length must be n_block_rows+1={n_block_rows + 1}, got {indptr.numel()}" diff --git a/tests/pytest/test_spmv_bsr_accuracy.py b/tests/pytest/test_spmv_bsr_accuracy.py index 504097c..5696c1a 100644 --- a/tests/pytest/test_spmv_bsr_accuracy.py +++ b/tests/pytest/test_spmv_bsr_accuracy.py @@ -9,7 +9,7 @@ spmv_bsr_mod = importlib.import_module("flagsparse.sparse_operations.spmv_bsr") pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") -BSR_MN_SHAPES = ((8, 8), (16, 32), (64, 96)) +BSR_MN_SHAPES = ((7, 8), (12, 9), (16, 32), (64, 96)) def _value_dtype_cases(): @@ -49,9 +49,7 @@ def _reference_dtype(dtype): def _dense_to_bsr(dense, index_dtype, block_dim): device = dense.device M, N = dense.shape - if M % block_dim != 0 or N % block_dim != 0: - raise ValueError("shape is not divisible by block_dim for standard BSR") - n_block_rows = M // block_dim + n_block_rows = (M + block_dim - 1) // block_dim rows, cols = torch.nonzero(dense != 0, as_tuple=True) blocks = {} for row, col in zip(rows.tolist(), cols.tolist()): @@ -145,7 +143,7 @@ def test_spmv_bsr_matches_dense_reference(M, N, name, dtype, index_dtype, block_ @pytest.mark.spmv_bsr def test_spmv_bsr_prepared_path_matches_dense_reference(): device = torch.device("cuda") - M, N = 8, 12 + M, N = 7, 8 dtype = torch.complex64 block_dim = 4 data, indices, indptr, dense = _random_bsr_mn( @@ -186,9 +184,9 @@ def test_spmv_bsr_x_length_mismatch_rejected(): def test_spmv_bsr_int64_auto_fallback_to_int32(monkeypatch): device = torch.device("cuda") data, indices, indptr, dense = _random_bsr_mn( - 12, 10, torch.float32, torch.int64, 2, device + 12, 9, torch.float32, torch.int64, 2, device ) - x = torch.randn(10, dtype=torch.float32, device=device) + x = torch.randn(9, dtype=torch.float32, device=device) ref = dense.to(torch.float64) @ x.to(torch.float64) state = {"forced_once": False} original = spmv_bsr_mod._triton_spmv_bsr_kernel @@ -205,7 +203,7 @@ def fail_int64_once(prepared, x_in, op_code): indices, indptr, x, - shape=(12, 10), + shape=(12, 9), block_dim=2, index_fallback_policy="auto", ) @@ -241,13 +239,14 @@ def fail_int64(prepared, x_in, op_code): @pytest.mark.spmv_bsr -def test_spmv_bsr_rejects_non_divisible_shape(): +def test_spmv_bsr_non_divisible_shape_matches_dense_reference(): device = torch.device("cuda") - data = torch.ones((1, 4, 4), dtype=torch.float32, device=device) - indices = torch.tensor([0], dtype=torch.int32, device=device) - indptr = torch.tensor([0, 1], dtype=torch.int32, device=device) + M, N = 7, 8 + data, indices, indptr, dense = _random_bsr_mn( + M, N, torch.float32, torch.int32, 4, device + ) x = torch.randn(8, dtype=torch.float32, device=device) - with pytest.raises(ValueError, match="shape is not divisible by block_dim"): - prepare_spmv_bsr(data, indices, indptr, (7, 8), 4) - with pytest.raises(ValueError, match="shape is not divisible by block_dim"): - flagsparse_spmv_bsr(data, indices, indptr, x, shape=(7, 8), block_dim=4) + ref = dense.to(torch.float64) @ x.to(torch.float64) + prepared = prepare_spmv_bsr(data, indices, indptr, (M, N), 4) + out = flagsparse_spmv_bsr(x=x, prepared=prepared) + _assert_close(out, ref.to(torch.float32), torch.float32) diff --git a/tests/test_spmv_bsr.py b/tests/test_spmv_bsr.py index 2c4d16f..1389913 100644 --- a/tests/test_spmv_bsr.py +++ b/tests/test_spmv_bsr.py @@ -46,8 +46,11 @@ def _cupy_bsr_unavailable_reason(): def _print_baseline_notes(run_cusparse=True): print( - "PyTorch baseline: PT(ms) uses torch.sparse_bsr_tensor when shape is divisible by block_dim; " - "unsupported shapes are recorded in pytorch_error CSV." + "FlagSparse BSR follows AlphaSparse/cuSPARSE-style ceil block grid with boundary masks." + ) + print( + "PyTorch baseline: PT(ms) uses torch.sparse_bsr_tensor only when shape is divisible by block_dim; " + "otherwise PT(ms)=N/A and pytorch_error records the PyTorch limitation." ) if run_cusparse: reason = _cupy_bsr_unavailable_reason() @@ -152,14 +155,9 @@ def _zero_value(dtype): def _choose_auto_block_dim(entries, shape): - n_rows, n_cols = shape nnz = max(1, len(entries)) - first_divisible = None + best = None for block_dim in (16, 8, 4, 2): - if n_rows % block_dim != 0 or n_cols % block_dim != 0: - continue - if first_divisible is None: - first_divisible = block_dim blocks = { (int(row) // block_dim, int(col) // block_dim) for row, col in entries.keys() @@ -167,18 +165,14 @@ def _choose_auto_block_dim(entries, shape): stored = len(blocks) * block_dim * block_dim if stored <= 2.0 * nnz: return block_dim - return first_divisible - - -def _shape_divisible_by_block_dim(shape, block_dim): - return int(shape[0]) % int(block_dim) == 0 and int(shape[1]) % int(block_dim) == 0 + if best is None or stored < best[0]: + best = (stored, block_dim) + return best[1] if best is not None else 2 def _entries_to_bsr_torch(entries, shape, dtype, index_dtype, block_dim, device): n_rows, n_cols = int(shape[0]), int(shape[1]) block_dim = int(block_dim) - if not _shape_divisible_by_block_dim((n_rows, n_cols), block_dim): - raise ValueError("shape is not divisible by block_dim for standard BSR") blocks = {} for (row, col), value in entries.items(): brow = int(row) // block_dim @@ -190,7 +184,7 @@ def _entries_to_bsr_torch(entries, shape, dtype, index_dtype, block_dim, device) [_zero_value(dtype) for _ in range(block_dim * block_dim)], ) block[inner_row * block_dim + inner_col] += _mtx_value_for_dtype(value, dtype) - n_block_rows = n_rows // block_dim + n_block_rows = (n_rows + block_dim - 1) // block_dim rows = [[] for _ in range(n_block_rows)] for key in sorted(blocks): rows[key[0]].append(key) @@ -666,8 +660,7 @@ def add_entry(r, c, value): def _resolve_block_dims(block_dims, entries, shape): if block_dims == ["auto"]: - block_dim = _choose_auto_block_dim(entries, shape) - return [] if block_dim is None else [block_dim] + return [_choose_auto_block_dim(entries, shape)] return block_dims @@ -745,35 +738,7 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim try: entries, shape = load_mtx_entries(path) resolved_block_dims = _resolve_block_dims(block_dims, entries, shape) - if not resolved_block_dims: - row = _skip_row( - os.path.basename(path), - dtype, - index_dtype, - op, - shape, - "auto", - len(entries), - "shape is not divisible by any supported standard BSR block_dim", - ) - rows.append(row) - _print_row(row, timing=timing) - continue for block_dim in resolved_block_dims: - if not _shape_divisible_by_block_dim(shape, int(block_dim)): - row = _skip_row( - os.path.basename(path), - dtype, - index_dtype, - op, - shape, - int(block_dim), - len(entries), - "shape is not divisible by block_dim for standard BSR", - ) - rows.append(row) - _print_row(row, timing=timing) - continue data, indices, indptr = _entries_to_bsr_torch( entries, shape, dtype, index_dtype, int(block_dim), device ) From 58e728f0c6fc07c88b8ddc8b65c1241398914ac3 Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Sun, 12 Jul 2026 00:14:21 +0800 Subject: [PATCH 10/13] spmv_bsr --- src/flagsparse/sparse_operations/spmv_bsr.py | 44 +++-- tests/pytest/test_spmv_bsr_accuracy.py | 41 +++- tests/test_spmv_bsr.py | 194 +++++++++++++++++-- 3 files changed, 245 insertions(+), 34 deletions(-) diff --git a/src/flagsparse/sparse_operations/spmv_bsr.py b/src/flagsparse/sparse_operations/spmv_bsr.py index 6e88fd7..bec4a8b 100644 --- a/src/flagsparse/sparse_operations/spmv_bsr.py +++ b/src/flagsparse/sparse_operations/spmv_bsr.py @@ -83,6 +83,8 @@ class PreparedBsrSpmv: "shape", "n_rows", "n_cols", + "padded_n_rows", + "padded_n_cols", "block_dim", "n_block_rows", "n_block_cols", @@ -127,6 +129,8 @@ def __init__( self.block_dim = int(block_dim) self.n_block_rows = int(n_block_rows) self.n_block_cols = int(n_block_cols) + self.padded_n_rows = self.n_block_rows * self.block_dim + self.padded_n_cols = self.n_block_cols * self.block_dim self.nnzb = int(data.shape[0]) self.stored_nnz = int(data.numel()) if block_row_lengths is None: @@ -161,8 +165,6 @@ def _spmv_bsr_non_real_kernel( if brow >= n_block_rows: return row = brow * BLOCK_DIM + inner_row - if row >= n_rows: - return start = tl.load(indptr_ptr + brow) end = tl.load(indptr_ptr + brow + 1) offs = start + SEG * BLOCK_NNZ + tl.arange(0, BLOCK_NNZ) @@ -175,7 +177,7 @@ def _spmv_bsr_non_real_kernel( ) * 0 for inner_col in tl.static_range(0, BLOCK_DIM): col = bcols * BLOCK_DIM + inner_col - valid = mask & (col < n_cols) + valid = mask vals = tl.load( data_ptr + offs * BLOCK_DIM * BLOCK_DIM + inner_row * BLOCK_DIM + inner_col, mask=mask, @@ -205,8 +207,6 @@ def _spmv_bsr_non_complex_kernel( if brow >= n_block_rows: return row = brow * BLOCK_DIM + inner_row - if row >= n_rows: - return start = tl.load(indptr_ptr + brow) end = tl.load(indptr_ptr + brow + 1) offs = start + SEG * BLOCK_NNZ + tl.arange(0, BLOCK_NNZ) @@ -224,7 +224,7 @@ def _spmv_bsr_non_complex_kernel( ) * 0 for inner_col in tl.static_range(0, BLOCK_DIM): col = bcols * BLOCK_DIM + inner_col - valid = mask & (col < n_cols) + valid = mask elem = offs * BLOCK_DIM * BLOCK_DIM + inner_row * BLOCK_DIM + inner_col a_re = tl.load(data_ri_ptr + elem * 2, mask=mask, other=0.0) a_im = tl.load(data_ri_ptr + elem * 2 + 1, mask=mask, other=0.0) @@ -367,18 +367,30 @@ def _validate_spmv_bsr_x(x, prepared, op_code): raise ValueError("x must be a CUDA tensor") if x.dtype != prepared.data.dtype: raise TypeError("x dtype must match sparse matrix dtype") - expected = prepared.n_rows if _spmv_bsr_op_transposes(op_code) else prepared.n_cols - if x.numel() != expected: - raise ValueError(f"x length must be {expected}, got {x.numel()}") + logical_expected = prepared.n_rows if _spmv_bsr_op_transposes(op_code) else prepared.n_cols + padded_expected = ( + prepared.padded_n_rows + if _spmv_bsr_op_transposes(op_code) + else prepared.padded_n_cols + ) + if x.numel() not in (logical_expected, padded_expected): + raise ValueError( + f"x length must be {logical_expected} or padded length {padded_expected}, got {x.numel()}" + ) if x.device != prepared.data.device: raise ValueError("x must be on the same device as sparse matrix data") - return x.contiguous() + x = x.contiguous() + if x.numel() == padded_expected: + return x + padded = torch.zeros(padded_expected, dtype=x.dtype, device=x.device) + padded[: x.numel()].copy_(x) + return padded def _triton_spmv_bsr_kernel(prepared, x, op_code): _ensure_spmv_bsr_supported_op(op_code) dtype = prepared.data.dtype - y = torch.zeros(prepared.n_rows, dtype=dtype, device=prepared.data.device) + y = torch.zeros(prepared.padded_n_rows, dtype=dtype, device=prepared.data.device) if prepared.nnzb == 0: return y for seg in range(prepared.max_segments): @@ -393,8 +405,8 @@ def _triton_spmv_bsr_kernel(prepared, x, op_code): prepared.kernel_indptr, x_ri, y_ri, - prepared.n_rows, - prepared.n_cols, + prepared.padded_n_rows, + prepared.padded_n_cols, prepared.n_block_rows, BLOCK_DIM=prepared.block_dim, BLOCK_NNZ=prepared.block_nnz, @@ -407,8 +419,8 @@ def _triton_spmv_bsr_kernel(prepared, x, op_code): prepared.kernel_indptr, x, y, - prepared.n_rows, - prepared.n_cols, + prepared.padded_n_rows, + prepared.padded_n_cols, prepared.n_block_rows, BLOCK_DIM=prepared.block_dim, BLOCK_NNZ=prepared.block_nnz, @@ -561,6 +573,8 @@ def flagsparse_spmv_bsr( meta = { "op": _spmv_bsr_op_to_name(op_code), "block_dim": prepared.block_dim, + "logical_shape": prepared.shape, + "padded_shape": (prepared.padded_n_rows, prepared.padded_n_cols), "n_block_rows": prepared.n_block_rows, "n_block_cols": prepared.n_block_cols, "nnzb": prepared.nnzb, diff --git a/tests/pytest/test_spmv_bsr_accuracy.py b/tests/pytest/test_spmv_bsr_accuracy.py index 5696c1a..abc98c8 100644 --- a/tests/pytest/test_spmv_bsr_accuracy.py +++ b/tests/pytest/test_spmv_bsr_accuracy.py @@ -103,6 +103,14 @@ def _make_x(length, dtype, device): return _random_values((length,), dtype, device) +def _padded_rows(M, block_dim): + return ((int(M) + int(block_dim) - 1) // int(block_dim)) * int(block_dim) + + +def _padded_cols(N, block_dim): + return ((int(N) + int(block_dim) - 1) // int(block_dim)) * int(block_dim) + + def _assert_close(actual, expected, dtype): rtol, atol = close_tolerances(dtype) ref_dtype = _reference_dtype(dtype) @@ -137,7 +145,8 @@ def test_spmv_bsr_matches_dense_reference(M, N, name, dtype, index_dtype, block_ block_dim=block_dim, index_fallback_policy="auto", ) - _assert_close(out, ref, dtype) + assert out.numel() == _padded_rows(M, block_dim) + _assert_close(out[:M], ref, dtype) @pytest.mark.spmv_bsr @@ -153,7 +162,8 @@ def test_spmv_bsr_prepared_path_matches_dense_reference(): x = _make_x(N, dtype, device) ref = (dense.to(torch.complex128) @ x.to(torch.complex128)).to(dtype) out = flagsparse_spmv_bsr(x=x, prepared=prepared) - _assert_close(out, ref, dtype) + assert out.numel() == _padded_rows(M, block_dim) + _assert_close(out[:M], ref, dtype) @pytest.mark.spmv_bsr @@ -208,7 +218,8 @@ def fail_int64_once(prepared, x_in, op_code): index_fallback_policy="auto", ) assert state["forced_once"] - _assert_close(out, ref.to(torch.float32), torch.float32) + assert out.numel() == _padded_rows(12, 2) + _assert_close(out[:12], ref.to(torch.float32), torch.float32) @pytest.mark.spmv_bsr @@ -249,4 +260,26 @@ def test_spmv_bsr_non_divisible_shape_matches_dense_reference(): ref = dense.to(torch.float64) @ x.to(torch.float64) prepared = prepare_spmv_bsr(data, indices, indptr, (M, N), 4) out = flagsparse_spmv_bsr(x=x, prepared=prepared) - _assert_close(out, ref.to(torch.float32), torch.float32) + assert out.numel() == _padded_rows(M, 4) + _assert_close(out[:M], ref.to(torch.float32), torch.float32) + + +@pytest.mark.spmv_bsr +def test_spmv_bsr_accepts_logical_or_padded_x(): + device = torch.device("cuda") + M, N = 7, 9 + block_dim = 4 + data, indices, indptr, dense = _random_bsr_mn( + M, N, torch.float32, torch.int32, block_dim, device + ) + x = torch.randn(N, dtype=torch.float32, device=device) + x_padded = torch.zeros(_padded_cols(N, block_dim), dtype=torch.float32, device=device) + x_padded[:N].copy_(x) + ref = dense.to(torch.float64) @ x.to(torch.float64) + prepared = prepare_spmv_bsr(data, indices, indptr, (M, N), block_dim) + out_logical_x = flagsparse_spmv_bsr(x=x, prepared=prepared) + out_padded_x = flagsparse_spmv_bsr(x=x_padded, prepared=prepared) + assert out_logical_x.numel() == _padded_rows(M, block_dim) + assert out_padded_x.numel() == _padded_rows(M, block_dim) + _assert_close(out_logical_x, out_padded_x, torch.float32) + _assert_close(out_logical_x[:M], ref.to(torch.float32), torch.float32) diff --git a/tests/test_spmv_bsr.py b/tests/test_spmv_bsr.py index 1389913..952b3c0 100644 --- a/tests/test_spmv_bsr.py +++ b/tests/test_spmv_bsr.py @@ -46,11 +46,11 @@ def _cupy_bsr_unavailable_reason(): def _print_baseline_notes(run_cusparse=True): print( - "FlagSparse BSR follows AlphaSparse/cuSPARSE-style ceil block grid with boundary masks." + "FlagSparse BSR follows AlphaSparse/cuSPARSE-style padded block-grid semantics; native output is padded and correctness checks slice back to logical rows." ) print( "PyTorch baseline: PT(ms) uses torch.sparse_bsr_tensor only when shape is divisible by block_dim; " - "otherwise PT(ms)=N/A and pytorch_error records the PyTorch limitation." + "PTPad(ms) uses padded shape and slices back to logical rows for diagnostics." ) if run_cusparse: reason = _cupy_bsr_unavailable_reason() @@ -170,6 +170,26 @@ def _choose_auto_block_dim(entries, shape): return best[1] if best is not None else 2 +def _padded_shape(shape, block_dim): + block_dim = int(block_dim) + n_rows, n_cols = int(shape[0]), int(shape[1]) + return ( + ((n_rows + block_dim - 1) // block_dim) * block_dim, + ((n_cols + block_dim - 1) // block_dim) * block_dim, + ) + + +def _pad_vector(x, length): + length = int(length) + if x.numel() == length: + return x.contiguous() + if x.numel() > length: + raise ValueError(f"cannot pad vector of length {x.numel()} to shorter length {length}") + out = torch.zeros(length, dtype=x.dtype, device=x.device) + out[: x.numel()].copy_(x) + return out + + def _entries_to_bsr_torch(entries, shape, dtype, index_dtype, block_dim, device): n_rows, n_cols = int(shape[0]), int(shape[1]) block_dim = int(block_dim) @@ -270,12 +290,37 @@ def _pytorch_reference(data, indices, indptr, x, shape, dtype, block_dim): return torch.sparse.mm(A, x.to(ref_dtype).unsqueeze(1)).squeeze(1).to(dtype) -def _allclose_error_ratio(actual, expected, atol, rtol): +def _error_stats(actual, expected, atol, rtol): if expected.numel() == 0: - return 0.0 + return { + "ratio": 0.0, + "max_abs": 0.0, + "max_rel": 0.0, + "index": 0, + "actual": None, + "expected": None, + } diff = torch.abs(actual - expected).to(torch.float64) denom = atol + rtol * torch.abs(expected).to(torch.float64) - return float(torch.max(diff / denom).item()) + ratio_values = diff / denom + flat_ratio = ratio_values.reshape(-1) + max_pos = int(torch.argmax(flat_ratio).item()) + expected_abs = torch.abs(expected).to(torch.float64).reshape(-1) + rel = diff.reshape(-1) / torch.clamp(expected_abs, min=1.0e-30) + actual_flat = actual.reshape(-1) + expected_flat = expected.reshape(-1) + return { + "ratio": float(flat_ratio[max_pos].item()), + "max_abs": float(diff.reshape(-1)[max_pos].item()), + "max_rel": float(rel[max_pos].item()), + "index": max_pos, + "actual": actual_flat[max_pos].detach().cpu().item(), + "expected": expected_flat[max_pos].detach().cpu().item(), + } + + +def _allclose_error_ratio(actual, expected, atol, rtol): + return _error_stats(actual, expected, atol, rtol)["ratio"] def _cuda_event_benchmark(op, warmup, iters): @@ -296,8 +341,10 @@ def _cuda_event_benchmark(op, warmup, iters): def _time_flagsparse_bsr(data, indices, indptr, x, shape, block_dim, warmup, iters, timing=False): prepared = fs.prepare_spmv_bsr(data, indices, indptr, shape, block_dim, op="non") + _padded_rows, padded_cols = _padded_shape(shape, block_dim) + x_for_bsr = _pad_vector(x, padded_cols) out, gpu_ms = _cuda_event_benchmark( - lambda: fs.flagsparse_spmv_bsr(x=x, prepared=prepared), + lambda: fs.flagsparse_spmv_bsr(x=x_for_bsr, prepared=prepared), warmup, iters, ) @@ -316,6 +363,7 @@ def _time_pytorch(data, indices, indptr, x, shape, block_dim, warmup, iters): return ( None, "PyTorch BSR baseline requires both matrix dimensions to be divisible by block_dim", + None, ) with warnings.catch_warnings(): warnings.filterwarnings( @@ -332,8 +380,30 @@ def _time_pytorch(data, indices, indptr, x, shape, block_dim, warmup, iters): dtype=data.dtype, ) fn = lambda: torch.sparse.mm(A, x.unsqueeze(1)).squeeze(1) - _, ms = _cuda_event_benchmark(fn, warmup, iters) - return ms, None + out, ms = _cuda_event_benchmark(fn, warmup, iters) + return ms, None, out + + +def _time_pytorch_padded(data, indices, indptr, x, shape, block_dim, warmup, iters): + padded_shape = _padded_shape(shape, block_dim) + padded_x = _pad_vector(x, padded_shape[1]) + with warnings.catch_warnings(): + warnings.filterwarnings( + "ignore", + message="Sparse BSR tensor support is in beta state.*", + category=UserWarning, + ) + A = torch.sparse_bsr_tensor( + indptr, + indices, + data, + size=padded_shape, + device=data.device, + dtype=data.dtype, + ) + fn = lambda: torch.sparse.mm(A, padded_x.unsqueeze(1)).squeeze(1) + out, ms = _cuda_event_benchmark(fn, warmup, iters) + return ms, None, out def _time_cusparse(data, indices, indptr, x, shape, block_dim, warmup, iters): @@ -399,9 +469,9 @@ def _status(ok): def _header(timing=False): split = f" {'ProcGPU':>9} {'Compute':>9}" if timing else "" return ( - f"{'Matrix':<28} {'Op':>5} {'BDim':>5} {'Out':>7} {'Rows':>7} {'Cols':>7} {'NNZB':>9} {'Pad':>7} " + f"{'Matrix':<28} {'Op':>5} {'BDim':>5} {'Out':>7} {'PadOut':>7} {'PadRows':>7} {'Rows':>7} {'Cols':>7} {'NNZB':>9} {'Pad':>7} " f"{'BSR(ms)':>9} {'BSRGPU':>9} {'CPUProc':>9}{split} " - f"{'PT(ms)':>9} {'CU(ms)':>9} {'BSR/PT':>8} {'BSR/CU':>8} " + f"{'PT(ms)':>9} {'PTPad':>9} {'CU(ms)':>9} {'BSR/PT':>8} {'BSR/CU':>8} " f"{'Err':>10} {'Status':>6}" ) @@ -420,15 +490,30 @@ def _print_row(row, timing=False): else "" ) print( - f"{name:<28} {row['op']:>5} {row['block_dim']:>5} {row['out_size']:>7} {row['n_rows']:>7} {row['n_cols']:>7} {row['nnzb']:>9} {row['padding_ratio']:>7} " + f"{name:<28} {row['op']:>5} {row['block_dim']:>5} {row['out_size']:>7} {row['padded_out_size']:>7} {row['pad_rows']:>7} {row['n_rows']:>7} {row['n_cols']:>7} {row['nnzb']:>9} {row['padding_ratio']:>7} " f"{_fmt(row['bsr_ms']):>9} {_fmt(row['bsr_gpu_ms']):>9} {_fmt(row['process_cpu_ms']):>9}{split} " - f"{_fmt(row['pytorch_ms']):>9} {_fmt(row['cusparse_ms']):>9} " + f"{_fmt(row['pytorch_ms']):>9} {_fmt(row['pytorch_padded_ms']):>9} {_fmt(row['cusparse_ms']):>9} " f"{_spd(row['pytorch_ms'], row['bsr_ms']):>8} {_spd(row['cusparse_ms'], row['bsr_ms']):>8} " f"{_fmt_err(row['err']):>10} {row['status']:>6}" ) error = row.get("error") if error: print(f" error: {str(error)[:240]}") + if row.get("pytorch_ms") is None and row.get("pytorch_error"): + print(f" pt: {str(row['pytorch_error'])[:240]}") + if row.get("pytorch_padded_ms") is None and row.get("pytorch_padded_error"): + print(f" pt_padded: {str(row['pytorch_padded_error'])[:240]}") + if row.get("status") == "FAIL": + print( + " debug: " + f"max_abs={_fmt_err(row.get('max_abs_err'))}, " + f"max_rel={_fmt_err(row.get('max_rel_err'))}, " + f"max_idx={row.get('max_err_index')}, " + f"actual={row.get('actual_at_max')}, " + f"expected={row.get('expected_at_max')}, " + f"pytorch_err={_fmt_err(row.get('pytorch_err'))}, " + f"pytorch_padded_err={_fmt_err(row.get('pytorch_padded_err'))}" + ) def _base_row( @@ -445,6 +530,7 @@ def _base_row( nnzb = int(data.shape[0]) stored_nnz = int(data.numel()) logical_nnz = max(1, int(logical_nnz if logical_nnz is not None else stored_nnz)) + padded_rows, _padded_cols = _padded_shape(shape, block_dim) return { "matrix": matrix_name, "value_dtype": _dtype_name(dtype), @@ -452,6 +538,8 @@ def _base_row( "op": op, "block_dim": int(block_dim), "out_size": int(shape[0]) if op == "non" else "UNSUP", + "padded_out_size": padded_rows if op == "non" else "UNSUP", + "pad_rows": max(0, padded_rows - int(shape[0])) if op == "non" else "UNSUP", "n_rows": int(shape[0]), "n_cols": int(shape[1]), "nnzb": nnzb, @@ -465,9 +553,18 @@ def _base_row( "compute_ms": None, "pytorch_ms": None, "pytorch_error": None, + "pytorch_padded_ms": None, + "pytorch_padded_error": None, + "pytorch_padded_err": None, "cusparse_ms": None, "cusparse_error": None, "err": None, + "max_abs_err": None, + "max_rel_err": None, + "max_err_index": None, + "actual_at_max": None, + "expected_at_max": None, + "pytorch_err": None, "status": status, "error": None, } @@ -478,6 +575,10 @@ def _skip_row(matrix_name, dtype, index_dtype, op, shape, block_dim, logical_nnz block_dim_value = int(block_dim) except (TypeError, ValueError): block_dim_value = block_dim + try: + padded_rows, _padded_cols = _padded_shape(shape, block_dim_value) + except Exception: + padded_rows = "SKIP" return { "matrix": matrix_name, "value_dtype": _dtype_name(dtype), @@ -485,6 +586,8 @@ def _skip_row(matrix_name, dtype, index_dtype, op, shape, block_dim, logical_nnz "op": op, "block_dim": block_dim_value, "out_size": int(shape[0]) if op == "non" else "UNSUP", + "padded_out_size": padded_rows if op == "non" else "UNSUP", + "pad_rows": (max(0, int(padded_rows) - int(shape[0])) if isinstance(padded_rows, int) and op == "non" else "UNSUP"), "n_rows": int(shape[0]), "n_cols": int(shape[1]), "nnzb": "SKIP", @@ -498,9 +601,18 @@ def _skip_row(matrix_name, dtype, index_dtype, op, shape, block_dim, logical_nnz "compute_ms": None, "pytorch_ms": None, "pytorch_error": None, + "pytorch_padded_ms": None, + "pytorch_padded_error": None, + "pytorch_padded_err": None, "cusparse_ms": None, "cusparse_error": None, "err": None, + "max_abs_err": None, + "max_rel_err": None, + "max_err_index": None, + "actual_at_max": None, + "expected_at_max": None, + "pytorch_err": None, "status": "SKIP", "error": error, } @@ -558,12 +670,26 @@ def _run_one_case( ) try: y_ref = _pytorch_reference(data, indices, indptr, x, shape, dtype, block_dim) - err = _allclose_error_ratio(bsr["out"], y_ref, atol, rtol) + row["padded_out_size"] = int(bsr["out"].numel()) + row["pad_rows"] = max(0, int(bsr["out"].numel()) - int(shape[0])) + y_bsr = bsr["out"][: int(shape[0])] + stats = _error_stats(y_bsr, y_ref, atol, rtol) + err = stats["ratio"] + row.update( + { + "err": err, + "max_abs_err": stats["max_abs"], + "max_rel_err": stats["max_rel"], + "max_err_index": stats["index"], + "actual_at_max": stats["actual"], + "expected_at_max": stats["expected"], + } + ) except Exception as exc: row["error"] = f"reference failed after BSR run: {exc}" return row try: - row["pytorch_ms"], row["pytorch_error"] = _time_pytorch( + row["pytorch_ms"], row["pytorch_error"], pytorch_out = _time_pytorch( data, indices, indptr, @@ -573,8 +699,31 @@ def _run_one_case( warmup, iters, ) + if pytorch_out is not None: + row["pytorch_err"] = _allclose_error_ratio(pytorch_out, y_ref, atol, rtol) except Exception as exc: row["pytorch_error"] = str(exc) + try: + ( + row["pytorch_padded_ms"], + row["pytorch_padded_error"], + pytorch_padded_out, + ) = _time_pytorch_padded( + data, + indices, + indptr, + x, + shape, + block_dim, + warmup, + iters, + ) + if pytorch_padded_out is not None: + row["pytorch_padded_err"] = _allclose_error_ratio( + pytorch_padded_out[: int(shape[0])], y_ref, atol, rtol + ) + except Exception as exc: + row["pytorch_padded_error"] = str(exc) if run_cusparse: try: row["cusparse_ms"], row["cusparse_error"] = _time_cusparse( @@ -590,7 +739,6 @@ def _run_one_case( except Exception as exc: row["cusparse_error"] = str(exc) ok = (not math.isnan(err)) and err <= 1.0 - row["err"] = err row["status"] = _status(ok) row["error"] = None if ok else "correctness check failed" return row @@ -772,6 +920,8 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "op": op, "block_dim": "ERR", "out_size": "ERR", + "padded_out_size": "ERR", + "pad_rows": "ERR", "n_rows": "ERR", "n_cols": "ERR", "nnzb": "ERR", @@ -785,6 +935,9 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "compute_ms": None, "pytorch_ms": None, "pytorch_error": None, + "pytorch_padded_ms": None, + "pytorch_padded_error": None, + "pytorch_padded_err": None, "cusparse_ms": None, "cusparse_error": None, "err": None, @@ -801,6 +954,8 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "op", "block_dim", "out_size", + "padded_out_size", + "pad_rows", "n_rows", "n_cols", "nnzb", @@ -814,6 +969,15 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "compute_ms", "pytorch_ms", "pytorch_error", + "pytorch_err", + "pytorch_padded_ms", + "pytorch_padded_error", + "pytorch_padded_err", + "max_abs_err", + "max_rel_err", + "max_err_index", + "actual_at_max", + "expected_at_max", "cusparse_ms", "cusparse_error", "err", From 13adaacc837c8eef7377525b2127579e112a6e86 Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Mon, 13 Jul 2026 17:55:52 +0800 Subject: [PATCH 11/13] spmv_bsr --- tests/test_spmv_bsr.py | 28 ++++++++++++++++++++-------- 1 file changed, 20 insertions(+), 8 deletions(-) diff --git a/tests/test_spmv_bsr.py b/tests/test_spmv_bsr.py index 952b3c0..1405cd9 100644 --- a/tests/test_spmv_bsr.py +++ b/tests/test_spmv_bsr.py @@ -48,6 +48,9 @@ def _print_baseline_notes(run_cusparse=True): print( "FlagSparse BSR follows AlphaSparse/cuSPARSE-style padded block-grid semantics; native output is padded and correctness checks slice back to logical rows." ) + print( + "Accuracy reference: Ref=spmv-coo expands the same BSR arrays to COO and runs sparse matvec; PyTorch BSR is a baseline, not the reference." + ) print( "PyTorch baseline: PT(ms) uses torch.sparse_bsr_tensor only when shape is divisible by block_dim; " "PTPad(ms) uses padded shape and slices back to logical rows for diagnostics." @@ -278,7 +281,7 @@ def _bsr_to_torch_coo(data, indices, indptr, shape, block_dim): ).coalesce() -def _pytorch_reference(data, indices, indptr, x, shape, dtype, block_dim): +def _spmv_coo_reference(data, indices, indptr, x, shape, dtype, block_dim): ref_dtype = _reference_dtype(dtype) A = _bsr_to_torch_coo( data.to(ref_dtype), @@ -469,10 +472,10 @@ def _status(ok): def _header(timing=False): split = f" {'ProcGPU':>9} {'Compute':>9}" if timing else "" return ( - f"{'Matrix':<28} {'Op':>5} {'BDim':>5} {'Out':>7} {'PadOut':>7} {'PadRows':>7} {'Rows':>7} {'Cols':>7} {'NNZB':>9} {'Pad':>7} " + f"{'Matrix':<28} {'Op':>5} {'BDim':>5} {'Ref':>8} {'Out':>7} {'PadOut':>7} {'PadRows':>7} {'Rows':>7} {'Cols':>7} {'NNZB':>9} {'Pad':>7} " f"{'BSR(ms)':>9} {'BSRGPU':>9} {'CPUProc':>9}{split} " f"{'PT(ms)':>9} {'PTPad':>9} {'CU(ms)':>9} {'BSR/PT':>8} {'BSR/CU':>8} " - f"{'Err':>10} {'Status':>6}" + f"{'BSRErr':>10} {'PTPadErr':>10} {'Status':>6}" ) @@ -490,18 +493,18 @@ def _print_row(row, timing=False): else "" ) print( - f"{name:<28} {row['op']:>5} {row['block_dim']:>5} {row['out_size']:>7} {row['padded_out_size']:>7} {row['pad_rows']:>7} {row['n_rows']:>7} {row['n_cols']:>7} {row['nnzb']:>9} {row['padding_ratio']:>7} " + f"{name:<28} {row['op']:>5} {row['block_dim']:>5} {row['reference']:>8} {row['out_size']:>7} {row['padded_out_size']:>7} {row['pad_rows']:>7} {row['n_rows']:>7} {row['n_cols']:>7} {row['nnzb']:>9} {row['padding_ratio']:>7} " f"{_fmt(row['bsr_ms']):>9} {_fmt(row['bsr_gpu_ms']):>9} {_fmt(row['process_cpu_ms']):>9}{split} " f"{_fmt(row['pytorch_ms']):>9} {_fmt(row['pytorch_padded_ms']):>9} {_fmt(row['cusparse_ms']):>9} " f"{_spd(row['pytorch_ms'], row['bsr_ms']):>8} {_spd(row['cusparse_ms'], row['bsr_ms']):>8} " - f"{_fmt_err(row['err']):>10} {row['status']:>6}" + f"{_fmt_err(row['err']):>10} {_fmt_err(row.get('pytorch_padded_err')):>10} {row['status']:>6}" ) error = row.get("error") if error: print(f" error: {str(error)[:240]}") - if row.get("pytorch_ms") is None and row.get("pytorch_error"): + if row.get("pytorch_error"): print(f" pt: {str(row['pytorch_error'])[:240]}") - if row.get("pytorch_padded_ms") is None and row.get("pytorch_padded_error"): + if row.get("pytorch_padded_error"): print(f" pt_padded: {str(row['pytorch_padded_error'])[:240]}") if row.get("status") == "FAIL": print( @@ -536,6 +539,7 @@ def _base_row( "value_dtype": _dtype_name(dtype), "index_dtype": _dtype_name(index_dtype), "op": op, + "reference": "spmv-coo", "block_dim": int(block_dim), "out_size": int(shape[0]) if op == "non" else "UNSUP", "padded_out_size": padded_rows if op == "non" else "UNSUP", @@ -558,6 +562,7 @@ def _base_row( "pytorch_padded_err": None, "cusparse_ms": None, "cusparse_error": None, + "bsr_err": None, "err": None, "max_abs_err": None, "max_rel_err": None, @@ -584,6 +589,7 @@ def _skip_row(matrix_name, dtype, index_dtype, op, shape, block_dim, logical_nnz "value_dtype": _dtype_name(dtype), "index_dtype": _dtype_name(index_dtype), "op": op, + "reference": "spmv-coo", "block_dim": block_dim_value, "out_size": int(shape[0]) if op == "non" else "UNSUP", "padded_out_size": padded_rows if op == "non" else "UNSUP", @@ -606,6 +612,7 @@ def _skip_row(matrix_name, dtype, index_dtype, op, shape, block_dim, logical_nnz "pytorch_padded_err": None, "cusparse_ms": None, "cusparse_error": None, + "bsr_err": None, "err": None, "max_abs_err": None, "max_rel_err": None, @@ -669,7 +676,7 @@ def _run_one_case( } ) try: - y_ref = _pytorch_reference(data, indices, indptr, x, shape, dtype, block_dim) + y_ref = _spmv_coo_reference(data, indices, indptr, x, shape, dtype, block_dim) row["padded_out_size"] = int(bsr["out"].numel()) row["pad_rows"] = max(0, int(bsr["out"].numel()) - int(shape[0])) y_bsr = bsr["out"][: int(shape[0])] @@ -678,6 +685,7 @@ def _run_one_case( row.update( { "err": err, + "bsr_err": err, "max_abs_err": stats["max_abs"], "max_rel_err": stats["max_rel"], "max_err_index": stats["index"], @@ -918,6 +926,7 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "value_dtype": _dtype_name(dtype), "index_dtype": _dtype_name(index_dtype), "op": op, + "reference": "spmv-coo", "block_dim": "ERR", "out_size": "ERR", "padded_out_size": "ERR", @@ -940,6 +949,7 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "pytorch_padded_err": None, "cusparse_ms": None, "cusparse_error": None, + "bsr_err": None, "err": None, "status": "ERROR", "error": str(exc), @@ -952,6 +962,7 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "value_dtype", "index_dtype", "op", + "reference", "block_dim", "out_size", "padded_out_size", @@ -980,6 +991,7 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "expected_at_max", "cusparse_ms", "cusparse_error", + "bsr_err", "err", "status", "error", From bb5bafe4432c533b14c2de6363c378c5c796e01c Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Wed, 15 Jul 2026 15:46:14 +0800 Subject: [PATCH 12/13] spmv_bsr --- README.md | 8 + ops_support.csv | 54 ++++++ ops_support.py | 2 +- run_flagsparse_pytest.py | 2 + src/flagsparse/sparse_operations/spmv_bsr.py | 163 ++++++++++++---- tests/ci/test_runtime_policies.py | 8 +- tests/pytest/Note.md | 1 + tests/pytest/test_spmv_bsr_accuracy.py | 124 +++++++----- tests/test_spmv_bsr.py | 191 ++++++++++++++----- tools/ci/run_gpu_benchmark.py | 2 + 10 files changed, 421 insertions(+), 134 deletions(-) diff --git a/README.md b/README.md index ca82314..d9baf0d 100644 --- a/README.md +++ b/README.md @@ -74,6 +74,14 @@ python tests/test_spmv_opt.py [...] python tests/test_spmv_opt.py --csv out.csv ``` +**test_spmv_bsr.py** - native BSR SpMV with padded block-grid output: + +```bash +python tests/test_spmv_bsr.py --synthetic --ops non,trans,conj +python tests/test_spmv_bsr.py --csv-bsr out.csv --block-dims 2,4 --ops non,trans,conj +# correctness uses BSR-expanded COO as the exact reference; PyTorch BSR is a baseline only. +``` + **test_spmm.py** - CSR SpMM (`.mtx` batch, synthetic, or `--csv`): ```bash diff --git a/ops_support.csv b/ops_support.csv index c766d17..6058e7d 100644 --- a/ops_support.csv +++ b/ops_support.csv @@ -1,6 +1,8 @@ operator,format,index_dtype,value_dtype,op,route,status alpha_spmm_alg1,N/A,N/A,float32,N/A,N/A,DISCOVERED_UNMAPPED alpha_spmm_alg1,N/A,N/A,float64,N/A,N/A,DISCOVERED_UNMAPPED +alpha_spmm_alg1,N/A,N/A,complex64,N/A,N/A,DISCOVERED_UNMAPPED +alpha_spmm_alg1,N/A,N/A,complex128,N/A,N/A,DISCOVERED_UNMAPPED gather,index,int32,float16,non,triton,SUPPORTED gather,index,int32,bfloat16,non,triton,SUPPORTED gather,index,int32,float32,non,triton,SUPPORTED @@ -86,9 +88,11 @@ spmm,CSR,int32,float64,non,triton_opt_alg2,SUPPORTED spmm,CSR,int32,float64,trans,triton,SUPPORTED spmm,CSR,int32,float64,conj,triton,SUPPORTED spmm,CSR,int32,complex64,non,triton,SUPPORTED +spmm,CSR,int32,complex64,non,triton_opt_alg2,SUPPORTED spmm,CSR,int32,complex64,trans,triton,SUPPORTED spmm,CSR,int32,complex64,conj,triton,SUPPORTED spmm,CSR,int32,complex128,non,triton,SUPPORTED +spmm,CSR,int32,complex128,non,triton_opt_alg2,SUPPORTED spmm,CSR,int32,complex128,trans,triton,SUPPORTED spmm,CSR,int32,complex128,conj,triton,SUPPORTED spmm,CSR,int64,float16,non,triton,SUPPORTED @@ -108,11 +112,37 @@ spmm,CSR,int64,float64,non,triton_opt_alg2,SUPPORTED spmm,CSR,int64,float64,trans,triton,SUPPORTED spmm,CSR,int64,float64,conj,triton,SUPPORTED spmm,CSR,int64,complex64,non,triton,SUPPORTED +spmm,CSR,int64,complex64,non,triton_opt_alg2,SUPPORTED spmm,CSR,int64,complex64,trans,triton,SUPPORTED spmm,CSR,int64,complex64,conj,triton,SUPPORTED spmm,CSR,int64,complex128,non,triton,SUPPORTED +spmm,CSR,int64,complex128,non,triton_opt_alg2,SUPPORTED spmm,CSR,int64,complex128,trans,triton,SUPPORTED spmm,CSR,int64,complex128,conj,triton,SUPPORTED +spmv,BSR,int32,float32,non,triton,SUPPORTED +spmv,BSR,int32,float32,trans,triton,SUPPORTED +spmv,BSR,int32,float32,conj,triton,SUPPORTED +spmv,BSR,int32,float64,non,triton,SUPPORTED +spmv,BSR,int32,float64,trans,triton,SUPPORTED +spmv,BSR,int32,float64,conj,triton,SUPPORTED +spmv,BSR,int32,complex64,non,triton,SUPPORTED +spmv,BSR,int32,complex64,trans,triton,SUPPORTED +spmv,BSR,int32,complex64,conj,triton,SUPPORTED +spmv,BSR,int32,complex128,non,triton,SUPPORTED +spmv,BSR,int32,complex128,trans,triton,SUPPORTED +spmv,BSR,int32,complex128,conj,triton,SUPPORTED +spmv,BSR,int64,float32,non,triton,SUPPORTED +spmv,BSR,int64,float32,trans,triton,SUPPORTED +spmv,BSR,int64,float32,conj,triton,SUPPORTED +spmv,BSR,int64,float64,non,triton,SUPPORTED +spmv,BSR,int64,float64,trans,triton,SUPPORTED +spmv,BSR,int64,float64,conj,triton,SUPPORTED +spmv,BSR,int64,complex64,non,triton,SUPPORTED +spmv,BSR,int64,complex64,trans,triton,SUPPORTED +spmv,BSR,int64,complex64,conj,triton,SUPPORTED +spmv,BSR,int64,complex128,non,triton,SUPPORTED +spmv,BSR,int64,complex128,trans,triton,SUPPORTED +spmv,BSR,int64,complex128,conj,triton,SUPPORTED spmv,COO,int32,float32,non,triton,SUPPORTED spmv,COO,int32,float32,trans,triton,SUPPORTED spmv,COO,int32,float32,conj,triton,SUPPORTED @@ -149,6 +179,30 @@ spmv,COO->CSR,int64,float32,non,triton,SUPPORTED spmv,COO->CSR,int64,float64,non,triton,SUPPORTED spmv,COO->CSR,int64,complex64,non,triton,SUPPORTED spmv,COO->CSR,int64,complex128,non,triton,SUPPORTED +spmv,CSC,int32,float32,non,triton,SUPPORTED +spmv,CSC,int32,float32,trans,triton,SUPPORTED +spmv,CSC,int32,float32,conj,triton,SUPPORTED +spmv,CSC,int32,float64,non,triton,SUPPORTED +spmv,CSC,int32,float64,trans,triton,SUPPORTED +spmv,CSC,int32,float64,conj,triton,SUPPORTED +spmv,CSC,int32,complex64,non,triton,SUPPORTED +spmv,CSC,int32,complex64,trans,triton,SUPPORTED +spmv,CSC,int32,complex64,conj,triton,SUPPORTED +spmv,CSC,int32,complex128,non,triton,SUPPORTED +spmv,CSC,int32,complex128,trans,triton,SUPPORTED +spmv,CSC,int32,complex128,conj,triton,SUPPORTED +spmv,CSC,int64,float32,non,triton,SUPPORTED +spmv,CSC,int64,float32,trans,triton,SUPPORTED +spmv,CSC,int64,float32,conj,triton,SUPPORTED +spmv,CSC,int64,float64,non,triton,SUPPORTED +spmv,CSC,int64,float64,trans,triton,SUPPORTED +spmv,CSC,int64,float64,conj,triton,SUPPORTED +spmv,CSC,int64,complex64,non,triton,SUPPORTED +spmv,CSC,int64,complex64,trans,triton,SUPPORTED +spmv,CSC,int64,complex64,conj,triton,SUPPORTED +spmv,CSC,int64,complex128,non,triton,SUPPORTED +spmv,CSC,int64,complex128,trans,triton,SUPPORTED +spmv,CSC,int64,complex128,conj,triton,SUPPORTED spmv,CSR,int32,float16,non,triton,SUPPORTED spmv,CSR,int32,float16,trans,triton,SUPPORTED spmv,CSR,int32,float16,conj,triton,SUPPORTED diff --git a/ops_support.py b/ops_support.py index 8f4099a..c8e21ed 100644 --- a/ops_support.py +++ b/ops_support.py @@ -275,7 +275,7 @@ def registry(modules: dict[str, SourceModule]) -> tuple[ApiSpec, ...]: ApiSpec("spmv", "flagsparse_spmv_csr", "spmv_csr", "CSR", "triton", value_const="SUPPORTED_SPMV_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmv_ops, notes="op supports non/trans/conj; conj on real dtypes is transpose-equivalent"), ApiSpec("spmv", "flagsparse_spmv_coo", "spmv_coo", "COO", "triton", value_const="SUPPORTED_SPMV_COO_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmv_coo_ops, notes="COO path stores canonical row/col tensors and supports non/trans/conj"), ApiSpec("spmv", "flagsparse_spmv_csc", "spmv_csc", "CSC", "triton", value_const="SUPPORTED_SPMV_CSC_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmv_csc_ops, notes="native CSC path supports non/trans/conj without CSR/COO conversion"), - ApiSpec("spmv", "flagsparse_spmv_bsr", "spmv_bsr", "BSR", "triton", value_const="SUPPORTED_SPMV_BSR_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmv_bsr_ops, notes="native BSR v1 supports non; trans/conj are reserved but unsupported"), + ApiSpec("spmv", "flagsparse_spmv_bsr", "spmv_bsr", "BSR", "triton", value_const="SUPPORTED_SPMV_BSR_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmv_bsr_ops, notes="native BSR path supports non/trans/conj with padded block-grid output; trans/conj directly read BSR arrays without CSR/COO/CSC conversion"), ApiSpec("spmv", "flagsparse_spmv_coo_tocsr", "spmv_csr", "COO->CSR", "triton", value_const="SUPPORTED_SPMV_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=("non",), notes="COO input is converted to CSR before compute"), ApiSpec("spmm", "flagsparse_spmm_csr", "spmm_csr", "CSR", "triton", value_const="SUPPORTED_SPMM_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=spmm_ops, notes="op supports non/trans/conj; conj on real dtypes is transpose-equivalent"), ApiSpec("spmm", "flagsparse_spmm_csr_opt", "spmm_csr", "CSR", "triton_opt", values=("float32", "float64"), index_const="SUPPORTED_INDEX_DTYPES", ops=("non",), notes="bucketed opt path only supports float32/float64"), diff --git a/run_flagsparse_pytest.py b/run_flagsparse_pytest.py index e1a222d..4f0111b 100644 --- a/run_flagsparse_pytest.py +++ b/run_flagsparse_pytest.py @@ -183,6 +183,8 @@ class OperatorTestConfig: "{input}", "--csv-bsr", "{csv}", + "--ops", + "non,trans,conj", "--warmup", "{warmup}", "--iters", diff --git a/src/flagsparse/sparse_operations/spmv_bsr.py b/src/flagsparse/sparse_operations/spmv_bsr.py index bec4a8b..0779a5d 100644 --- a/src/flagsparse/sparse_operations/spmv_bsr.py +++ b/src/flagsparse/sparse_operations/spmv_bsr.py @@ -21,7 +21,7 @@ SPMV_BSR_OP_TRANS: "trans", SPMV_BSR_OP_CONJ_TRANS: "conj", } -SPMV_BSR_SUPPORTED_OP_NAMES = ("non",) +SPMV_BSR_SUPPORTED_OP_NAMES = ("non", "trans", "conj") _SPMV_BSR_OP_NAME_TO_CODE = { name: code for code, name in SPMV_BSR_OP_NAMES.items() } @@ -60,10 +60,7 @@ def _spmv_bsr_op_transposes(op): def _ensure_spmv_bsr_supported_op(op_code): - if op_code != SPMV_BSR_OP_NON: - raise NotImplementedError( - "BSR SpMV v1 only supports op='non'; trans/conj are reserved for a future native BSR kernel" - ) + _normalize_spmv_bsr_op(op_code) def _normalize_spmv_bsr_index_fallback_policy(index_fallback_policy): @@ -238,6 +235,79 @@ def _spmv_bsr_non_complex_kernel( tl.atomic_add(y_ri_ptr + row * 2 + 1, acc_im) +@triton.jit +def _spmv_bsr_trans_real_kernel( + data_ptr, + indices_ptr, + indptr_ptr, + x_ptr, + y_ptr, + n_block_rows, + BLOCK_DIM: tl.constexpr, + BLOCK_NNZ: tl.constexpr, + SEG: tl.constexpr, +): + brow = tl.program_id(0) + inner_row = tl.program_id(1) + if brow >= n_block_rows: + return + row = brow * BLOCK_DIM + inner_row + x_val = tl.load(x_ptr + row) + start = tl.load(indptr_ptr + brow) + end = tl.load(indptr_ptr + brow + 1) + offs = start + SEG * BLOCK_NNZ + tl.arange(0, BLOCK_NNZ) + mask = offs < end + bcols = tl.load(indices_ptr + offs, mask=mask, other=0) + for inner_col in tl.static_range(0, BLOCK_DIM): + col = bcols * BLOCK_DIM + inner_col + vals = tl.load( + data_ptr + offs * BLOCK_DIM * BLOCK_DIM + inner_row * BLOCK_DIM + inner_col, + mask=mask, + other=0.0, + ) + tl.atomic_add(y_ptr + col, vals * x_val, mask=mask) + + +@triton.jit +def _spmv_bsr_trans_complex_kernel( + data_ri_ptr, + indices_ptr, + indptr_ptr, + x_ri_ptr, + y_ri_ptr, + n_block_rows, + BLOCK_DIM: tl.constexpr, + BLOCK_NNZ: tl.constexpr, + SEG: tl.constexpr, + CONJ: tl.constexpr, +): + brow = tl.program_id(0) + inner_row = tl.program_id(1) + if brow >= n_block_rows: + return + row = brow * BLOCK_DIM + inner_row + x_re = tl.load(x_ri_ptr + row * 2) + x_im = tl.load(x_ri_ptr + row * 2 + 1) + start = tl.load(indptr_ptr + brow) + end = tl.load(indptr_ptr + brow + 1) + offs = start + SEG * BLOCK_NNZ + tl.arange(0, BLOCK_NNZ) + mask = offs < end + bcols = tl.load(indices_ptr + offs, mask=mask, other=0) + for inner_col in tl.static_range(0, BLOCK_DIM): + col = bcols * BLOCK_DIM + inner_col + elem = offs * BLOCK_DIM * BLOCK_DIM + inner_row * BLOCK_DIM + inner_col + a_re = tl.load(data_ri_ptr + elem * 2, mask=mask, other=0.0) + a_im_raw = tl.load(data_ri_ptr + elem * 2 + 1, mask=mask, other=0.0) + if CONJ: + a_im = -a_im_raw + else: + a_im = a_im_raw + prod_re = a_re * x_re - a_im * x_im + prod_im = a_re * x_im + a_im * x_re + tl.atomic_add(y_ri_ptr + col * 2, prod_re, mask=mask) + tl.atomic_add(y_ri_ptr + col * 2 + 1, prod_im, mask=mask) + + def _prepare_spmv_bsr_matrix(data, indices, indptr, shape, block_dim): if not all(torch.is_tensor(t) for t in (data, indices, indptr)): raise TypeError("data, indices, indptr must all be torch.Tensor") @@ -390,7 +460,9 @@ def _validate_spmv_bsr_x(x, prepared, op_code): def _triton_spmv_bsr_kernel(prepared, x, op_code): _ensure_spmv_bsr_supported_op(op_code) dtype = prepared.data.dtype - y = torch.zeros(prepared.padded_n_rows, dtype=dtype, device=prepared.data.device) + trans = _spmv_bsr_op_transposes(op_code) + out_len = prepared.padded_n_cols if trans else prepared.padded_n_rows + y = torch.zeros(out_len, dtype=dtype, device=prepared.data.device) if prepared.nnzb == 0: return y for seg in range(prepared.max_segments): @@ -399,33 +471,60 @@ def _triton_spmv_bsr_kernel(prepared, x, op_code): data_ri = torch.view_as_real(prepared.data).reshape(-1) x_ri = torch.view_as_real(x).reshape(-1) y_ri = torch.view_as_real(y).reshape(-1) - _spmv_bsr_non_complex_kernel[grid]( - data_ri, - prepared.kernel_indices, - prepared.kernel_indptr, - x_ri, - y_ri, - prepared.padded_n_rows, - prepared.padded_n_cols, - prepared.n_block_rows, - BLOCK_DIM=prepared.block_dim, - BLOCK_NNZ=prepared.block_nnz, - SEG=seg, - ) + if trans: + _spmv_bsr_trans_complex_kernel[grid]( + data_ri, + prepared.kernel_indices, + prepared.kernel_indptr, + x_ri, + y_ri, + prepared.n_block_rows, + BLOCK_DIM=prepared.block_dim, + BLOCK_NNZ=prepared.block_nnz, + SEG=seg, + CONJ=(op_code == SPMV_BSR_OP_CONJ_TRANS), + ) + else: + _spmv_bsr_non_complex_kernel[grid]( + data_ri, + prepared.kernel_indices, + prepared.kernel_indptr, + x_ri, + y_ri, + prepared.padded_n_rows, + prepared.padded_n_cols, + prepared.n_block_rows, + BLOCK_DIM=prepared.block_dim, + BLOCK_NNZ=prepared.block_nnz, + SEG=seg, + ) else: - _spmv_bsr_non_real_kernel[grid]( - prepared.data, - prepared.kernel_indices, - prepared.kernel_indptr, - x, - y, - prepared.padded_n_rows, - prepared.padded_n_cols, - prepared.n_block_rows, - BLOCK_DIM=prepared.block_dim, - BLOCK_NNZ=prepared.block_nnz, - SEG=seg, - ) + if trans: + _spmv_bsr_trans_real_kernel[grid]( + prepared.data, + prepared.kernel_indices, + prepared.kernel_indptr, + x, + y, + prepared.n_block_rows, + BLOCK_DIM=prepared.block_dim, + BLOCK_NNZ=prepared.block_nnz, + SEG=seg, + ) + else: + _spmv_bsr_non_real_kernel[grid]( + prepared.data, + prepared.kernel_indices, + prepared.kernel_indptr, + x, + y, + prepared.padded_n_rows, + prepared.padded_n_cols, + prepared.n_block_rows, + BLOCK_DIM=prepared.block_dim, + BLOCK_NNZ=prepared.block_nnz, + SEG=seg, + ) return y diff --git a/tests/ci/test_runtime_policies.py b/tests/ci/test_runtime_policies.py index 5d3c688..115d27f 100644 --- a/tests/ci/test_runtime_policies.py +++ b/tests/ci/test_runtime_policies.py @@ -152,12 +152,14 @@ def test_spmv_bsr_op_transpose_contract(op): ) -@pytest.mark.parametrize("op", ["trans", "conj"]) -def test_spmv_bsr_unsupported_ops_rejected_by_policy(op): - with pytest.raises(NotImplementedError, match="only supports op='non'"): +@pytest.mark.parametrize("op", ["non", "trans", "conj"]) +def test_spmv_bsr_supported_ops_accepted_by_policy(op): + assert ( spmv_bsr_ops._ensure_spmv_bsr_supported_op( spmv_bsr_ops._normalize_spmv_bsr_op(op) ) + is None + ) def test_scatter_policy_validator_rejects_unknown_policy(): diff --git a/tests/pytest/Note.md b/tests/pytest/Note.md index 6a9b092..295c6c1 100644 --- a/tests/pytest/Note.md +++ b/tests/pytest/Note.md @@ -78,6 +78,7 @@ All parametrized accuracy tests build synthetic tensors on `torch.device("cuda") - Gather/Scatter use `GATHER_SCATTER_SHAPES` and `GATHER_SCATTER_FLOAT_DTYPES`; references are PyTorch indexing and `index_copy_`. - CSR/COO SpMV use synthetic sparse matrices; references use `torch.sparse.mm`. +- BSR SpMV covers `non` / `trans` / `conj`; native output uses padded block-grid length and tests compare the logical slice against a BSR-expanded COO reference. PyTorch BSR is recorded only as a same-format baseline, never as the golden reference. - CSR/COO SpMM, SpSM, SpGEMM, and SDDMM use small synthetic matrices and PyTorch dense/sparse or sampled dense references. - SpSV uses diagonally strengthened triangular matrices; references use `torch.linalg.solve_triangular`; optional CuPy/cuSPARSE references run only when available. diff --git a/tests/pytest/test_spmv_bsr_accuracy.py b/tests/pytest/test_spmv_bsr_accuracy.py index abc98c8..b2446d4 100644 --- a/tests/pytest/test_spmv_bsr_accuracy.py +++ b/tests/pytest/test_spmv_bsr_accuracy.py @@ -111,6 +111,37 @@ def _padded_cols(N, block_dim): return ((int(N) + int(block_dim) - 1) // int(block_dim)) * int(block_dim) +def _op_transposes(op): + return op in ("trans", "conj") + + +def _make_ref(dense, x, op, dtype): + ref_dtype = _reference_dtype(dtype) + A = dense.to(ref_dtype) + x_ref = x.to(ref_dtype) + if op == "trans": + return (A.T @ x_ref).to(dtype) + if op == "conj": + return (A.conj().T @ x_ref).to(dtype) + return (A @ x_ref).to(dtype) + + +def _logical_x_len(M, N, op): + return M if _op_transposes(op) else N + + +def _logical_out_len(M, N, op): + return N if _op_transposes(op) else M + + +def _padded_out_len(M, N, block_dim, op): + return _padded_cols(N, block_dim) if _op_transposes(op) else _padded_rows(M, block_dim) + + +def _padded_x_len(M, N, block_dim, op): + return _padded_rows(M, block_dim) if _op_transposes(op) else _padded_cols(N, block_dim) + + def _assert_close(actual, expected, dtype): rtol, atol = close_tolerances(dtype) ref_dtype = _reference_dtype(dtype) @@ -128,14 +159,14 @@ def _assert_close(actual, expected, dtype): "index_dtype", [torch.int32, torch.int64], ids=["int32", "int64"] ) @pytest.mark.parametrize("block_dim", [2, 4], ids=["block2", "block4"]) -def test_spmv_bsr_matches_dense_reference(M, N, name, dtype, index_dtype, block_dim): +@pytest.mark.parametrize("op", ["non", "trans", "conj"], ids=["non", "trans", "conj"]) +def test_spmv_bsr_matches_dense_reference(M, N, name, dtype, index_dtype, block_dim, op): device = torch.device("cuda") data, indices, indptr, dense = _random_bsr_mn( M, N, dtype, index_dtype, block_dim, device ) - x = _make_x(N, dtype, device) - ref_dtype = _reference_dtype(dtype) - ref = (dense.to(ref_dtype) @ x.to(ref_dtype)).to(dtype) + x = _make_x(_logical_x_len(M, N, op), dtype, device) + ref = _make_ref(dense, x, op, dtype) out = flagsparse_spmv_bsr( data, indices, @@ -143,14 +174,17 @@ def test_spmv_bsr_matches_dense_reference(M, N, name, dtype, index_dtype, block_ x, shape=(M, N), block_dim=block_dim, + op=op, index_fallback_policy="auto", ) - assert out.numel() == _padded_rows(M, block_dim) - _assert_close(out[:M], ref, dtype) + logical_out = _logical_out_len(M, N, op) + assert out.numel() == _padded_out_len(M, N, block_dim, op) + _assert_close(out[:logical_out], ref, dtype) @pytest.mark.spmv_bsr -def test_spmv_bsr_prepared_path_matches_dense_reference(): +@pytest.mark.parametrize("op", ["non", "trans", "conj"], ids=["non", "trans", "conj"]) +def test_spmv_bsr_prepared_path_matches_dense_reference(op): device = torch.device("cuda") M, N = 7, 8 dtype = torch.complex64 @@ -158,46 +192,41 @@ def test_spmv_bsr_prepared_path_matches_dense_reference(): data, indices, indptr, dense = _random_bsr_mn( M, N, dtype, torch.int32, block_dim, device ) - prepared = prepare_spmv_bsr(data, indices, indptr, (M, N), block_dim, op="non") - x = _make_x(N, dtype, device) - ref = (dense.to(torch.complex128) @ x.to(torch.complex128)).to(dtype) + prepared = prepare_spmv_bsr(data, indices, indptr, (M, N), block_dim, op=op) + x = _make_x(_logical_x_len(M, N, op), dtype, device) + ref = _make_ref(dense, x, op, dtype) out = flagsparse_spmv_bsr(x=x, prepared=prepared) - assert out.numel() == _padded_rows(M, block_dim) - _assert_close(out[:M], ref, dtype) - - -@pytest.mark.spmv_bsr -@pytest.mark.parametrize("op", ["trans", "conj"], ids=["trans", "conj"]) -def test_spmv_bsr_unsupported_ops_are_rejected(op): - device = torch.device("cuda") - data, indices, indptr, _dense = _random_bsr_mn( - 8, 12, torch.float32, torch.int32, 2, device - ) - x = torch.randn(8, dtype=torch.float32, device=device) - with pytest.raises(NotImplementedError, match="only supports op='non'"): - flagsparse_spmv_bsr(data, indices, indptr, x, shape=(8, 12), block_dim=2, op=op) + logical_out = _logical_out_len(M, N, op) + assert out.numel() == _padded_out_len(M, N, block_dim, op) + _assert_close(out[:logical_out], ref, dtype) @pytest.mark.spmv_bsr -def test_spmv_bsr_x_length_mismatch_rejected(): +@pytest.mark.parametrize("op", ["non", "trans", "conj"], ids=["non", "trans", "conj"]) +def test_spmv_bsr_x_length_mismatch_rejected(op): device = torch.device("cuda") + M, N = 8, 12 data, indices, indptr, _dense = _random_bsr_mn( - 8, 12, torch.float32, torch.int32, 2, device + M, N, torch.float32, torch.int32, 2, device ) - prepared = prepare_spmv_bsr(data, indices, indptr, (8, 12), 2) - x = torch.randn(8, dtype=torch.float32, device=device) - with pytest.raises(ValueError, match="x length must be 12"): + prepared = prepare_spmv_bsr(data, indices, indptr, (M, N), 2, op=op) + bad_len = N if _op_transposes(op) else M + x = torch.randn(bad_len, dtype=torch.float32, device=device) + expected = _logical_x_len(M, N, op) + with pytest.raises(ValueError, match=f"x length must be {expected}"): flagsparse_spmv_bsr(x=x, prepared=prepared) @pytest.mark.spmv_bsr -def test_spmv_bsr_int64_auto_fallback_to_int32(monkeypatch): +@pytest.mark.parametrize("op", ["non", "trans"], ids=["non", "trans"]) +def test_spmv_bsr_int64_auto_fallback_to_int32(monkeypatch, op): device = torch.device("cuda") + M, N = 12, 9 data, indices, indptr, dense = _random_bsr_mn( - 12, 9, torch.float32, torch.int64, 2, device + M, N, torch.float32, torch.int64, 2, device ) - x = torch.randn(9, dtype=torch.float32, device=device) - ref = dense.to(torch.float64) @ x.to(torch.float64) + x = torch.randn(_logical_x_len(M, N, op), dtype=torch.float32, device=device) + ref = _make_ref(dense, x, op, torch.float32) state = {"forced_once": False} original = spmv_bsr_mod._triton_spmv_bsr_kernel @@ -213,13 +242,15 @@ def fail_int64_once(prepared, x_in, op_code): indices, indptr, x, - shape=(12, 9), + shape=(M, N), block_dim=2, + op=op, index_fallback_policy="auto", ) assert state["forced_once"] - assert out.numel() == _padded_rows(12, 2) - _assert_close(out[:12], ref.to(torch.float32), torch.float32) + logical_out = _logical_out_len(M, N, op) + assert out.numel() == _padded_out_len(M, N, 2, op) + _assert_close(out[:logical_out], ref.to(torch.float32), torch.float32) @pytest.mark.spmv_bsr @@ -265,21 +296,24 @@ def test_spmv_bsr_non_divisible_shape_matches_dense_reference(): @pytest.mark.spmv_bsr -def test_spmv_bsr_accepts_logical_or_padded_x(): +@pytest.mark.parametrize("op", ["non", "trans", "conj"], ids=["non", "trans", "conj"]) +def test_spmv_bsr_accepts_logical_or_padded_x(op): device = torch.device("cuda") M, N = 7, 9 block_dim = 4 data, indices, indptr, dense = _random_bsr_mn( M, N, torch.float32, torch.int32, block_dim, device ) - x = torch.randn(N, dtype=torch.float32, device=device) - x_padded = torch.zeros(_padded_cols(N, block_dim), dtype=torch.float32, device=device) - x_padded[:N].copy_(x) - ref = dense.to(torch.float64) @ x.to(torch.float64) - prepared = prepare_spmv_bsr(data, indices, indptr, (M, N), block_dim) + logical_x = _logical_x_len(M, N, op) + x = torch.randn(logical_x, dtype=torch.float32, device=device) + x_padded = torch.zeros(_padded_x_len(M, N, block_dim, op), dtype=torch.float32, device=device) + x_padded[:logical_x].copy_(x) + ref = _make_ref(dense, x, op, torch.float32) + prepared = prepare_spmv_bsr(data, indices, indptr, (M, N), block_dim, op=op) out_logical_x = flagsparse_spmv_bsr(x=x, prepared=prepared) out_padded_x = flagsparse_spmv_bsr(x=x_padded, prepared=prepared) - assert out_logical_x.numel() == _padded_rows(M, block_dim) - assert out_padded_x.numel() == _padded_rows(M, block_dim) + logical_out = _logical_out_len(M, N, op) + assert out_logical_x.numel() == _padded_out_len(M, N, block_dim, op) + assert out_padded_x.numel() == _padded_out_len(M, N, block_dim, op) _assert_close(out_logical_x, out_padded_x, torch.float32) - _assert_close(out_logical_x[:M], ref.to(torch.float32), torch.float32) + _assert_close(out_logical_x[:logical_out], ref.to(torch.float32), torch.float32) diff --git a/tests/test_spmv_bsr.py b/tests/test_spmv_bsr.py index 1405cd9..61a6e85 100644 --- a/tests/test_spmv_bsr.py +++ b/tests/test_spmv_bsr.py @@ -29,7 +29,7 @@ VALUE_DTYPES = (torch.float32, torch.float64, torch.complex64, torch.complex128) INDEX_DTYPES = (torch.int32, torch.int64) OPS = ("non", "trans", "conj") -SUPPORTED_OPS = ("non",) +SUPPORTED_OPS = OPS TEST_SIZES = ((64, 96), (160, 1024), (128, 256)) DEFAULT_BLOCK_DIMS = (4,) WARMUP = 10 @@ -182,6 +182,28 @@ def _padded_shape(shape, block_dim): ) +def _op_transposes(op): + return str(op).lower() in ("trans", "conj") + + +def _logical_x_size(shape, op): + return int(shape[0]) if _op_transposes(op) else int(shape[1]) + + +def _padded_x_size(shape, block_dim, op): + padded_rows, padded_cols = _padded_shape(shape, block_dim) + return padded_rows if _op_transposes(op) else padded_cols + + +def _logical_out_size(shape, op): + return int(shape[1]) if _op_transposes(op) else int(shape[0]) + + +def _padded_out_size(shape, block_dim, op): + padded_rows, padded_cols = _padded_shape(shape, block_dim) + return padded_cols if _op_transposes(op) else padded_rows + + def _pad_vector(x, length): length = int(length) if x.numel() == length: @@ -281,7 +303,7 @@ def _bsr_to_torch_coo(data, indices, indptr, shape, block_dim): ).coalesce() -def _spmv_coo_reference(data, indices, indptr, x, shape, dtype, block_dim): +def _spmv_coo_reference(data, indices, indptr, x, shape, dtype, block_dim, op): ref_dtype = _reference_dtype(dtype) A = _bsr_to_torch_coo( data.to(ref_dtype), @@ -290,6 +312,10 @@ def _spmv_coo_reference(data, indices, indptr, x, shape, dtype, block_dim): shape, block_dim, ) + if op == "trans": + A = A.transpose(0, 1) + elif op == "conj": + A = A.conj().transpose(0, 1) return torch.sparse.mm(A, x.to(ref_dtype).unsqueeze(1)).squeeze(1).to(dtype) @@ -342,10 +368,9 @@ def _cuda_event_benchmark(op, warmup, iters): return out, start.elapsed_time(end) / count -def _time_flagsparse_bsr(data, indices, indptr, x, shape, block_dim, warmup, iters, timing=False): - prepared = fs.prepare_spmv_bsr(data, indices, indptr, shape, block_dim, op="non") - _padded_rows, padded_cols = _padded_shape(shape, block_dim) - x_for_bsr = _pad_vector(x, padded_cols) +def _time_flagsparse_bsr(data, indices, indptr, x, shape, block_dim, op, warmup, iters, timing=False): + prepared = fs.prepare_spmv_bsr(data, indices, indptr, shape, block_dim, op=op) + x_for_bsr = _pad_vector(x, _padded_x_size(shape, block_dim, op)) out, gpu_ms = _cuda_event_benchmark( lambda: fs.flagsparse_spmv_bsr(x=x_for_bsr, prepared=prepared), warmup, @@ -361,7 +386,15 @@ def _time_flagsparse_bsr(data, indices, indptr, x, shape, block_dim, warmup, ite } -def _time_pytorch(data, indices, indptr, x, shape, block_dim, warmup, iters): +def _apply_pytorch_op(A, x, op): + if op == "trans": + A = A.transpose(0, 1) + elif op == "conj": + A = A.conj().transpose(0, 1) + return torch.sparse.mm(A, x.unsqueeze(1)).squeeze(1) + + +def _time_pytorch(data, indices, indptr, x, shape, block_dim, op, warmup, iters): if int(shape[0]) % int(block_dim) != 0 or int(shape[1]) % int(block_dim) != 0: return ( None, @@ -382,14 +415,14 @@ def _time_pytorch(data, indices, indptr, x, shape, block_dim, warmup, iters): device=data.device, dtype=data.dtype, ) - fn = lambda: torch.sparse.mm(A, x.unsqueeze(1)).squeeze(1) + fn = lambda: _apply_pytorch_op(A, x, op) out, ms = _cuda_event_benchmark(fn, warmup, iters) return ms, None, out -def _time_pytorch_padded(data, indices, indptr, x, shape, block_dim, warmup, iters): +def _time_pytorch_padded(data, indices, indptr, x, shape, block_dim, op, warmup, iters): padded_shape = _padded_shape(shape, block_dim) - padded_x = _pad_vector(x, padded_shape[1]) + padded_x = _pad_vector(x, padded_shape[0] if _op_transposes(op) else padded_shape[1]) with warnings.catch_warnings(): warnings.filterwarnings( "ignore", @@ -404,12 +437,12 @@ def _time_pytorch_padded(data, indices, indptr, x, shape, block_dim, warmup, ite device=data.device, dtype=data.dtype, ) - fn = lambda: torch.sparse.mm(A, padded_x.unsqueeze(1)).squeeze(1) + fn = lambda: _apply_pytorch_op(A, padded_x, op) out, ms = _cuda_event_benchmark(fn, warmup, iters) return ms, None, out -def _time_cusparse(data, indices, indptr, x, shape, block_dim, warmup, iters): +def _time_cusparse(data, indices, indptr, x, shape, block_dim, op, warmup, iters): if cp is None or cpx_sparse is None: return None, "CuPy/cupyx.scipy.sparse is not available" if not hasattr(cpx_sparse, "bsr_matrix"): @@ -423,7 +456,12 @@ def run_with_index_dtype(index_dtype, fallback_note=None): ptr_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(indptr.to(index_dtype))) x_cp = cp.from_dlpack(torch.utils.dlpack.to_dlpack(x)) A = cpx_sparse.bsr_matrix((data_cp, ind_cp, ptr_cp), shape=shape) - fn = lambda: A @ x_cp + if op == "trans": + fn = lambda: A.T @ x_cp + elif op == "conj": + fn = lambda: A.conj().T @ x_cp + else: + fn = lambda: A @ x_cp for _ in range(max(0, int(warmup))): _ = fn() cp.cuda.runtime.deviceSynchronize() @@ -474,8 +512,8 @@ def _header(timing=False): return ( f"{'Matrix':<28} {'Op':>5} {'BDim':>5} {'Ref':>8} {'Out':>7} {'PadOut':>7} {'PadRows':>7} {'Rows':>7} {'Cols':>7} {'NNZB':>9} {'Pad':>7} " f"{'BSR(ms)':>9} {'BSRGPU':>9} {'CPUProc':>9}{split} " - f"{'PT(ms)':>9} {'PTPad':>9} {'CU(ms)':>9} {'BSR/PT':>8} {'BSR/CU':>8} " - f"{'BSRErr':>10} {'PTPadErr':>10} {'Status':>6}" + f"{'PT(ms)':>9} {'PTPad':>9} {'PTPMode':>9} {'CU(ms)':>9} {'BSR/PT':>8} {'BSR/CU':>8} " + f"{'BSRErr':>10} {'PTPadErr':>10} {'B/PT':>10} {'B/PTPad':>10} {'Status':>6}" ) @@ -495,9 +533,10 @@ def _print_row(row, timing=False): print( f"{name:<28} {row['op']:>5} {row['block_dim']:>5} {row['reference']:>8} {row['out_size']:>7} {row['padded_out_size']:>7} {row['pad_rows']:>7} {row['n_rows']:>7} {row['n_cols']:>7} {row['nnzb']:>9} {row['padding_ratio']:>7} " f"{_fmt(row['bsr_ms']):>9} {_fmt(row['bsr_gpu_ms']):>9} {_fmt(row['process_cpu_ms']):>9}{split} " - f"{_fmt(row['pytorch_ms']):>9} {_fmt(row['pytorch_padded_ms']):>9} {_fmt(row['cusparse_ms']):>9} " + f"{_fmt(row['pytorch_ms']):>9} {_fmt(row['pytorch_padded_ms']):>9} {str(row.get('pytorch_padded_mode') or 'N/A')[:9]:>9} {_fmt(row['cusparse_ms']):>9} " f"{_spd(row['pytorch_ms'], row['bsr_ms']):>8} {_spd(row['cusparse_ms'], row['bsr_ms']):>8} " - f"{_fmt_err(row['err']):>10} {_fmt_err(row.get('pytorch_padded_err')):>10} {row['status']:>6}" + f"{_fmt_err(row['err']):>10} {_fmt_err(row.get('pytorch_padded_err')):>10} " + f"{_fmt_err(row.get('bsr_vs_pytorch_err')):>10} {_fmt_err(row.get('bsr_vs_pytorch_padded_err')):>10} {row['status']:>6}" ) error = row.get("error") if error: @@ -515,7 +554,9 @@ def _print_row(row, timing=False): f"actual={row.get('actual_at_max')}, " f"expected={row.get('expected_at_max')}, " f"pytorch_err={_fmt_err(row.get('pytorch_err'))}, " - f"pytorch_padded_err={_fmt_err(row.get('pytorch_padded_err'))}" + f"pytorch_padded_err={_fmt_err(row.get('pytorch_padded_err'))}, " + f"bsr_vs_pytorch_err={_fmt_err(row.get('bsr_vs_pytorch_err'))}, " + f"bsr_vs_pytorch_padded_err={_fmt_err(row.get('bsr_vs_pytorch_padded_err'))}" ) @@ -533,7 +574,8 @@ def _base_row( nnzb = int(data.shape[0]) stored_nnz = int(data.numel()) logical_nnz = max(1, int(logical_nnz if logical_nnz is not None else stored_nnz)) - padded_rows, _padded_cols = _padded_shape(shape, block_dim) + logical_out = _logical_out_size(shape, op) if op in SUPPORTED_OPS else "UNSUP" + padded_out = _padded_out_size(shape, block_dim, op) if op in SUPPORTED_OPS else "UNSUP" return { "matrix": matrix_name, "value_dtype": _dtype_name(dtype), @@ -541,9 +583,9 @@ def _base_row( "op": op, "reference": "spmv-coo", "block_dim": int(block_dim), - "out_size": int(shape[0]) if op == "non" else "UNSUP", - "padded_out_size": padded_rows if op == "non" else "UNSUP", - "pad_rows": max(0, padded_rows - int(shape[0])) if op == "non" else "UNSUP", + "out_size": logical_out, + "padded_out_size": padded_out, + "pad_rows": (max(0, int(padded_out) - int(logical_out)) if op in SUPPORTED_OPS else "UNSUP"), "n_rows": int(shape[0]), "n_cols": int(shape[1]), "nnzb": nnzb, @@ -560,6 +602,9 @@ def _base_row( "pytorch_padded_ms": None, "pytorch_padded_error": None, "pytorch_padded_err": None, + "pytorch_padded_mode": None, + "bsr_vs_pytorch_err": None, + "bsr_vs_pytorch_padded_err": None, "cusparse_ms": None, "cusparse_error": None, "bsr_err": None, @@ -581,9 +626,13 @@ def _skip_row(matrix_name, dtype, index_dtype, op, shape, block_dim, logical_nnz except (TypeError, ValueError): block_dim_value = block_dim try: - padded_rows, _padded_cols = _padded_shape(shape, block_dim_value) + logical_out = _logical_out_size(shape, op) if op in SUPPORTED_OPS else "UNSUP" + padded_out = _padded_out_size(shape, block_dim_value, op) if op in SUPPORTED_OPS else "UNSUP" + pad_rows = max(0, int(padded_out) - int(logical_out)) if op in SUPPORTED_OPS else "UNSUP" except Exception: - padded_rows = "SKIP" + logical_out = "SKIP" + padded_out = "SKIP" + pad_rows = "SKIP" return { "matrix": matrix_name, "value_dtype": _dtype_name(dtype), @@ -591,9 +640,9 @@ def _skip_row(matrix_name, dtype, index_dtype, op, shape, block_dim, logical_nnz "op": op, "reference": "spmv-coo", "block_dim": block_dim_value, - "out_size": int(shape[0]) if op == "non" else "UNSUP", - "padded_out_size": padded_rows if op == "non" else "UNSUP", - "pad_rows": (max(0, int(padded_rows) - int(shape[0])) if isinstance(padded_rows, int) and op == "non" else "UNSUP"), + "out_size": logical_out, + "padded_out_size": padded_out, + "pad_rows": pad_rows, "n_rows": int(shape[0]), "n_cols": int(shape[1]), "nnzb": "SKIP", @@ -610,6 +659,9 @@ def _skip_row(matrix_name, dtype, index_dtype, op, shape, block_dim, logical_nnz "pytorch_padded_ms": None, "pytorch_padded_error": None, "pytorch_padded_err": None, + "pytorch_padded_mode": None, + "bsr_vs_pytorch_err": None, + "bsr_vs_pytorch_padded_err": None, "cusparse_ms": None, "cusparse_error": None, "bsr_err": None, @@ -657,12 +709,12 @@ def _run_one_case( row["process_gpu_ms"] = 0.0 if timing else None if op not in SUPPORTED_OPS: row["status"] = "SKIP" - row["error"] = "BSR SpMV v1 only supports op=non" + row["error"] = "unsupported BSR SpMV op" return row - x = _random_values((int(shape[1]),), dtype, data.device) + x = _random_values((_logical_x_size(shape, op),), dtype, data.device) atol, rtol = _reference_tolerance(dtype) try: - bsr = _time_flagsparse_bsr(data, indices, indptr, x, shape, block_dim, warmup, iters, timing=timing) + bsr = _time_flagsparse_bsr(data, indices, indptr, x, shape, block_dim, op, warmup, iters, timing=timing) except Exception as exc: row["error"] = f"flagsparse_spmv_bsr failed: {exc}" return row @@ -676,10 +728,12 @@ def _run_one_case( } ) try: - y_ref = _spmv_coo_reference(data, indices, indptr, x, shape, dtype, block_dim) + y_ref = _spmv_coo_reference(data, indices, indptr, x, shape, dtype, block_dim, op) row["padded_out_size"] = int(bsr["out"].numel()) - row["pad_rows"] = max(0, int(bsr["out"].numel()) - int(shape[0])) - y_bsr = bsr["out"][: int(shape[0])] + logical_out = _logical_out_size(shape, op) + row["out_size"] = logical_out + row["pad_rows"] = max(0, int(bsr["out"].numel()) - int(logical_out)) + y_bsr = bsr["out"][: logical_out] stats = _error_stats(y_bsr, y_ref, atol, rtol) err = stats["ratio"] row.update( @@ -696,6 +750,7 @@ def _run_one_case( except Exception as exc: row["error"] = f"reference failed after BSR run: {exc}" return row + pytorch_out = None try: row["pytorch_ms"], row["pytorch_error"], pytorch_out = _time_pytorch( data, @@ -704,34 +759,57 @@ def _run_one_case( x, shape, block_dim, + op, warmup, iters, ) if pytorch_out is not None: row["pytorch_err"] = _allclose_error_ratio(pytorch_out, y_ref, atol, rtol) + row["bsr_vs_pytorch_err"] = _allclose_error_ratio( + y_bsr, pytorch_out, atol, rtol + ) except Exception as exc: row["pytorch_error"] = str(exc) - try: - ( - row["pytorch_padded_ms"], - row["pytorch_padded_error"], - pytorch_padded_out, - ) = _time_pytorch_padded( - data, - indices, - indptr, - x, - shape, - block_dim, - warmup, - iters, - ) - if pytorch_padded_out is not None: - row["pytorch_padded_err"] = _allclose_error_ratio( - pytorch_padded_out[: int(shape[0])], y_ref, atol, rtol + padded_shape = _padded_shape(shape, block_dim) + if ( + pytorch_out is not None + and int(padded_shape[0]) == int(shape[0]) + and int(padded_shape[1]) == int(shape[1]) + ): + row["pytorch_padded_ms"] = row["pytorch_ms"] + row["pytorch_padded_error"] = None + row["pytorch_padded_err"] = row["pytorch_err"] + row["bsr_vs_pytorch_padded_err"] = row["bsr_vs_pytorch_err"] + row["pytorch_padded_mode"] = "same_as_pt" + else: + try: + ( + row["pytorch_padded_ms"], + row["pytorch_padded_error"], + pytorch_padded_out, + ) = _time_pytorch_padded( + data, + indices, + indptr, + x, + shape, + block_dim, + op, + warmup, + iters, ) - except Exception as exc: - row["pytorch_padded_error"] = str(exc) + row["pytorch_padded_mode"] = "padded_shape" + if pytorch_padded_out is not None: + pytorch_padded_logical = pytorch_padded_out[: logical_out] + row["pytorch_padded_err"] = _allclose_error_ratio( + pytorch_padded_logical, y_ref, atol, rtol + ) + row["bsr_vs_pytorch_padded_err"] = _allclose_error_ratio( + y_bsr, pytorch_padded_logical, atol, rtol + ) + except Exception as exc: + row["pytorch_padded_error"] = str(exc) + row["pytorch_padded_mode"] = "padded_shape" if run_cusparse: try: row["cusparse_ms"], row["cusparse_error"] = _time_cusparse( @@ -741,6 +819,7 @@ def _run_one_case( x, shape, block_dim, + op, warmup, iters, ) @@ -947,6 +1026,9 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "pytorch_padded_ms": None, "pytorch_padded_error": None, "pytorch_padded_err": None, + "pytorch_padded_mode": None, + "bsr_vs_pytorch_err": None, + "bsr_vs_pytorch_padded_err": None, "cusparse_ms": None, "cusparse_error": None, "bsr_err": None, @@ -984,6 +1066,9 @@ def run_csv(mtx_paths, csv_path, value_dtypes=None, index_dtypes=None, block_dim "pytorch_padded_ms", "pytorch_padded_error", "pytorch_padded_err", + "pytorch_padded_mode", + "bsr_vs_pytorch_err", + "bsr_vs_pytorch_padded_err", "max_abs_err", "max_rel_err", "max_err_index", diff --git a/tools/ci/run_gpu_benchmark.py b/tools/ci/run_gpu_benchmark.py index 7bf4277..b908dad 100644 --- a/tools/ci/run_gpu_benchmark.py +++ b/tools/ci/run_gpu_benchmark.py @@ -107,6 +107,8 @@ def _command_specs( "spmv-bsr": [ "tests/test_spmv_bsr.py", "--synthetic", + "--ops", + "non,trans,conj", "--warmup", str(args.warmup), "--iters", From 2f6f005c82d23f93b957dd206e8c17c4f7829091 Mon Sep 17 00:00:00 2001 From: zyq1105331849 <1105331849@qq.com> Date: Wed, 15 Jul 2026 19:42:03 +0800 Subject: [PATCH 13/13] merge --- README.md | 7 + README_cn.md | 6 + conf/operators.yaml | 17 + ops_support.csv | 4 + ops_support.py | 1 + pytest.ini | 1 + run_flagsparse_pytest.py | 11 + src/flagsparse/__init__.py | 2 + src/flagsparse/sparse_formats.py | 79 ++- src/flagsparse/sparse_operations/__init__.py | 2 + src/flagsparse/sparse_operations/spsv.py | 185 +++++- tests/ci/test_cli_help.py | 1 + tests/ci/test_gather_contract.py | 74 +++ tests/ci/test_operator_registry.py | 1 + tests/ci/test_public_api.py | 1 + tests/test_spsv_sell.py | 651 +++++++++++++++++++ 16 files changed, 1028 insertions(+), 15 deletions(-) create mode 100644 tests/ci/test_gather_contract.py create mode 100644 tests/test_spsv_sell.py diff --git a/README.md b/README.md index d9baf0d..80ddd53 100644 --- a/README.md +++ b/README.md @@ -127,10 +127,17 @@ python tests/test_spgemm.py --csv results.csv # optional: --dtype flo **test_spsv.py** - SpSV (triangular solve; **square** matrices only). CSR and COO share this script; there is **no** `test_spsv_coo.py`. +**test_spsv_sell.py** - lower, real, native column-major SELL SpSV. Its CSV and +terminal fields follow the CSR SpSV output. `FlagSparse_ms` and `cuSPARSE_ms` +both cover every per-call preparation/analysis plus solve; static descriptors +and SELL conversion are outside the timed interval. + ```bash python tests/test_spsv.py --synthetic python tests/test_spsv.py --csv-csr spsv.csv python tests/test_spsv.py --csv-coo out.csv # same CSV columns as CSR +pytest -q -s tests/test_spsv_sell.py +python tests/test_spsv_sell.py --csv sell.csv --slice-size 32 ``` **test_spsm.py** - SpSM (triangular matrix-matrix solve; **square** matrices only): diff --git a/README_cn.md b/README_cn.md index a413a63..051d5fb 100644 --- a/README_cn.md +++ b/README_cn.md @@ -117,10 +117,16 @@ python tests/test_spgemm.py <目录/> --csv results.csv # 可选:--dtype fl **test_spsv.py** - SpSV(三角求解;**仅方阵**)。CSR 与 COO 共用本脚本;**不存在** `test_spsv_coo.py`。 +**test_spsv_sell.py** - 下三角、实数、原生列主序 SELL SpSV。CSV 和终端字段 +遵循 CSR SpSV 输出;`FlagSparse_ms` 和 `cuSPARSE_ms` 都覆盖每次调用的准备/ +分析加求解,静态 descriptor 与 SELL 转换不计时。 + ```bash python tests/test_spsv.py --synthetic python tests/test_spsv.py <目录/> --csv-csr spsv.csv python tests/test_spsv.py <目录/> --csv-coo out.csv # 列与 CSR 相同 +pytest -q -s tests/test_spsv_sell.py +python tests/test_spsv_sell.py <目录或文件.mtx> --csv sell.csv --slice-size 32 ``` **test_spsm.py** - SpSM(三角矩阵-稠密矩阵求解;**仅方阵**): diff --git a/conf/operators.yaml b/conf/operators.yaml index b366035..eb1503e 100644 --- a/conf/operators.yaml +++ b/conf/operators.yaml @@ -295,6 +295,23 @@ ops: stages: - beta: "1.0" + - id: spsv_sell + description: | + Solves sparse triangular systems represented in cuSPARSE-compatible + column-major Sliced ELLPACK format using a native Triton kernel. + for: + - flagsparse_spsv_sell + labels: + - flagsparse + - sparse + - sell + - triton + - public-api + kind: + - SparseSolver + stages: + - beta: "1.0" + - id: spsv_descriptor_api description: | Provides descriptor, buffer-size, analysis, preprocess, and solve APIs for SpSV workflows. diff --git a/ops_support.csv b/ops_support.csv index 6058e7d..da51c7e 100644 --- a/ops_support.csv +++ b/ops_support.csv @@ -279,3 +279,7 @@ spsv,CSR,int64,complex64,non,triton,SUPPORTED spsv,CSR,int64,complex64,trans,triton,SUPPORTED spsv,CSR,int64,complex128,non,triton,SUPPORTED spsv,CSR,int64,complex128,trans,triton,SUPPORTED +spsv,SELL,int32,float32,non,triton,SUPPORTED +spsv,SELL,int32,float64,non,triton,SUPPORTED +spsv,SELL,int64,float32,non,triton,SUPPORTED +spsv,SELL,int64,float64,non,triton,SUPPORTED diff --git a/ops_support.py b/ops_support.py index c8e21ed..3849f07 100644 --- a/ops_support.py +++ b/ops_support.py @@ -285,6 +285,7 @@ def registry(modules: dict[str, SourceModule]) -> tuple[ApiSpec, ...]: ApiSpec("sddmm", "flagsparse_sddmm_csr", "sddmm_csr", "CSR", "triton", value_const="SUPPORTED_SDDMM_VALUE_DTYPES", index_const="SUPPORTED_INDEX_DTYPES", ops=("non",)), ApiSpec("spsv", "flagsparse_spsv_csr", "spsv", "CSR", "triton", value_const="SUPPORTED_SPSV_VALUE_DTYPES", index_const="SUPPORTED_SPSV_INDEX_DTYPES", ops=("NON_TRANS", "TRANS"), notes="TRANS support is narrower than NON_TRANS; see combo constants"), ApiSpec("spsv", "flagsparse_spsv_coo", "spsv", "COO", "triton", value_const="SUPPORTED_SPSV_VALUE_DTYPES", index_const="SUPPORTED_SPSV_INDEX_DTYPES", ops=("NON_TRANS", "TRANS"), notes="TRANS support is narrower than NON_TRANS; see combo constants"), + ApiSpec("spsv", "flagsparse_spsv_sell", "spsv", "SELL", "triton", values=("float32", "float64"), indices=("int32", "int64"), ops=("NON_TRANS",), notes="lower triangular; cuSPARSE-compatible column-major SELL storage; native Triton solve"), ApiSpec("spsm", "flagsparse_spsm_csr", "spsm", "CSR", "triton", value_const="SUPPORTED_SPSM_VALUE_DTYPES", index_const="SUPPORTED_SPSM_INDEX_DTYPES", values=("float32", "float64", "complex64", "complex128"), indices=("int32",), ops=("NON_TRANS",), notes="opA/opB must both be NON_TRANS; row-major dense layout only"), ApiSpec("spsm", "flagsparse_spsm_coo", "spsm", "COO", "triton", value_const="SUPPORTED_SPSM_VALUE_DTYPES", index_const="SUPPORTED_SPSM_INDEX_DTYPES", values=("float32", "float64", "complex64", "complex128"), indices=("int32",), ops=("NON_TRANS",), notes="opA/opB must both be NON_TRANS; row-major dense layout only"), ) diff --git a/pytest.ini b/pytest.ini index bb035fc..3cd9029 100644 --- a/pytest.ini +++ b/pytest.ini @@ -22,6 +22,7 @@ markers = spsv: CSR/COO SpSV accuracy (tests/pytest) spsv_csr: CSR SpSV accuracy (tests/pytest) spsv_coo: COO SpSV accuracy (tests/pytest) + spsv_sell: native lower-real Triton SELL SpSV accuracy and cuSPARSE comparison spsm: CSR/COO SpSM accuracy (tests/pytest) spsm_csr: CSR SpSM accuracy (tests/pytest) spsm_coo: COO SpSM accuracy (tests/pytest) diff --git a/run_flagsparse_pytest.py b/run_flagsparse_pytest.py index 4f0111b..5af76c9 100644 --- a/run_flagsparse_pytest.py +++ b/run_flagsparse_pytest.py @@ -304,6 +304,16 @@ class OperatorTestConfig: "--iters", "{iters}", ), + "spsv_sell": ( + "tests/test_spsv_sell.py", + "{input}", + "--csv", + "{csv}", + "--warmup", + "{warmup}", + "--iters", + "{iters}", + ), "spsm_csr": ( "tests/test_spsm.py", "{input}", @@ -355,6 +365,7 @@ class OperatorTestConfig: "sddmm_csr": OperatorTestConfig("sddmm_csr", PERFORMANCE_COMMANDS["sddmm_csr"]), "spsv_csr": OperatorTestConfig("spsv_csr", PERFORMANCE_COMMANDS["spsv_csr"]), "spsv_coo": OperatorTestConfig("spsv_coo", PERFORMANCE_COMMANDS["spsv_coo"]), + "spsv_sell": OperatorTestConfig("spsv_sell", PERFORMANCE_COMMANDS["spsv_sell"]), "spsm_csr": OperatorTestConfig("spsm_csr", PERFORMANCE_COMMANDS["spsm_csr"]), "spsm_coo": OperatorTestConfig("spsm_coo", PERFORMANCE_COMMANDS["spsm_coo"]), } diff --git a/src/flagsparse/__init__.py b/src/flagsparse/__init__.py index 9714dcd..5f4b879 100644 --- a/src/flagsparse/__init__.py +++ b/src/flagsparse/__init__.py @@ -67,6 +67,7 @@ "flagsparse_spmm_csr_opt_alg2_preprocess", "flagsparse_spsv_csr", "flagsparse_spsv_coo", + "flagsparse_spsv_sell", "flagsparse_spsv_buffer_size", "flagsparse_spsv_buffer_size_ex", "flagsparse_spsv_analysis_csr", @@ -201,6 +202,7 @@ "flagsparse_spmm_csr_opt_alg2_preprocess", "flagsparse_spsv_csr", "flagsparse_spsv_coo", + "flagsparse_spsv_sell", "flagsparse_spsv_buffer_size", "flagsparse_spsv_buffer_size_ex", "flagsparse_spsv_analysis_csr", diff --git a/src/flagsparse/sparse_formats.py b/src/flagsparse/sparse_formats.py index 5646f86..7c375c3 100644 --- a/src/flagsparse/sparse_formats.py +++ b/src/flagsparse/sparse_formats.py @@ -204,15 +204,39 @@ def __repr__(self): class SELLMatrix: """ - Sliced ELLPACK format. Stores: values, indices (column), slice_ptr, rows_per_slice. - CuPy-compatible interface; backend uses CuPy arrays. + cuSPARSE-compatible Sliced ELLPACK format. + + Entries inside every slice use column-major SELL order:: + + offset + slot * slice_size + row_in_slice + + The column index ``-1`` denotes padding. ``rows_per_slice`` is retained + for the high-level conversion helpers; the physical stride is always + ``slice_size``, including the final partial slice. """ - def __init__(self, values, indices, slice_ptr, rows_per_slice, shape, dtype=None): + def __init__( + self, + values, + indices, + slice_ptr, + rows_per_slice, + shape, + dtype=None, + slice_size=None, + ): self._values = _to_cupy_array(values, dtype=_resolve_dtype(dtype)) self._indices = _to_cupy_array(indices, dtype=cp.int64) self._slice_ptr = _to_cupy_array(slice_ptr, dtype=cp.int64) self._rows_per_slice = _to_cupy_array(rows_per_slice, dtype=cp.int64) self._shape = tuple(shape) + if slice_size is None: + if self._rows_per_slice.size: + slice_size = int(cp.max(self._rows_per_slice).item()) + else: + slice_size = 1 + self._slice_size = int(slice_size) + if self._slice_size <= 0: + raise ValueError("slice_size must be a positive integer") @property def values(self): @@ -230,6 +254,10 @@ def slice_ptr(self): def rows_per_slice(self): return self._rows_per_slice + @property + def slice_size(self): + return self._slice_size + @property def shape(self): return self._shape @@ -365,11 +393,12 @@ def _sell_to_coo(sell_mat): if rps <= 0: base_row += rps continue - max_nnz = (end - start) // rps + slice_size = int(sell_mat.slice_size) + max_nnz = (end - start) // slice_size for r in range(rps): row = base_row + r for k in range(max_nnz): - idx = start + r * max_nnz + k + idx = start + k * slice_size + r col = int(indices[idx]) val = values[idx] nonzero = (val != 0).item() if hasattr(val, "item") else (val != 0) @@ -461,10 +490,10 @@ def _coo_to_sell_impl(rows, cols, data, shape, slice_size): max_nnz = int(cp.max(nnz_per_row[r0:r1])) else: max_nnz = 0 - total_entries += rps * max_nnz + total_entries += slice_size * max_nnz slice_ptr[s + 1] = total_entries values = cp.zeros(total_entries, dtype=data.dtype) - indices = cp.zeros(total_entries, dtype=cp.int64) + indices = cp.full(total_entries, -1, dtype=cp.int64) row_start = cp.zeros(n_rows + 1, dtype=cp.int64) row_start[1:] = cp.cumsum(nnz_per_row) base = 0 @@ -474,18 +503,26 @@ def _coo_to_sell_impl(rows, cols, data, shape, slice_size): rps = int(rows_per_slice[s]) if rps == 0: continue - max_nnz = (int(slice_ptr[s + 1]) - int(slice_ptr[s])) // rps + max_nnz = (int(slice_ptr[s + 1]) - int(slice_ptr[s])) // slice_size for r in range(rps): row = r0 + r start = int(row_start[row]) end = int(row_start[row + 1]) nnz = end - start - dst_start = base + r * max_nnz if nnz > 0: - values[dst_start : dst_start + nnz] = data[start:end] - indices[dst_start : dst_start + nnz] = cols[start:end] + dst = base + cp.arange(nnz, dtype=cp.int64) * slice_size + r + values[dst] = data[start:end] + indices[dst] = cols[start:end] base = int(slice_ptr[s + 1]) - return SELLMatrix(values, indices, slice_ptr, rows_per_slice, shape, dtype=data.dtype) + return SELLMatrix( + values, + indices, + slice_ptr, + rows_per_slice, + shape, + dtype=data.dtype, + slice_size=slice_size, + ) def _coo_to_blocked_ell_impl(rows, cols, data, shape, block_shape): @@ -590,10 +627,24 @@ def create_bsr_matrix(data, indices, indptr, shape, blocksize, dtype=None): return BSRMatrix(data, indices, indptr, shape, blocksize=blocksize, dtype=dtype) -def create_sell_matrix(values, indices, slice_ptr, rows_per_slice, shape, dtype=None): +def create_sell_matrix( + values, + indices, + slice_ptr, + rows_per_slice, + shape, + dtype=None, + slice_size=None, +): """Create SELL from values, indices, slice_ptr, rows_per_slice, shape.""" return SELLMatrix( - values, indices, slice_ptr, rows_per_slice, shape, dtype=dtype + values, + indices, + slice_ptr, + rows_per_slice, + shape, + dtype=dtype, + slice_size=slice_size, ) diff --git a/src/flagsparse/sparse_operations/__init__.py b/src/flagsparse/sparse_operations/__init__.py index bed39a2..2ba2665 100644 --- a/src/flagsparse/sparse_operations/__init__.py +++ b/src/flagsparse/sparse_operations/__init__.py @@ -102,6 +102,7 @@ flagsparse_spsv_coo, flagsparse_spsv_create_workspace, flagsparse_spsv_csr, + flagsparse_spsv_sell, flagsparse_spsv_preprocess_coo, flagsparse_spsv_preprocess_csr, flagsparse_spsv_solve_ex, @@ -194,6 +195,7 @@ "flagsparse_spsv_coo", "flagsparse_spsv_create_workspace", "flagsparse_spsv_csr", + "flagsparse_spsv_sell", "flagsparse_spsv_preprocess_coo", "flagsparse_spsv_preprocess_csr", "flagsparse_spsv_solve_ex", diff --git a/src/flagsparse/sparse_operations/spsv.py b/src/flagsparse/sparse_operations/spsv.py index 802e516..b3e3638 100644 --- a/src/flagsparse/sparse_operations/spsv.py +++ b/src/flagsparse/sparse_operations/spsv.py @@ -1,4 +1,4 @@ -"""Sparse triangular solve (SpSV) CSR/COO.""" +"""Sparse triangular solve (SpSV) for CSR, COO, and SELL matrices.""" from ._common import * @@ -319,6 +319,66 @@ def _prepare_spsv_inputs(data, indices, indptr, b, shape): n_cols, ) + +def _prepare_spsv_sell_inputs( + values, + col_indices, + slice_offsets, + b, + shape, + slice_size, +): + """Validate cuSPARSE-compatible column-major SELL inputs.""" + + tensors = (values, col_indices, slice_offsets, b) + if not all(torch.is_tensor(t) for t in tensors): + raise TypeError("SELL SpSV inputs must be torch.Tensor") + if any(not t.is_cuda or t.ndim != 1 for t in tensors): + raise ValueError("SELL SpSV inputs must be 1D CUDA tensors") + if len({t.device for t in tensors}) != 1: + raise ValueError("SELL SpSV inputs must use one CUDA device") + + n_rows, n_cols = int(shape[0]), int(shape[1]) + if n_rows != n_cols: + raise ValueError("SELL SpSV requires a square matrix") + slice_size = int(slice_size) + if slice_size <= 0: + raise ValueError("slice_size must be positive") + n_slices = (n_rows + slice_size - 1) // slice_size + if slice_offsets.numel() != n_slices + 1: + raise ValueError("invalid slice_offsets length") + if values.numel() != col_indices.numel() or b.numel() != n_rows: + raise ValueError("invalid SELL values, columns, or right-hand-side length") + if values.dtype not in (torch.float32, torch.float64): + raise TypeError("SELL values must be float32 or float64") + if ( + col_indices.dtype not in (torch.int32, torch.int64) + or slice_offsets.dtype != col_indices.dtype + ): + raise TypeError("SELL columns and offsets must share int32 or int64 dtype") + if b.dtype != values.dtype: + raise TypeError("b dtype must match SELL values") + + offsets = slice_offsets.contiguous() + cols = col_indices.contiguous() + slice_lengths = offsets[1:] - offsets[:-1] + if ( + int(offsets[0].item()) != 0 + or int(offsets[-1].item()) != values.numel() + or bool(torch.any(slice_lengths < 0).item()) + or bool(torch.any(slice_lengths % slice_size != 0).item()) + ): + raise ValueError("invalid SELL slice offsets") + if cols.numel() > 0: + if bool(torch.any(cols < -1).item()): + raise IndexError("SELL padding must use column index -1") + valid_cols = cols >= 0 + if bool(torch.any(cols[valid_cols] >= n_cols).item()): + raise IndexError("SELL column index is out of range") + + return values.contiguous(), cols, offsets, b.contiguous(), n_rows, slice_size + + def _spsv_diag_eps_for_dtype(value_dtype): return 1e-12 if value_dtype in (torch.float64, torch.complex128) else 1e-6 @@ -1987,6 +2047,68 @@ def _spsv_csr_cw_kernel_complex( logical_row = tl.atomic_add(row_counter_ptr, 1) +@triton.jit +def _spsv_sell_cw_kernel( + values_ptr, + col_indices_ptr, + slice_offsets_ptr, + b_ptr, + x_ptr, + ready_ptr, + row_counter_ptr, + n_rows, + SLICE_SIZE: tl.constexpr, + USE_FP64_ACC: tl.constexpr, +): + """Persistent dependency solve over cuSPARSE column-major SELL storage.""" + + logical_row = tl.atomic_add(row_counter_ptr, 1) + while logical_row < n_rows: + row = logical_row + slice_id = row // SLICE_SIZE + row_in_slice = row - slice_id * SLICE_SIZE + slice_start = tl.load(slice_offsets_ptr + slice_id) + slice_end = tl.load(slice_offsets_ptr + slice_id + 1) + width = (slice_end - slice_start) // SLICE_SIZE + if USE_FP64_ACC: + rhs = tl.load(b_ptr + row).to(tl.float64) + tmp_sum = tl.zeros((), dtype=tl.float64) + diag = tl.zeros((), dtype=tl.float64) + else: + rhs = tl.load(b_ptr + row).to(tl.float32) + tmp_sum = tl.zeros((), dtype=tl.float32) + diag = tl.zeros((), dtype=tl.float32) + slot = 0 + while slot < width: + offset = slice_start + slot * SLICE_SIZE + row_in_slice + col = tl.load(col_indices_ptr + offset) + valid = (col >= 0) & (col < n_rows) + if valid: + if col == row: + if USE_FP64_ACC: + diag = tl.load(values_ptr + offset).to(tl.float64) + else: + diag = tl.load(values_ptr + offset).to(tl.float32) + else: + is_dependency = col < row + if is_dependency: + dep_ready = tl.atomic_add(ready_ptr + col, 0) + while dep_ready != 1: + dep_ready = tl.atomic_add(ready_ptr + col, 0) + if USE_FP64_ACC: + a = tl.load(values_ptr + offset).to(tl.float64) + x_dep = tl.load(x_ptr + col).to(tl.float64) + else: + a = tl.load(values_ptr + offset).to(tl.float32) + x_dep = tl.load(x_ptr + col).to(tl.float32) + tmp_sum += a * x_dep + slot += 1 + x_row = (rhs - tmp_sum) / diag + tl.store(x_ptr + row, x_row) + _publish_ready_flag_i32(ready_ptr, row) + logical_row = tl.atomic_add(row_counter_ptr, 1) + + @triton.jit def _spsv_csr_transpose_cw_kernel( data_ptr, @@ -2958,6 +3080,38 @@ def _triton_spsv_csr_cw_vector_complex( return x +def _launch_spsv_sell( + values, + col_indices, + slice_offsets, + b_vec, + n_rows, + *, + slice_size, + out, + ready, + row_counter, +): + ready.zero_() + row_counter.zero_() + if n_rows == 0: + return out + worker_count = _snap_cw_worker_count(min(n_rows, 32), n_rows) + _spsv_sell_cw_kernel[(int(worker_count),)]( + values, + col_indices, + slice_offsets, + b_vec, + out, + ready, + row_counter, + n_rows, + SLICE_SIZE=int(slice_size), + USE_FP64_ACC=values.dtype == torch.float64, + ) + return out + + def _triton_spsv_csr_u_lo_cw_vector(*args, **kwargs): return _triton_spsv_csr_cw_vector(*args, lower=True, unit_diagonal=True, **kwargs) @@ -4529,6 +4683,35 @@ def flagsparse_spsv_solve_ex( raise ValueError("matA.format must be 'csr' or 'coo'") +def flagsparse_spsv_sell( + values, + col_indices, + slice_offsets, + b, + shape, + *, + slice_size, +): + """Solve a real non-unit lower triangle in column-major SELL format.""" + + values, cols, offsets, b, n_rows, slice_size = ( + _prepare_spsv_sell_inputs( + values, col_indices, slice_offsets, b, shape, slice_size + ) + ) + return _launch_spsv_sell( + values, + cols, + offsets, + b, + n_rows, + slice_size=slice_size, + out=torch.empty_like(b), + ready=torch.empty(n_rows, dtype=torch.int32, device=b.device), + row_counter=torch.empty(1, dtype=torch.int32, device=b.device), + ) + + def flagsparse_spsv_csr( data, indices, diff --git a/tests/ci/test_cli_help.py b/tests/ci/test_cli_help.py index 601c568..6657488 100644 --- a/tests/ci/test_cli_help.py +++ b/tests/ci/test_cli_help.py @@ -25,6 +25,7 @@ "tests/test_spmm.py", "tests/test_spgemm.py", "tests/test_spsv.py", + "tests/test_spsv_sell.py", "tests/test_spsm.py", ] diff --git a/tests/ci/test_gather_contract.py b/tests/ci/test_gather_contract.py new file mode 100644 index 0000000..e83b79f --- /dev/null +++ b/tests/ci/test_gather_contract.py @@ -0,0 +1,74 @@ +"""Static contract checks for the gather benchmark entrypoint.""" + +import ast +from pathlib import Path + + +PROJECT_ROOT = Path(__file__).resolve().parents[2] +BENCHMARKS_PATH = PROJECT_ROOT / "src" / "flagsparse" / "sparse_operations" / "benchmarks.py" +GATHER_TEST_PATH = PROJECT_ROOT / "tests" / "test_gather.py" + + +def _tree(path): + return ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) + + +def _benchmark_gather_parameters(): + for node in _tree(BENCHMARKS_PATH).body: + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): + if node.name == "benchmark_gather_case": + return {arg.arg for arg in node.args.args} + raise AssertionError("benchmark_gather_case definition not found") + + +def _gather_test_keywords(): + keywords = set() + for node in ast.walk(_tree(GATHER_TEST_PATH)): + if not isinstance(node, ast.Call): + continue + func = node.func + if ( + isinstance(func, ast.Attribute) + and isinstance(func.value, ast.Name) + and func.value.id == "ast" + and func.attr == "benchmark_gather_case" + ): + keywords.update(keyword.arg for keyword in node.keywords if keyword.arg) + assert keywords, "tests/test_gather.py does not call ast.benchmark_gather_case" + return keywords + + +def test_gather_cli_keywords_match_benchmark_signature(): + assert _gather_test_keywords() <= _benchmark_gather_parameters() + + +def test_gather_csv_rows_cover_all_summary_fields(): + tree = _tree(GATHER_TEST_PATH) + summary_fields = None + row_key_sets = [] + + for node in ast.walk(tree): + if not isinstance(node, ast.Assign) or len(node.targets) != 1: + continue + target = node.targets[0] + if not isinstance(target, ast.Name): + continue + if target.id == "summary_fields" and isinstance(node.value, ast.List): + summary_fields = { + item.value + for item in node.value.elts + if isinstance(item, ast.Constant) and isinstance(item.value, str) + } + if target.id == "row" and isinstance(node.value, ast.Dict): + keys = { + key.value + for key in node.value.keys + if isinstance(key, ast.Constant) and isinstance(key.value, str) + } + if "case_id" in keys and "status" in keys: + row_key_sets.append(keys) + + assert summary_fields, "summary_fields list not found" + assert len(row_key_sets) == 2, "expected success and error summary row dictionaries" + for row_keys in row_key_sets: + assert summary_fields <= row_keys diff --git a/tests/ci/test_operator_registry.py b/tests/ci/test_operator_registry.py index d88b784..283dd32 100644 --- a/tests/ci/test_operator_registry.py +++ b/tests/ci/test_operator_registry.py @@ -51,6 +51,7 @@ def test_operator_registry_interfaces_are_public_exports(): "public-api", "csr", "coo", + "sell", "optimized", "alpha-sparse", "descriptor", diff --git a/tests/ci/test_public_api.py b/tests/ci/test_public_api.py index 6843dbd..64014c3 100644 --- a/tests/ci/test_public_api.py +++ b/tests/ci/test_public_api.py @@ -17,6 +17,7 @@ "flagsparse_sddmm_csr", "flagsparse_spsv_csr", "flagsparse_spsv_coo", + "flagsparse_spsv_sell", "flagsparse_spsm_csr", "flagsparse_spsm_coo", "create_csr_matrix", diff --git a/tests/test_spsv_sell.py b/tests/test_spsv_sell.py new file mode 100644 index 0000000..0c78ce4 --- /dev/null +++ b/tests/test_spsv_sell.py @@ -0,0 +1,651 @@ +"""SELL SpSV correctness and fair Triton/cuSPARSE timing checks.""" + +import argparse +import csv +import ctypes +import ctypes.util +import glob +import os +import sys +from pathlib import Path + +_PROJECT_ROOT = Path(__file__).resolve().parents[1] +_SRC_ROOT = _PROJECT_ROOT / "src" +for path in (_PROJECT_ROOT, _SRC_ROOT): + if str(path) not in sys.path: + sys.path.insert(0, str(path)) + +import torch + +if __name__ != "__main__": + import pytest + +import flagsparse as fs +from flagsparse.sparse_operations import spsv as spsv_impl +from tests.test_spsv import ( + _allinone_filtered_avg_ms, + _apply_csr_op, + _build_random_triangular_csr, + _dtype_name, + _fmt_err, + _fmt_ms, + _fmt_ratio, + _load_mtx_to_csr_torch, + _random_rhs_for_spsv, +) + + +if __name__ != "__main__": + pytestmark = pytest.mark.spsv_sell + +VALUE_DTYPES = (torch.float32, torch.float64) +INDEX_DTYPES = (torch.int32, torch.int64) +WARMUP = 1 +ITERS = 1 + +CSV_FIELDS = [ + "matrix", + "value_dtype", + "index_dtype", + "opA", + "n_rows", + "n_cols", + "nnz", + "FlagSparse_ms", + "cuSPARSE_ms", + "PyTorch_ms", + "FlagSparse_vs_cuSPARSE_speedup", + "FlagSparse_vs_PyTorch_speedup", + "status", + "err_ref", + "err_res", + "err_pt", + "err_cu", + "pytorch_reason", + "error", +] + +_SUCCESS = 0 +_NON_TRANSPOSE = 0 +_INDEX_BASE_ZERO = 0 +_INDEX_32I = 2 +_INDEX_64I = 3 +_SPMAT_FILL_MODE = 0 +_SPMAT_DIAG_TYPE = 1 +_FILL_MODE_LOWER = 0 +_DIAG_TYPE_NON_UNIT = 0 +_SPSV_ALG_DEFAULT = 0 +_CUDA_R_32F = 0 +_CUDA_R_64F = 1 + + +def _check(status, name): + if int(status) != _SUCCESS: + raise RuntimeError(f"{name} failed with cuSPARSE status {int(status)}") + + +def _cuda_dtype(dtype): + return {torch.float32: _CUDA_R_32F, torch.float64: _CUDA_R_64F}[dtype] + + +def _index_dtype(dtype): + return {torch.int32: _INDEX_32I, torch.int64: _INDEX_64I}[dtype] + + +def _configure_cusparse(lib): + p = ctypes.c_void_p + pp = ctypes.POINTER(p) + i = ctypes.c_int + i64 = ctypes.c_int64 + + signatures = { + "cusparseCreate": ([pp], i), + "cusparseDestroy": ([p], i), + "cusparseSetStream": ([p, p], i), + "cusparseCreateSlicedEll": ( + [pp, i64, i64, i64, i64, i64, p, p, p, i, i, i, i], + i, + ), + "cusparseDestroySpMat": ([p], i), + "cusparseSpMatSetAttribute": ([p, i, p, ctypes.c_size_t], i), + "cusparseCreateDnVec": ([pp, i64, p, i], i), + "cusparseDestroyDnVec": ([p], i), + "cusparseSpSV_createDescr": ([pp], i), + "cusparseSpSV_destroyDescr": ([p], i), + "cusparseSpSV_bufferSize": ( + [p, i, p, p, p, p, i, i, p, ctypes.POINTER(ctypes.c_size_t)], + i, + ), + "cusparseSpSV_analysis": ([p, i, p, p, p, p, i, i, p, p], i), + "cusparseSpSV_solve": ([p, i, p, p, p, p, i, i, p], i), + } + for name, (argtypes, restype) in signatures.items(): + function = getattr(lib, name) + function.argtypes = argtypes + function.restype = restype + + +def _load_cusparse(): + name = ctypes.util.find_library("cusparse") or "libcusparse.so.12" + lib = ctypes.CDLL(name) + _configure_cusparse(lib) + return lib + + +def _stream_ptr(): + return ctypes.c_void_p(int(torch.cuda.current_stream().cuda_stream)) + + +def _csr_to_sell(values, cols, row_ptr, n_rows, slice_size): + """Convert CSR to cuSPARSE's column-major Sliced ELLPACK layout.""" + + slice_size = int(slice_size) + n_slices = (n_rows + slice_size - 1) // slice_size + widths = [] + for slice_id in range(n_slices): + row0 = slice_id * slice_size + row1 = min(row0 + slice_size, n_rows) + widths.append( + max( + int(row_ptr[row + 1].item() - row_ptr[row].item()) + for row in range(row0, row1) + ) + ) + + offsets = torch.zeros( + n_slices + 1, dtype=row_ptr.dtype, device=row_ptr.device + ) + if widths: + increments = torch.tensor( + [width * slice_size for width in widths], + dtype=row_ptr.dtype, + device=row_ptr.device, + ) + offsets[1:] = torch.cumsum(increments, dim=0) + + padded_size = int(offsets[-1].item()) + sell_values = torch.zeros( + padded_size, dtype=values.dtype, device=values.device + ) + sell_cols = torch.full( + (padded_size,), -1, dtype=cols.dtype, device=cols.device + ) + for slice_id in range(n_slices): + row0 = slice_id * slice_size + row1 = min(row0 + slice_size, n_rows) + base = int(offsets[slice_id].item()) + for row in range(row0, row1): + start = int(row_ptr[row].item()) + end = int(row_ptr[row + 1].item()) + count = end - start + dst = ( + base + + torch.arange(count, device=values.device) * slice_size + + row + - row0 + ) + sell_values[dst] = values[start:end] + sell_cols[dst] = cols[start:end] + return sell_values, sell_cols, offsets + + +def _time_cuda(run, warmup=None, iters=None): + warmup = WARMUP if warmup is None else int(warmup) + iters = ITERS if iters is None else int(iters) + for _ in range(warmup): + run() + torch.cuda.synchronize() + + samples = [] + output = None + for _ in range(iters): + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + output = run() + end.record() + end.synchronize() + samples.append(float(start.elapsed_time(end))) + return output, _allinone_filtered_avg_ms(samples, fmt="SELL") + + +class _CusparseSellSpSV: + """Minimal native cuSPARSE SELL baseline with reusable descriptors/workspace.""" + + def __init__(self, values, cols, offsets, b, n_rows, nnz, slice_size): + self.lib = _load_cusparse() + self.values = values + self.cols = cols + self.offsets = offsets + self.b = b + self.x = torch.empty_like(b) + self.value_type = _cuda_dtype(values.dtype) + self.alpha = ( + ctypes.c_float(1.0) + if values.dtype == torch.float32 + else ctypes.c_double(1.0) + ) + self.handle = ctypes.c_void_p() + self.matrix = ctypes.c_void_p() + self.vec_b = ctypes.c_void_p() + self.vec_x = ctypes.c_void_p() + self.descr = ctypes.c_void_p() + self.workspace = None + + _check(self.lib.cusparseCreate(ctypes.byref(self.handle)), "cusparseCreate") + _check( + self.lib.cusparseSetStream(self.handle, _stream_ptr()), + "cusparseSetStream", + ) + _check( + self.lib.cusparseCreateSlicedEll( + ctypes.byref(self.matrix), + ctypes.c_int64(n_rows), + ctypes.c_int64(n_rows), + ctypes.c_int64(nnz), + ctypes.c_int64(values.numel()), + ctypes.c_int64(slice_size), + ctypes.c_void_p(offsets.data_ptr()), + ctypes.c_void_p(cols.data_ptr()), + ctypes.c_void_p(values.data_ptr()), + ctypes.c_int(_index_dtype(offsets.dtype)), + ctypes.c_int(_index_dtype(cols.dtype)), + ctypes.c_int(_INDEX_BASE_ZERO), + ctypes.c_int(self.value_type), + ), + "cusparseCreateSlicedEll", + ) + + fill = ctypes.c_int(_FILL_MODE_LOWER) + diag = ctypes.c_int(_DIAG_TYPE_NON_UNIT) + for attribute, value in ( + (_SPMAT_FILL_MODE, fill), + (_SPMAT_DIAG_TYPE, diag), + ): + _check( + self.lib.cusparseSpMatSetAttribute( + self.matrix, + ctypes.c_int(attribute), + ctypes.byref(value), + ctypes.sizeof(value), + ), + "cusparseSpMatSetAttribute", + ) + + for target, tensor in (("vec_b", b), ("vec_x", self.x)): + descriptor = ctypes.c_void_p() + _check( + self.lib.cusparseCreateDnVec( + ctypes.byref(descriptor), + ctypes.c_int64(n_rows), + ctypes.c_void_p(tensor.data_ptr()), + ctypes.c_int(self.value_type), + ), + "cusparseCreateDnVec", + ) + setattr(self, target, descriptor) + + _check( + self.lib.cusparseSpSV_createDescr(ctypes.byref(self.descr)), + "cusparseSpSV_createDescr", + ) + size = ctypes.c_size_t() + _check( + self.lib.cusparseSpSV_bufferSize( + self.handle, + ctypes.c_int(_NON_TRANSPOSE), + ctypes.byref(self.alpha), + self.matrix, + self.vec_b, + self.vec_x, + ctypes.c_int(self.value_type), + ctypes.c_int(_SPSV_ALG_DEFAULT), + self.descr, + ctypes.byref(size), + ), + "cusparseSpSV_bufferSize", + ) + self.workspace = torch.empty( + max(1, size.value), dtype=torch.uint8, device=values.device + ) + self.workspace_ptr = ( + ctypes.c_void_p(self.workspace.data_ptr()) + if size.value + else ctypes.c_void_p() + ) + + def analysis_and_solve(self): + _check( + self.lib.cusparseSpSV_analysis( + self.handle, + ctypes.c_int(_NON_TRANSPOSE), + ctypes.byref(self.alpha), + self.matrix, + self.vec_b, + self.vec_x, + ctypes.c_int(self.value_type), + ctypes.c_int(_SPSV_ALG_DEFAULT), + self.descr, + self.workspace_ptr, + ), + "cusparseSpSV_analysis", + ) + _check( + self.lib.cusparseSpSV_solve( + self.handle, + ctypes.c_int(_NON_TRANSPOSE), + ctypes.byref(self.alpha), + self.matrix, + self.vec_b, + self.vec_x, + ctypes.c_int(self.value_type), + ctypes.c_int(_SPSV_ALG_DEFAULT), + self.descr, + ), + "cusparseSpSV_solve", + ) + return self.x + + def close(self): + if self.descr.value: + self.lib.cusparseSpSV_destroyDescr(self.descr) + if self.vec_x.value: + self.lib.cusparseDestroyDnVec(self.vec_x) + if self.vec_b.value: + self.lib.cusparseDestroyDnVec(self.vec_b) + if self.matrix.value: + self.lib.cusparseDestroySpMat(self.matrix) + if self.handle.value: + self.lib.cusparseDestroy(self.handle) + + +def _benchmark_triton(values, cols, offsets, b, n_rows, slice_size): + ready = torch.empty(n_rows, dtype=torch.int32, device=b.device) + row_counter = torch.empty(1, dtype=torch.int32, device=b.device) + out = torch.empty_like(b) + + def analysis_and_solve(): + return spsv_impl._launch_spsv_sell( + values, + cols, + offsets, + b, + n_rows, + slice_size=slice_size, + out=out, + ready=ready, + row_counter=row_counter, + ) + + return _time_cuda(analysis_and_solve) + + +def _run_case( + matrix, + values, + cols, + row_ptr, + b, + expected, + slice_size, +): + n_rows = int(row_ptr.numel() - 1) + sell_values, sell_cols, offsets = _csr_to_sell( + values, cols, row_ptr, n_rows, slice_size + ) + public_result = fs.flagsparse_spsv_sell( + sell_values, + sell_cols, + offsets, + b, + (n_rows, n_rows), + slice_size=slice_size, + ) + triton_result, triton_ms = _benchmark_triton( + sell_values, sell_cols, offsets, b, n_rows, slice_size + ) + baseline = _CusparseSellSpSV( + sell_values, + sell_cols, + offsets, + b, + n_rows, + values.numel(), + slice_size, + ) + try: + cusparse_result, cusparse_ms = _time_cuda(baseline.analysis_and_solve) + finally: + baseline.close() + + err_ref = float(torch.max(torch.abs(public_result - expected)).item()) + err_res = float( + torch.max( + torch.abs( + _apply_csr_op( + values, + cols, + row_ptr, + triton_result, + (n_rows, n_rows), + "NON", + lower=True, + ) + - b + ) + ).item() + ) + err_cu = float(torch.max(torch.abs(triton_result - cusparse_result)).item()) + atol = 2e-5 if values.dtype == torch.float32 else 1e-11 + rtol = 2e-5 if values.dtype == torch.float32 else 1e-11 + status = ( + "PASS" + if torch.allclose(triton_result, expected, atol=atol, rtol=rtol) + and torch.allclose(triton_result, cusparse_result, atol=atol, rtol=rtol) + else "FAIL" + ) + record = { + "matrix": matrix, + "value_dtype": _dtype_name(values.dtype), + "index_dtype": _dtype_name(cols.dtype), + "opA": "NON", + "n_rows": n_rows, + "n_cols": n_rows, + "nnz": int(values.numel()), + "FlagSparse_ms": triton_ms, + "cuSPARSE_ms": cusparse_ms, + "PyTorch_ms": None, + "FlagSparse_vs_cuSPARSE_speedup": cusparse_ms / triton_ms, + "FlagSparse_vs_PyTorch_speedup": None, + "status": status, + "err_ref": err_ref, + "err_res": err_res, + "err_pt": None, + "err_cu": err_cu, + "pytorch_reason": "not used for SELL", + "error": None, + } + return record, triton_result, cusparse_result + + +def _print_header(slice_size, value_dtype, index_dtype): + print("=" * 144) + print( + f"Value dtype: {_dtype_name(value_dtype)} | " + f"Index dtype: {_dtype_name(index_dtype)} | SELL | " + f"triA=LOWER | opA=NON | slice_size={slice_size}" + ) + print( + f"Benchmark schedule: warmup={WARMUP}, iter={ITERS}; " + "FS.ms and CU.ms both include every per-call analysis/preparation + solve." + ) + print("CU.spd = CU.ms / FS.ms; PT.spd = PT.ms / FS.ms.") + print("-" * 144) + print( + f"{'Matrix':<28} {'N_rows':>7} {'N_cols':>7} {'NNZ':>10} " + f"{'FS.ms':>10} {'CU.ms':>10} {'PT.ms':>10} " + f"{'CU.spd':>10} {'PT.spd':>10} {'Status':>6} " + f"{'Eref':>10} {'Eres':>10} {'Ept':>10} {'Ecu':>10}" + ) + print("-" * 144) + + +def _print_record(record): + name = str(record["matrix"]) + name = name[:27] + ("…" if len(name) > 27 else "") + print( + f"{name:<28} {record['n_rows']:>7} {record['n_cols']:>7} " + f"{record['nnz']:>10} " + f"{_fmt_ms(record['FlagSparse_ms']):>10} " + f"{_fmt_ms(record['cuSPARSE_ms']):>10} " + f"{_fmt_ms(record['PyTorch_ms']):>10} " + f"{_fmt_ratio(record['FlagSparse_vs_cuSPARSE_speedup']):>10} " + f"{_fmt_ratio(record['FlagSparse_vs_PyTorch_speedup']):>10} " + f"{record['status']:>6} {_fmt_err(record['err_ref']):>10} " + f"{_fmt_err(record['err_res']):>10} {_fmt_err(record['err_pt']):>10} " + f"{_fmt_err(record['err_cu']):>10}" + ) + + +def test_spsv_sell_matches_cusparse(value_dtype, index_dtype, slice_size): + if not torch.cuda.is_available(): + pytest.skip("CUDA is unavailable") + + n_rows = 64 + values, cols, row_ptr, shape = _build_random_triangular_csr( + n_rows, + value_dtype, + index_dtype, + torch.device("cuda"), + lower=True, + ) + row_ptr = row_ptr.to(index_dtype) + expected = _random_rhs_for_spsv( + shape, value_dtype, values.device, op_mode="NON", seed=1234 + ) + b = _apply_csr_op( + values, cols, row_ptr, expected, shape, "NON", lower=True + ) + atol = 2e-5 if value_dtype == torch.float32 else 1e-11 + rtol = 2e-5 if value_dtype == torch.float32 else 1e-11 + try: + record, triton_result, cusparse_result = _run_case( + "synthetic-64", values, cols, row_ptr, b, expected, slice_size + ) + except (AttributeError, OSError, RuntimeError) as exc: + pytest.skip(f"native cuSPARSE SELL SpSV is unavailable: {exc}") + assert torch.allclose(triton_result, cusparse_result, atol=atol, rtol=rtol) + assert record["status"] == "PASS" + assert record["FlagSparse_ms"] > 0.0 + assert record["cuSPARSE_ms"] > 0.0 + _print_header(slice_size, value_dtype, index_dtype) + _print_record(record) + + +if __name__ != "__main__": + test_spsv_sell_matches_cusparse = pytest.mark.parametrize( + "value_dtype", VALUE_DTYPES + )( + pytest.mark.parametrize("index_dtype", INDEX_DTYPES)( + pytest.mark.parametrize("slice_size", (8, 32))( + test_spsv_sell_matches_cusparse + ) + ) + ) + + +def _expand_mtx_paths(inputs): + paths = [] + for value in inputs: + if os.path.isdir(value): + paths.extend(sorted(glob.glob(os.path.join(value, "*.mtx")))) + elif value.endswith(".mtx"): + paths.append(value) + return paths + + +def main(): + global WARMUP, ITERS + parser = argparse.ArgumentParser( + description="Lower real SELL SpSV: Triton versus cuSPARSE analysis+solve" + ) + parser.add_argument("mtx", nargs="+", help=".mtx files or directories") + parser.add_argument("--csv", required=True, help="output CSV path") + parser.add_argument("--slice-size", type=int, default=32) + parser.add_argument("--warmup", type=int, default=WARMUP) + parser.add_argument("--iters", type=int, default=ITERS) + args = parser.parse_args() + + if not torch.cuda.is_available(): + raise SystemExit("CUDA is unavailable") + WARMUP = max(0, args.warmup) + ITERS = max(1, args.iters) + paths = _expand_mtx_paths(args.mtx) + if not paths: + raise SystemExit("No .mtx files found") + + records = [] + for value_dtype in VALUE_DTYPES: + for index_dtype in INDEX_DTYPES: + _print_header(args.slice_size, value_dtype, index_dtype) + for path in paths: + try: + values, cols, row_ptr, shape = _load_mtx_to_csr_torch( + path, + dtype=value_dtype, + device=torch.device("cuda"), + lower=True, + ) + if int(shape[0]) != int(shape[1]): + raise ValueError(f"SpSV requires a square matrix, got {shape}") + cols = cols.to(index_dtype) + row_ptr = row_ptr.to(index_dtype) + expected = torch.ones( + shape[0], dtype=value_dtype, device=values.device + ) + b = _apply_csr_op( + values, + cols, + row_ptr, + expected, + shape, + "NON", + lower=True, + ) + record, _, _ = _run_case( + os.path.basename(path), + values, + cols, + row_ptr, + b, + expected, + args.slice_size, + ) + records.append(record) + _print_record(record) + except Exception as exc: + record = {key: None for key in CSV_FIELDS} + record.update( + matrix=os.path.basename(path), + value_dtype=_dtype_name(value_dtype), + index_dtype=_dtype_name(index_dtype), + opA="NON", + status="ERROR", + error=str(exc), + ) + records.append(record) + print(f"{record['matrix']:<28} ERROR: {exc}") + + with open(args.csv, "w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=CSV_FIELDS) + writer.writeheader() + for record in records: + writer.writerow( + {key: "" if record.get(key) is None else record.get(key) for key in CSV_FIELDS} + ) + print("-" * 144) + print(f"Wrote {len(records)} rows to {args.csv}") + + +if __name__ == "__main__": + main()