Skip to content

Commit ca1e239

Browse files
committed
spsv-coo
1 parent 590256f commit ca1e239

7 files changed

Lines changed: 825 additions & 718 deletions

File tree

spsm_csr_non.csv

Lines changed: 61 additions & 0 deletions
Large diffs are not rendered by default.

spsv_coo_non.csv

Lines changed: 241 additions & 0 deletions
Large diffs are not rendered by default.

src/flagsparse/sparse_operations/spsv.py

Lines changed: 28 additions & 229 deletions
Original file line numberDiff line numberDiff line change
@@ -1432,57 +1432,6 @@ def _spsv_csr_transpose_cw_kernel_complex(
14321432
row = tl.atomic_add(row_counter_ptr, 1)
14331433

14341434

1435-
@triton.jit
1436-
def _spsv_coo_level_kernel_real(
1437-
data_ptr,
1438-
row_ptr_ptr,
1439-
col_ptr,
1440-
b_ptr,
1441-
x_ptr,
1442-
rows_ptr,
1443-
n_level_rows,
1444-
BLOCK_NNZ: tl.constexpr,
1445-
MAX_SEGMENTS: tl.constexpr,
1446-
LOWER: tl.constexpr,
1447-
UNIT_DIAG: tl.constexpr,
1448-
DIAG_EPS: tl.constexpr,
1449-
):
1450-
pid = tl.program_id(0)
1451-
if pid >= n_level_rows:
1452-
return
1453-
row = tl.load(rows_ptr + pid)
1454-
start = tl.load(row_ptr_ptr + row)
1455-
end = tl.load(row_ptr_ptr + row + 1)
1456-
acc = tl.load(data_ptr + start, mask=start < end, other=0.0) * 0
1457-
diag = tl.load(data_ptr + start, mask=start < end, other=0.0) * 0
1458-
if UNIT_DIAG:
1459-
diag = diag + 1.0
1460-
1461-
for seg in range(MAX_SEGMENTS):
1462-
idx = start + seg * BLOCK_NNZ
1463-
offsets = idx + tl.arange(0, BLOCK_NNZ)
1464-
mask = offsets < end
1465-
a = tl.load(data_ptr + offsets, mask=mask, other=0.0)
1466-
col = tl.load(col_ptr + offsets, mask=mask, other=0)
1467-
x_vals = tl.load(x_ptr + col, mask=mask, other=0.0)
1468-
1469-
if LOWER:
1470-
solved = col < row
1471-
else:
1472-
solved = col > row
1473-
is_diag = col == row
1474-
1475-
acc = acc + tl.sum(tl.where(mask & solved, a * x_vals, 0.0))
1476-
if not UNIT_DIAG:
1477-
diag = diag + tl.sum(tl.where(mask & is_diag, a, 0.0))
1478-
1479-
rhs = tl.load(b_ptr + row)
1480-
diag_safe = tl.where(tl.abs(diag) < DIAG_EPS, 1.0, diag)
1481-
x_row = (rhs - acc) / diag_safe
1482-
x_row = tl.where(x_row == x_row, x_row, 0.0)
1483-
tl.store(x_ptr + row, x_row)
1484-
1485-
14861435
def _build_spsv_levels(indptr, indices, n_rows, lower=True):
14871436
"""Build dependency levels for triangular solve so each level can run in parallel."""
14881437
if n_rows == 0:
@@ -2111,6 +2060,11 @@ def _prepare_spsv_coo_inputs(data, row, col, b, shape):
21112060
raise TypeError("row dtype must be torch.int32 or torch.int64")
21122061
if col.dtype not in SUPPORTED_SPSV_INDEX_DTYPES:
21132062
raise TypeError("col dtype must be torch.int32 or torch.int64")
2063+
input_index_dtype = (
2064+
torch.int64
2065+
if row.dtype == torch.int64 or col.dtype == torch.int64
2066+
else torch.int32
2067+
)
21142068
row64 = row.to(torch.int64).contiguous()
21152069
col64 = col.to(torch.int64).contiguous()
21162070
if col64.numel() > 0 and int(col64.max().item()) > _INDEX_LIMIT_INT32:
@@ -2129,9 +2083,9 @@ def _prepare_spsv_coo_inputs(data, row, col, b, shape):
21292083
if max_col >= n_cols:
21302084
raise IndexError(f"col indices out of range for n_cols={n_cols}")
21312085

2132-
_validate_spsv_non_trans_combo(data.dtype, torch.int32, "COO")
21332086
return (
21342087
data.contiguous(),
2088+
input_index_dtype,
21352089
row64,
21362090
col64,
21372091
b.contiguous(),
@@ -2160,16 +2114,6 @@ def _csr_transpose(data, indices64, indptr64, n_rows, n_cols, conjugate=False):
21602114
return data_t, indices_t, indptr_t
21612115

21622116

2163-
def _coo_is_sorted_unique(row64, col64, n_cols):
2164-
nnz = row64.numel()
2165-
if nnz <= 1:
2166-
return True
2167-
key = row64 * max(1, n_cols) + col64
2168-
is_sorted = bool(torch.all(key[1:] >= key[:-1]).item())
2169-
is_unique = bool(torch.all(key[1:] != key[:-1]).item())
2170-
return is_sorted and is_unique
2171-
2172-
21732117
def _build_coo_row_ptr(row_sorted, n_rows):
21742118
row_ptr = torch.zeros(n_rows + 1, dtype=torch.int64, device=row_sorted.device)
21752119
if row_sorted.numel() > 0:
@@ -2206,53 +2150,6 @@ def _coo_to_csr_sorted_unique(data, row64, col64, n_rows, n_cols):
22062150
return data_u, indices, indptr
22072151

22082152

2209-
def _triton_spsv_coo_vector(
2210-
data,
2211-
cols,
2212-
row_ptr,
2213-
b_vec,
2214-
n_rows,
2215-
lower=True,
2216-
unit_diagonal=False,
2217-
block_nnz=None,
2218-
max_segments=None,
2219-
diag_eps=1e-12,
2220-
levels=None,
2221-
block_nnz_use=None,
2222-
max_segments_use=None,
2223-
):
2224-
x = torch.zeros_like(b_vec)
2225-
if n_rows == 0:
2226-
return x
2227-
if levels is None:
2228-
levels = _build_spsv_levels(row_ptr, cols, n_rows, lower=lower)
2229-
if block_nnz_use is None or max_segments_use is None:
2230-
block_nnz_use, max_segments_use = _auto_spsv_launch_config(
2231-
row_ptr, block_nnz=block_nnz, max_segments=max_segments
2232-
)
2233-
2234-
for rows_lv in levels:
2235-
n_lv = rows_lv.numel()
2236-
if n_lv == 0:
2237-
continue
2238-
grid = (n_lv,)
2239-
_spsv_coo_level_kernel_real[grid](
2240-
data,
2241-
row_ptr,
2242-
cols,
2243-
b_vec,
2244-
x,
2245-
rows_lv,
2246-
n_level_rows=n_lv,
2247-
BLOCK_NNZ=block_nnz_use,
2248-
MAX_SEGMENTS=max_segments_use,
2249-
LOWER=lower,
2250-
UNIT_DIAG=unit_diagonal,
2251-
DIAG_EPS=diag_eps,
2252-
)
2253-
return x
2254-
2255-
22562153
def flagsparse_spsv_csr(
22572154
data,
22582155
indices,
@@ -2768,7 +2665,6 @@ def _analyze_spsv_csr(
27682665

27692666

27702667
def flagsparse_spsv_coo(
2771-
27722668
data,
27732669
row,
27742670
col,
@@ -2777,134 +2673,37 @@ def flagsparse_spsv_coo(
27772673
lower=True,
27782674
unit_diagonal=False,
27792675
transpose=False,
2780-
coo_mode="auto",
27812676
block_nnz=None,
27822677
max_segments=None,
27832678
out=None,
27842679
return_time=False,
27852680
):
2786-
"""COO SpSV with dual mode:
2787-
- direct: use COO level kernel directly (requires sorted+unique COO)
2788-
- csr: convert COO -> CSR (sorted+deduplicated) then call flagsparse_spsv_csr
2789-
- auto: pick direct when sorted+unique and supported, otherwise csr
2790-
2791-
Notes:
2792-
- direct mode currently supports only non-transposed real-valued inputs
2793-
- complex dtypes and TRANS/CONJ always route through the CSR implementation
2794-
"""
2795-
data, row64, col64, b, n_rows, n_cols = _prepare_spsv_coo_inputs(
2681+
"""COO SpSV by canonicalizing COO into CSR, then reusing CSR SpSV."""
2682+
data, input_index_dtype, row64, col64, b, n_rows, n_cols = _prepare_spsv_coo_inputs(
27962683
data, row, col, b, shape
27972684
)
27982685
if n_rows != n_cols:
27992686
raise ValueError(f"A must be square, got shape={shape}")
28002687

2801-
mode = str(coo_mode).lower()
2802-
if mode not in ("auto", "direct", "csr"):
2803-
raise ValueError("coo_mode must be one of: 'auto', 'direct', 'csr'")
2804-
2805-
sorted_unique = _coo_is_sorted_unique(row64, col64, n_cols)
28062688
trans_mode = _normalize_spsv_transpose_mode(transpose)
2807-
direct_supported = (trans_mode == "N") and (not torch.is_complex(data))
2808-
use_direct = direct_supported and (mode == "direct" or (mode == "auto" and sorted_unique))
2809-
if mode == "direct" and not direct_supported:
2810-
raise ValueError(
2811-
"coo_mode='direct' supports only non-transposed real-valued inputs; "
2812-
"use coo_mode='csr' or 'auto' for TRANS/CONJ or complex dtypes"
2813-
)
2814-
if mode == "direct" and not sorted_unique:
2815-
raise ValueError(
2816-
"coo_mode='direct' requires COO sorted by (row, col) with no duplicate coordinates; "
2817-
"use coo_mode='csr' or 'auto' for unsorted/duplicate COO input"
2818-
)
2819-
2820-
if not use_direct:
2821-
data_csr, indices_csr, indptr_csr = _coo_to_csr_sorted_unique(
2822-
data, row64, col64, n_rows, n_cols
2823-
)
2824-
return flagsparse_spsv_csr(
2825-
data_csr,
2826-
indices_csr,
2827-
indptr_csr,
2828-
b,
2829-
shape,
2830-
lower=lower,
2831-
unit_diagonal=unit_diagonal,
2832-
transpose=transpose,
2833-
block_nnz=block_nnz,
2834-
max_segments=max_segments,
2835-
out=out,
2836-
return_time=return_time,
2837-
)
2838-
2839-
kernel_cols = col64.to(torch.int32)
2840-
row_ptr = _build_coo_row_ptr(row64, n_rows)
2841-
2842-
compute_dtype = data.dtype
2843-
data_in = data
2844-
b_in = b
2845-
if data.dtype == torch.float32 and SPSV_PROMOTE_FP32_TO_FP64:
2846-
compute_dtype = torch.float64
2847-
data_in = data.to(torch.float64)
2848-
b_in = b.to(torch.float64)
2849-
levels = _build_spsv_levels(row_ptr, kernel_cols, n_rows, lower=lower)
2850-
block_nnz_use, max_segments_use = _auto_spsv_launch_config(
2851-
row_ptr, block_nnz=block_nnz, max_segments=max_segments
2852-
)
2853-
diag_eps = _spsv_diag_eps_for_dtype(compute_dtype)
2854-
2855-
if return_time:
2856-
torch.cuda.synchronize()
2857-
t0 = time.perf_counter()
2858-
if b_in.ndim == 1:
2859-
x = _triton_spsv_coo_vector(
2860-
data_in,
2861-
kernel_cols,
2862-
row_ptr,
2863-
b_in,
2864-
n_rows,
2865-
lower=lower,
2866-
unit_diagonal=unit_diagonal,
2867-
block_nnz=block_nnz,
2868-
max_segments=max_segments,
2869-
diag_eps=diag_eps,
2870-
levels=levels,
2871-
block_nnz_use=block_nnz_use,
2872-
max_segments_use=max_segments_use,
2873-
)
2689+
if trans_mode == "N":
2690+
_validate_spsv_non_trans_combo(data.dtype, input_index_dtype, "COO")
28742691
else:
2875-
b_cols = b_in if b_in.is_contiguous() else b_in.contiguous()
2876-
cols_out = []
2877-
for bj in torch.unbind(b_cols, dim=1):
2878-
cols_out.append(
2879-
_triton_spsv_coo_vector(
2880-
data_in,
2881-
kernel_cols,
2882-
row_ptr,
2883-
bj,
2884-
n_rows,
2885-
lower=lower,
2886-
unit_diagonal=unit_diagonal,
2887-
block_nnz=block_nnz,
2888-
max_segments=max_segments,
2889-
diag_eps=diag_eps,
2890-
levels=levels,
2891-
block_nnz_use=block_nnz_use,
2892-
max_segments_use=max_segments_use,
2893-
)
2894-
)
2895-
x = torch.stack(cols_out, dim=1)
2896-
if compute_dtype != data.dtype:
2897-
x = x.to(data.dtype)
2898-
if return_time:
2899-
torch.cuda.synchronize()
2900-
elapsed_ms = (time.perf_counter() - t0) * 1000.0
2901-
2902-
if out is not None:
2903-
if out.shape != x.shape or out.dtype != x.dtype:
2904-
raise ValueError("out shape/dtype must match result")
2905-
out.copy_(x)
2906-
x = out
2907-
2908-
if return_time:
2909-
return x, elapsed_ms
2910-
return x
2692+
_validate_spsv_trans_combo(data.dtype, input_index_dtype, "COO")
2693+
data_csr, indices_csr, indptr_csr = _coo_to_csr_sorted_unique(
2694+
data, row64, col64, n_rows, n_cols
2695+
)
2696+
return flagsparse_spsv_csr(
2697+
data_csr,
2698+
indices_csr,
2699+
indptr_csr,
2700+
b,
2701+
shape,
2702+
lower=lower,
2703+
unit_diagonal=unit_diagonal,
2704+
transpose=transpose,
2705+
block_nnz=block_nnz,
2706+
max_segments=max_segments,
2707+
out=out,
2708+
return_time=return_time,
2709+
)

0 commit comments

Comments
 (0)