diff --git a/README.md b/README.md index ca82314..80ddd53 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 @@ -119,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 a51c456..eb1503e 100644 --- a/conf/operators.yaml +++ b/conf/operators.yaml @@ -63,6 +63,40 @@ 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_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. @@ -261,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 c766d17..da51c7e 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 @@ -225,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 7c91f0c..3849f07 100644 --- a/ops_support.py +++ b/ops_support.py @@ -228,6 +228,16 @@ 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") + ) + 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 @@ -264,6 +274,8 @@ 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 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"), @@ -273,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 378a20b..3cd9029 100644 --- a/pytest.ini +++ b/pytest.ini @@ -11,6 +11,8 @@ 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_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) @@ -20,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 ebaebe6..5af76c9 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"), @@ -162,6 +168,28 @@ class OperatorTestConfig: "--iters", "{iters}", ), + "spmv_csc": ( + "tests/test_spmv_csc.py", + "{input}", + "--csv-csc", + "{csv}", + "--warmup", + "{warmup}", + "--iters", + "{iters}", + ), + "spmv_bsr": ( + "tests/test_spmv_bsr.py", + "{input}", + "--csv-bsr", + "{csv}", + "--ops", + "non,trans,conj", + "--warmup", + "{warmup}", + "--iters", + "{iters}", + ), "spmv_coo_tocsr": ( "tests/test_spmv_coo.py", "{input}", @@ -276,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}", @@ -304,6 +342,8 @@ 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_bsr": OperatorTestConfig("spmv_bsr", PERFORMANCE_COMMANDS["spmv_bsr"]), "spmv_coo_tocsr": OperatorTestConfig( "spmv_coo_tocsr", PERFORMANCE_COMMANDS["spmv_coo_tocsr"] ), @@ -325,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 5466aae..5f4b879 100644 --- a/src/flagsparse/__init__.py +++ b/src/flagsparse/__init__.py @@ -16,6 +16,9 @@ "comprehensive_gather_test", "comprehensive_scatter_test", "PreparedCoo", + "PreparedCooSpmmRoute", + "PreparedBsrSpmv", + "PreparedCscSpmv", "PreparedAlphaSpmmAlg1", "PreparedCsrSpmv", "PreparedCsrSpmmOpt", @@ -31,13 +34,16 @@ "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", "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", "flagsparse_alpha_spmm_alg1", "flagsparse_alpha_spmm_alg1_tle", @@ -50,7 +56,9 @@ "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", "flagsparse_spmm_csr_run", "flagsparse_spmm_csr_opt_alg1", @@ -59,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", @@ -88,6 +97,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", @@ -96,9 +106,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", @@ -136,6 +151,9 @@ "comprehensive_gather_test", "comprehensive_scatter_test", "PreparedCoo", + "PreparedCooSpmmRoute", + "PreparedBsrSpmv", + "PreparedCscSpmv", "PreparedAlphaSpmmAlg1", "PreparedCsrSpmv", "PreparedCsrSpmmOpt", @@ -151,13 +169,16 @@ "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", "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", "flagsparse_alpha_spmm_alg1", "flagsparse_alpha_spmm_alg1_tle", @@ -170,7 +191,9 @@ "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", "flagsparse_spmm_csr_run", "flagsparse_spmm_csr_opt_alg1", @@ -179,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", @@ -208,6 +232,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", @@ -216,9 +241,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_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 d54af97..2ba2665 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, @@ -64,6 +74,8 @@ 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, flagsparse_spmv_coo_tocsr, @@ -90,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, @@ -109,8 +122,11 @@ __all__ = [ "PreparedCoo", + "PreparedCooSpmmRoute", "PreparedAlphaSpmmAlg1", + "PreparedBsrSpmv", "PreparedCsrSpmv", + "PreparedCscSpmv", "PreparedCsrSpmmOpt", "PreparedCsrSpmmRoute", "PreparedCsrSpmmOptAlg2", @@ -121,6 +137,8 @@ "FlagSparseDnVecDescr", "SpmmCsrAlgorithm", "SpmmCsrAlgorithmUnavailable", + "SpmmCooAlgorithm", + "SpmmCooAlgorithmUnavailable", "FlagSparseSpMatDescr", "FlagSparseSpSVDescr", "FlagSparseSpSVHandle", @@ -150,6 +168,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", @@ -158,7 +177,9 @@ "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", "flagsparse_spsm_coo", "flagsparse_spsm_csr", @@ -174,12 +195,14 @@ "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", "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", @@ -198,12 +221,17 @@ "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", + "prepare_spmv_bsr", "prepare_spmv_coo_tocsr", + "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..a83b579 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,512 @@ 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, + # 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 + 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|64|128|128|256", + "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/src/flagsparse/sparse_operations/spmv_bsr.py b/src/flagsparse/sparse_operations/spmv_bsr.py new file mode 100644 index 0000000..0779a5d --- /dev/null +++ b/src/flagsparse/sparse_operations/spmv_bsr.py @@ -0,0 +1,692 @@ +"""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", "trans", "conj") +_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): + _normalize_spmv_bsr_op(op_code) + + +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", + "padded_n_rows", + "padded_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.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: + 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 + 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 + 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 + 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 + 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) + + +@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") + 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") + 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") + 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 + 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): + 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) + 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: + 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 + + +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, + "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, + "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/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/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 0e7c15c..6657488 100644 --- a/tests/ci/test_cli_help.py +++ b/tests/ci/test_cli_help.py @@ -20,9 +20,12 @@ "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", + "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_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..64014c3 100644 --- a/tests/ci/test_public_api.py +++ b/tests/ci/test_public_api.py @@ -7,12 +7,17 @@ "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", "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/ci/test_runtime_policies.py b/tests/ci/test_runtime_policies.py index 01e94e6..115d27f 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,52 @@ 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", ["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(): with pytest.raises( ValueError, match="index_fallback_policy must be 'auto' or 'strict'" 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 new file mode 100644 index 0000000..b2446d4 --- /dev/null +++ b/tests/pytest/test_spmv_bsr_accuracy.py @@ -0,0 +1,319 @@ +import importlib + +import pytest +import torch + +from flagsparse import flagsparse_spmv_bsr, prepare_spmv_bsr +from tests.pytest.accuracy_utils import close_tolerances + + +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 = ((7, 8), (12, 9), (16, 32), (64, 96)) + + +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 _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 _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) + assert torch.allclose( + actual.to(ref_dtype), expected.to(ref_dtype), rtol=rtol, atol=atol + ) + + +@pytest.mark.spmv_bsr +@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()] +) +@pytest.mark.parametrize( + "index_dtype", [torch.int32, torch.int64], ids=["int32", "int64"] +) +@pytest.mark.parametrize("block_dim", [2, 4], ids=["block2", "block4"]) +@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(_logical_x_len(M, N, op), dtype, device) + ref = _make_ref(dense, x, op, dtype) + out = flagsparse_spmv_bsr( + data, + indices, + indptr, + x, + shape=(M, N), + block_dim=block_dim, + op=op, + index_fallback_policy="auto", + ) + 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 +@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 + 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=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) + 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 +@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( + M, N, torch.float32, torch.int32, 2, device + ) + 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 +@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( + M, N, torch.float32, torch.int64, 2, device + ) + 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 + + 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=(M, N), + block_dim=2, + op=op, + index_fallback_policy="auto", + ) + assert state["forced_once"] + 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 +def test_spmv_bsr_int64_strict_no_fallback(monkeypatch): + device = torch.device("cuda") + data, indices, indptr, _dense = _random_bsr_mn( + 12, 10, torch.float32, torch.int64, 2, 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): + 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, 10), + block_dim=2, + index_fallback_policy="strict", + ) + + +@pytest.mark.spmv_bsr +def test_spmv_bsr_non_divisible_shape_matches_dense_reference(): + device = torch.device("cuda") + 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) + 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 out.numel() == _padded_rows(M, 4) + _assert_close(out[:M], ref.to(torch.float32), torch.float32) + + +@pytest.mark.spmv_bsr +@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 + ) + 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) + 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[:logical_out], ref.to(torch.float32), torch.float32) 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_spmm_coo.py b/tests/test_spmm_coo.py index b951440..51b8ca2 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,350 @@ 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, + 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, + row, + col, + shape, + B, + prepared=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, + dtype, + warmup, + 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( + 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): + stage_t0 = None + try: + ast.resolve_spmm_coo_algorithm(alg, op, dtype) + stage_t0 = _start(f"run {alg}") + result = _time_coo_algorithm( + prepared, + B, + alg, + warmup, + iters, + timing=timing, + 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, + 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 +2035,18 @@ 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 = [] + 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 @@ -1545,81 +2056,44 @@ 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)}", + flush=True, ) - 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 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, + _dtype_name(index_dtype), + index_dtype, + op_name, + layout_name, + alg_names, + n_dense_cols, + warmup, + iters, + run_cusparse, + timing, + diagnose, + progress=True, + ) + 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 ''}", + flush=True, ) - 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", - ] - 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()}) + _write_csv(csv_path, rows, fieldnames) + _write_csv(best_path, _best_rows(rows), BEST_FIELDS) + if diagnose: + _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 +2517,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 +2546,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 +2585,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 +2599,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 +2616,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 +2628,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__": diff --git a/tests/test_spmv_bsr.py b/tests/test_spmv_bsr.py new file mode 100644 index 0000000..61a6e85 --- /dev/null +++ b/tests/test_spmv_bsr.py @@ -0,0 +1,1180 @@ +"""Native BSR SpMV benchmark and correctness script.""" + +import argparse +import csv +import glob +import math +import os +import sys +import warnings +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 = OPS +TEST_SIZES = ((64, 96), (160, 1024), (128, 256)) +DEFAULT_BLOCK_DIMS = (4,) +WARMUP = 10 +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( + "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." + ) + 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.", "") + + +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): + nnz = max(1, len(entries)) + best = None + 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 + if best is None or stored < best[0]: + best = (stored, block_dim) + 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 _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: + 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) + 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 _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), + indices, + indptr, + 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) + + +def _error_stats(actual, expected, atol, rtol): + if expected.numel() == 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) + 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): + 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, 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, + 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 _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, + "PyTorch BSR baseline requires both matrix dimensions to be divisible by block_dim", + None, + ) + 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: _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, op, warmup, iters): + padded_shape = _padded_shape(shape, block_dim) + padded_x = _pad_vector(x, padded_shape[0] if _op_transposes(op) else 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: _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, 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"): + 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)}" + + 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) + 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() + 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): + 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} {'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} {'PTPMode':>9} {'CU(ms)':>9} {'BSR/PT':>8} {'BSR/CU':>8} " + f"{'BSRErr':>10} {'PTPadErr':>10} {'B/PT':>10} {'B/PTPad':>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['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} {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} " + 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: + print(f" error: {str(error)[:240]}") + if row.get("pytorch_error"): + print(f" pt: {str(row['pytorch_error'])[:240]}") + if 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'))}, " + 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'))}" + ) + + +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)) + 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), + "index_dtype": _dtype_name(index_dtype), + "op": op, + "reference": "spmv-coo", + "block_dim": int(block_dim), + "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, + "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, + "pytorch_error": None, + "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, + "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, + } + + +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 + try: + 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: + logical_out = "SKIP" + padded_out = "SKIP" + pad_rows = "SKIP" + return { + "matrix": matrix_name, + "value_dtype": _dtype_name(dtype), + "index_dtype": _dtype_name(index_dtype), + "op": op, + "reference": "spmv-coo", + "block_dim": block_dim_value, + "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", + "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, + "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, + "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, + } + + +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"] = "unsupported BSR SpMV op" + return row + 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, op, 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 = _spmv_coo_reference(data, indices, indptr, x, shape, dtype, block_dim, op) + row["padded_out_size"] = int(bsr["out"].numel()) + 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( + { + "err": err, + "bsr_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 + pytorch_out = None + try: + row["pytorch_ms"], row["pytorch_error"], pytorch_out = _time_pytorch( + data, + indices, + indptr, + 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) + 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, + ) + 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( + data, + indices, + indptr, + x, + shape, + block_dim, + op, + warmup, + iters, + ) + except Exception as exc: + row["cusparse_error"] = str(exc) + ok = (not math.isnan(err)) and err <= 1.0 + 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.") + _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}") + 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 = [] + _print_baseline_notes(run_cusparse=run_cusparse) + 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) + resolved_block_dims = _resolve_block_dims(block_dims, entries, shape) + for block_dim in resolved_block_dims: + 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, + "reference": "spmv-coo", + "block_dim": "ERR", + "out_size": "ERR", + "padded_out_size": "ERR", + "pad_rows": "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, + "pytorch_error": None, + "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, + "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", + "reference", + "block_dim", + "out_size", + "padded_out_size", + "pad_rows", + "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", + "pytorch_error", + "pytorch_err", + "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", + "actual_at_max", + "expected_at_max", + "cusparse_ms", + "cusparse_error", + "bsr_err", + "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/tests/test_spmv_csc.py b/tests/test_spmv_csc.py new file mode 100644 index 0000000..2ca948d --- /dev/null +++ b/tests/test_spmv_csc.py @@ -0,0 +1,702 @@ +"""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}" + ) + error = row.get("error") + if error: + print(f" error: {str(error)[:240]}") + + +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) + 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: + 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 + 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): + 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, + 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 + 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: + if fail_fast: + raise + 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), + } + 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)) + 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") + ] + 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 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") + 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") + 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, + fail_fast=args.fail_fast, + ) + 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, + fail_fast=args.fail_fast, + ) + 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, + fail_fast=args.fail_fast, + ) + + +if __name__ == "__main__": + main() 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() diff --git a/tools/ci/run_gpu_benchmark.py b/tools/ci/run_gpu_benchmark.py index 1442da3..b908dad 100644 --- a/tools/ci/run_gpu_benchmark.py +++ b/tools/ci/run_gpu_benchmark.py @@ -27,6 +27,8 @@ def _parse_args() -> argparse.Namespace: "scatter", "spmv", "spmv-coo", + "spmv-csc", + "spmv-bsr", "spmm", "spmm-coo", "spsv", @@ -93,6 +95,26 @@ 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, + ], + "spmv-bsr": [ + "tests/test_spmv_bsr.py", + "--synthetic", + "--ops", + "non,trans,conj", + "--warmup", + str(args.warmup), + "--iters", + str(args.iters), + *no_cusparse, + ], "spmm": [ "tests/test_spmm.py", "--synthetic", @@ -135,6 +157,8 @@ def _command_specs( "scatter", "spmv", "spmv-coo", + "spmv-csc", + "spmv-bsr", "spmm", "spmm-coo", "spsv",