@@ -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-
14861435def _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-
21732117def _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-
22562153def flagsparse_spsv_csr (
22572154 data ,
22582155 indices ,
@@ -2768,7 +2665,6 @@ def _analyze_spsv_csr(
27682665
27692666
27702667def 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