Skip to content

Commit 8d38162

Browse files
authored
recover changes (fortran-lang#1117)
1 parent 9591c0d commit 8d38162

6 files changed

Lines changed: 258 additions & 4 deletions

File tree

doc/specs/stdlib_sparse.md

Lines changed: 37 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -338,7 +338,7 @@ If the `diagonal` array has not been previously allocated, the `diag` subroutine
338338

339339
### Syntax
340340

341-
`call ` [[stdlib_sparse_conversion(module):csr2sellc(interface)]] `(csr,ell[,num_nz_rows])`
341+
`call ` [[stdlib_sparse_conversion(module):csr2ell(interface)]] `(csr,ell[,num_nz_rows])`
342342

343343
### Arguments
344344

@@ -361,4 +361,39 @@ If the `diagonal` array has not been previously allocated, the `diag` subroutine
361361
### Example
362362
```fortran
363363
{!example/linalg/example_sparse_spmv.f90!}
364-
```
364+
```
365+
366+
<!-- -- -- -- -- -- -- -- -- -- -- -- -- -- -- -- -- -- -- -->
367+
## Operator overloading (`+`, `-`, `*`, `/`) {#operators}
368+
369+
### Status
370+
371+
Experimental
372+
373+
### Description
374+
375+
The definition of all standard arithmetic operators have been overloaded to be applicable for the matrix types defined by `stdlib_sparse`. The operators have been overloaded to support the following tuple combinations of the left-hand-side of the operation: same type and kind matrix-matrix, matrix-scalar and scalar-matrix.
376+
377+
### Syntax
378+
379+
- Matrix-matrix operators :
380+
381+
`C = A + B`
382+
383+
`C = A - B`
384+
385+
`C = A * B`
386+
387+
`C = A / B`
388+
389+
- Matrix scalar operators :
390+
391+
`B = A + alpha` or `B = alpha + A`
392+
393+
`B = A - alpha` or `B = alpha - A`
394+
395+
`B = A * alpha` or `B = alpha * A`
396+
397+
`B = A / alpha` or `B = alpha / A`
398+
399+
*Note*: scalar addition and subtraction operators perform element-wise operations only on the stored (non-zero) values, not on the full mathematical matrix. Meaning, the sparsity pattern is preserved.

src/sparse/CMakeLists.txt

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ set(sparse_fppFiles
22
stdlib_sparse_constants.fypp
33
stdlib_sparse_conversion.fypp
44
stdlib_sparse_kinds.fypp
5+
stdlib_sparse_operators.fypp
56
stdlib_sparse_spmv.fypp
67
)
78

src/sparse/stdlib_sparse_conversion.fypp

Lines changed: 37 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -161,6 +161,7 @@ module stdlib_sparse_conversion
161161
#:for k1, t1, s1 in (KINDS_TYPES)
162162
module procedure :: coo_from_ijv_${s1}$
163163
module procedure :: csr_from_ijv_${s1}$
164+
module procedure :: csc_from_ijv_${s1}$
164165
module procedure :: ell_from_ijv_${s1}$
165166
module procedure :: sellc_from_ijv_${s1}$
166167
#:endfor
@@ -258,8 +259,9 @@ contains
258259
temp(1,1:COO%nnz) = COO%index(2,1:COO%nnz)
259260
temp(2,1:COO%nnz) = COO%index(1,1:COO%nnz)
260261
allocate(data, source = COO%data )
262+
261263
nnz = COO%nnz
262-
call sort_coo_unique_${s1}$( temp, data, nnz, COO%nrows, COO%ncols )
264+
call sort_coo_unique_${s1}$( temp, data, nnz, COO%ncols, COO%nrows )
263265

264266
if( allocated(CSC%row) ) then
265267
CSC%row(1:COO%nnz) = temp(2,1:COO%nnz)
@@ -740,6 +742,40 @@ contains
740742
end subroutine
741743
#:endfor
742744

745+
#:for k1, t1, s1 in (KINDS_TYPES)
746+
subroutine csc_from_ijv_${s1}$(CSC,row,col,data,nrows,ncols)
747+
type(CSC_${s1}$_type), intent(inout) :: CSC
748+
integer(ilp), intent(in) :: row(:)
749+
integer(ilp), intent(in) :: col(:)
750+
${t1}$, intent(in), optional :: data(:)
751+
integer(ilp), intent(in), optional :: nrows
752+
integer(ilp), intent(in), optional :: ncols
753+
754+
integer(ilp) :: nrows_, ncols_
755+
!---------------------------------------------------------
756+
if(present(nrows)) then
757+
nrows_ = nrows
758+
else
759+
nrows_ = maxval(row)
760+
end if
761+
if(present(ncols)) then
762+
ncols_ = ncols
763+
else
764+
ncols_ = maxval(col)
765+
end if
766+
!---------------------------------------------------------
767+
block
768+
type(COO_${s1}$_type) :: COO
769+
if(present(data)) then
770+
call from_ijv(COO,row,col,data=data,nrows=nrows_,ncols=ncols_)
771+
else
772+
call from_ijv(COO,row,col,nrows=nrows_,ncols=ncols_)
773+
end if
774+
call coo2csc(COO,CSC)
775+
end block
776+
end subroutine
777+
#:endfor
778+
743779
#:for k1, t1, s1 in (KINDS_TYPES)
744780
subroutine ell_from_ijv_${s1}$(ELL,row,col,data,nrows,ncols,num_nz_rows)
745781
type(ELL_${s1}$_type), intent(inout) :: ELL

src/sparse/stdlib_sparse_kinds.fypp

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,10 @@
33
#:set R_KINDS_TYPES = list(zip(REAL_KINDS, REAL_TYPES, REAL_SUFFIX))
44
#:set C_KINDS_TYPES = list(zip(CMPLX_KINDS, CMPLX_TYPES, CMPLX_SUFFIX))
55
#:set KINDS_TYPES = R_KINDS_TYPES+C_KINDS_TYPES
6+
#:set OP_NAMES = ["add","sub","mul","div"]
7+
#:set OP_SYMBOLS = ["+","-","*","/"]
8+
#:set OPERATORS = list(zip(OP_NAMES, OP_SYMBOLS))
9+
610
!! The `stdlib_sparse_kinds` module provides derived type definitions for different sparse matrices
711
!!
812
! This code was modified from https://github.com/jalvesz/FSPARSE by its author: Alves Jose
@@ -13,6 +17,8 @@ module stdlib_sparse_kinds
1317
private
1418
public :: sparse_full, sparse_lower, sparse_upper
1519
public :: sparse_op_none, sparse_op_transpose, sparse_op_hermitian
20+
public :: operator(+), operator(-), operator(*), operator(/)
21+
1622
!! version: experimental
1723
!!
1824
!! Base sparse type holding the meta data related to the storage capacity of a matrix.
@@ -128,6 +134,34 @@ module stdlib_sparse_kinds
128134
end type
129135
#:endfor
130136

137+
#:for op, sym in OPERATORS
138+
!! Overload the `${sym}$` operator for sparse matrices
139+
!! [Specifications](../page/specs/stdlib_sparse.html#operators)
140+
interface operator(${sym}$)
141+
#:for matrix in SPARSE_KINDS
142+
#:for k, t, s in (KINDS_TYPES)
143+
pure module function sparse_${op}$_${matrix}$_${s}$(a, b) result(c)
144+
type(${matrix}$_${s}$_type), intent(in) :: a, b
145+
type(${matrix}$_${s}$_type) :: c
146+
end function
147+
148+
pure module function sparse_${op}$_${matrix}$_scalar_${s}$(a, b) result(c)
149+
${t}$, intent(in) :: a
150+
type(${matrix}$_${s}$_type), intent(in) :: b
151+
type(${matrix}$_${s}$_type) :: c
152+
end function
153+
154+
pure module function sparse_${op}$_scalar_${matrix}$_${s}$(a, b) result(c)
155+
type(${matrix}$_${s}$_type), intent(in) :: a
156+
${t}$, intent(in) :: b
157+
type(${matrix}$_${s}$_type) :: c
158+
end function
159+
#:endfor
160+
#:endfor
161+
end interface
162+
163+
#:endfor
164+
131165
contains
132166

133167
!! (re)Allocate matrix memory for the COO type
Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
#:include "common.fypp"
2+
#:set R_KINDS_TYPES = list(zip(REAL_KINDS, REAL_TYPES, REAL_SUFFIX))
3+
#:set C_KINDS_TYPES = list(zip(CMPLX_KINDS, CMPLX_TYPES, CMPLX_SUFFIX))
4+
#:set KINDS_TYPES = R_KINDS_TYPES+C_KINDS_TYPES
5+
#:set OP_NAMES = ["add","sub","mul","div"]
6+
#:set OP_SYMBOLS = ["+","-","*","/"]
7+
#:set OPERATORS = list(zip(OP_NAMES, OP_SYMBOLS))
8+
9+
submodule(stdlib_sparse_kinds) stdlib_sparse_operators
10+
implicit none
11+
12+
contains
13+
14+
#:for op, sym in OPERATORS
15+
#:for matrix in SPARSE_KINDS
16+
#:for k, t, s in (KINDS_TYPES)
17+
pure module function sparse_${op}$_${matrix}$_${s}$(a, b) result(c)
18+
type(${matrix}$_${s}$_type), intent(in) :: a, b
19+
type(${matrix}$_${s}$_type) :: c
20+
c = a
21+
c%data = c%data ${sym}$ b%data
22+
end function
23+
24+
pure module function sparse_${op}$_${matrix}$_scalar_${s}$(a, b) result(c)
25+
${t}$, intent(in) :: a
26+
type(${matrix}$_${s}$_type), intent(in) :: b
27+
type(${matrix}$_${s}$_type) :: c
28+
c = b
29+
c%data = a ${sym}$ c%data
30+
end function
31+
32+
pure module function sparse_${op}$_scalar_${matrix}$_${s}$(a, b) result(c)
33+
type(${matrix}$_${s}$_type), intent(in) :: a
34+
${t}$, intent(in) :: b
35+
type(${matrix}$_${s}$_type) :: c
36+
c = a
37+
c%data = c%data ${sym}$ b
38+
end function
39+
40+
#:endfor
41+
#:endfor
42+
#:endfor
43+
44+
end submodule stdlib_sparse_operators

test/linalg/test_linalg_sparse.fypp

Lines changed: 105 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,9 @@
11
#:include "common.fypp"
22
#:set R_KINDS_TYPES = list(zip(REAL_KINDS, REAL_TYPES, REAL_SUFFIX))
33
#:set KINDS_TYPES = R_KINDS_TYPES
4+
#:set OP_NAMES = ["add","sub","mul","div"]
5+
#:set OP_SYMBOLS = ["+","-","*","/"]
6+
#:set OPERATORS = list(zip(OP_NAMES, OP_SYMBOLS))
47
module test_sparse_spmv
58
use testdrive, only : new_unittest, unittest_type, error_type, check, skip_test
69
use stdlib_kinds, only: sp, dp, xdp, qp, int8, int16, int32, int64
@@ -25,7 +28,8 @@ contains
2528
new_unittest('sellc', test_sellc), &
2629
new_unittest('symmetries', test_symmetries), &
2730
new_unittest('diagonal', test_diagonal), &
28-
new_unittest('add_get_values', test_add_get_values) &
31+
new_unittest('add_get_values', test_add_get_values), &
32+
new_unittest('sparse_operators', test_sparse_operators) &
2933
]
3034
end subroutine
3135

@@ -332,6 +336,7 @@ contains
332336
diagonal = 0.0
333337
call coo2csc( COO, CSC )
334338
call diag( CSC , diagonal )
339+
print *, 'diagonal csc:', diagonal
335340
call check(error, all(diagonal == [1,2,3,4]) )
336341
if (allocated(error)) return
337342
end block
@@ -383,6 +388,105 @@ contains
383388
#:endfor
384389
end subroutine
385390

391+
subroutine test_sparse_operators(error)
392+
!> Error handling
393+
type(error_type), allocatable, intent(out) :: error
394+
#:for matrix in ["COO","CSR","CSC"] # ELL and SELLC are supported but since they have explicit zeros in data array, the tests are less straightforward
395+
#:for k, t, s in (KINDS_TYPES)
396+
block
397+
integer, parameter :: wp = ${k}$
398+
real(wp), parameter :: tol = 1000*epsilon(0._wp)
399+
real(wp) :: scalar
400+
integer :: row(10), col(10)
401+
real(wp) :: data(10)
402+
type(${matrix}$_${s}$_type) :: a, b, c
403+
${t}$:: err
404+
405+
data(:) = real([9,-3,4,7,8,-1,8,4,5,6],kind=wp)
406+
col(:) = [1,5,1,2,2,3,4,1,3,4]
407+
row(:) = [1,1,2,2,3,3,3,4,4,4]
408+
409+
call from_ijv(a, row, col, data)
410+
call from_ijv(b, row, col, data)
411+
scalar = 2.0_wp
412+
413+
! Test sparse + sparse
414+
c = a + b
415+
err = sum( abs( c%data - 2.0_wp*a%data ) ) / size(data)
416+
call check(error, err <= tol, "error in sparse + sparse" )
417+
if (allocated(error)) return
418+
419+
! Test scalar + sparse
420+
c = scalar + a
421+
err = sum( abs( c%data - (scalar + a%data) ) ) / size(data)
422+
call check(error, err <= tol, "error in scalar + sparse" )
423+
if (allocated(error)) return
424+
425+
! Test sparse + scalar
426+
c = a + scalar
427+
err = sum( abs( c%data - (a%data + scalar) ) ) / size(data)
428+
call check(error, err <= tol, "error in sparse + scalar" )
429+
if (allocated(error)) return
430+
431+
! Test sparse * sparse
432+
c = a * b
433+
err = sum( abs( c%data - a%data*b%data ) ) / size(data)
434+
call check(error, err <= tol, "error in sparse * sparse" )
435+
if (allocated(error)) return
436+
437+
! Test scalar * sparse
438+
c = scalar * a
439+
err = sum( abs( c%data - (scalar*a%data) ) ) / size(data)
440+
call check(error, err <= tol, "error in scalar * sparse" )
441+
if (allocated(error)) return
442+
443+
! Test sparse * scalar
444+
c = a * scalar
445+
err = sum( abs( c%data - (a%data*scalar) ) ) / size(data)
446+
call check(error, err <= tol, "error in sparse * scalar" )
447+
if (allocated(error)) return
448+
449+
! Test sparse - sparse
450+
c = a - b
451+
err = sum( abs( c%data ) ) / size(data)
452+
call check(error, err <= tol, "error in sparse - sparse" )
453+
if (allocated(error)) return
454+
455+
! Test scalar - sparse
456+
c = scalar - a
457+
err = sum( abs( c%data - (scalar - a%data) ) ) / size(data)
458+
call check(error, err <= tol, "error in scalar - sparse" )
459+
if (allocated(error)) return
460+
461+
! Test sparse - scalar
462+
c = a - scalar
463+
err = sum( abs( c%data - (a%data - scalar) ) ) / size(data)
464+
call check(error, err <= tol, "error in sparse - scalar" )
465+
if (allocated(error)) return
466+
467+
! Test sparse / sparse
468+
c = a / b
469+
err = sum( abs( c%data - 1.0_wp ) ) / size(data)
470+
call check(error, err <= tol, "error in sparse / sparse" )
471+
if (allocated(error)) return
472+
473+
! Test scalar / sparse
474+
c = scalar / a
475+
err = sum( abs( c%data - (scalar / a%data) ) ) / size(data)
476+
call check(error, err <= tol, "error in scalar / sparse" )
477+
if (allocated(error)) return
478+
479+
! Test sparse / scalar
480+
c = a / scalar
481+
err = sum( abs( c%data - (a%data / scalar) ) ) / size(data)
482+
call check(error, err <= tol, "error in sparse / scalar" )
483+
if (allocated(error)) return
484+
485+
end block
486+
#:endfor
487+
#:endfor
488+
end subroutine
489+
386490
end module
387491

388492

0 commit comments

Comments
 (0)