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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 15 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,14 @@ python tests/test_spmv_opt.py <dir_or_file.mtx> [...]
python tests/test_spmv_opt.py <dir/> --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 <dir/> --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
Expand Down Expand Up @@ -119,10 +127,17 @@ python tests/test_spgemm.py <dir/> --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 <dir/> --csv-csr spsv.csv
python tests/test_spsv.py <dir/> --csv-coo out.csv # same CSV columns as CSR
pytest -q -s tests/test_spsv_sell.py
python tests/test_spsv_sell.py <dir_or_file.mtx> --csv sell.csv --slice-size 32
```

**test_spsm.py** - SpSM (triangular matrix-matrix solve; **square** matrices only):
Expand Down
6 changes: 6 additions & 0 deletions README_cn.md
Original file line number Diff line number Diff line change
Expand Up @@ -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(三角矩阵-稠密矩阵求解;**仅方阵**):
Expand Down
51 changes: 51 additions & 0 deletions conf/operators.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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.
Expand Down
58 changes: 58 additions & 0 deletions ops_support.csv
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
13 changes: 13 additions & 0 deletions ops_support.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"),
Expand All @@ -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"),
)
Expand Down
3 changes: 3 additions & 0 deletions pytest.ini
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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)
Expand Down
41 changes: 41 additions & 0 deletions run_flagsparse_pytest.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,8 @@
"pytorch_ms",
"cusparse_ms",
"cupy_ms",
"csc_ms",
"bsr_ms",
"base_ms",
"alg1_ms",
"alg2_ms",
Expand All @@ -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"),
Expand Down Expand Up @@ -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}",
Expand Down Expand Up @@ -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}",
Expand Down Expand Up @@ -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"]
),
Expand All @@ -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"]),
}
Expand Down
Loading
Loading