Skip to content
Closed
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
25 changes: 25 additions & 0 deletions src/core/dbcsr_types.F
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,7 @@ MODULE dbcsr_types

PUBLIC :: dbcsr_type_invalid, dbcsr_type_no_symmetry, dbcsr_type_symmetric, &
dbcsr_type_antisymmetric, dbcsr_type_hermitian, dbcsr_type_antihermitian
PUBLIC :: dbcsr_symmetry_data_type_compatible
PUBLIC :: dbcsr_no_transpose, dbcsr_transpose, dbcsr_conjugate_transpose
PUBLIC :: dbcsr_repl_none, dbcsr_repl_row, dbcsr_repl_col, dbcsr_repl_full

Expand Down Expand Up @@ -588,4 +589,28 @@ MODULE dbcsr_types
INTEGER(KIND=int_8), DIMENSION(SIZE(dbcsr_mpi_size_limits) + 1, 2, 2) :: data_size_breakdown = -1
END TYPE dbcsr_mpi_statistics_type

CONTAINS

PURE FUNCTION dbcsr_symmetry_data_type_compatible(matrix_type, data_type) RESULT(compatible)
!! Whether a matrix symmetry is supported for the given scalar type.

CHARACTER, INTENT(IN) :: matrix_type
INTEGER, INTENT(IN) :: data_type
LOGICAL :: compatible

compatible = .FALSE.
SELECT CASE (data_type)
CASE (dbcsr_type_real_4, dbcsr_type_real_8)
SELECT CASE (matrix_type)
CASE (dbcsr_type_no_symmetry, dbcsr_type_symmetric, dbcsr_type_antisymmetric)
compatible = .TRUE.
END SELECT
CASE (dbcsr_type_complex_4, dbcsr_type_complex_8)
SELECT CASE (matrix_type)
CASE (dbcsr_type_no_symmetry, dbcsr_type_hermitian, dbcsr_type_antihermitian)
compatible = .TRUE.
END SELECT
END SELECT
END FUNCTION dbcsr_symmetry_data_type_compatible

END MODULE dbcsr_types
15 changes: 11 additions & 4 deletions src/dbcsr_api.F
Original file line number Diff line number Diff line change
Expand Up @@ -123,9 +123,10 @@ MODULE dbcsr_api
dbcsr_func_inverse, dbcsr_func_tanh, dbcsr_iterator_prv => dbcsr_iterator, dbcsr_mp_obj, &
dbcsr_no_transpose, dbcsr_norm_column, dbcsr_norm_frobenius, dbcsr_norm_maxabsnorm, &
dbcsr_prv_type => dbcsr_type, dbcsr_scalar_type, dbcsr_transpose, &
dbcsr_type_antisymmetric, dbcsr_type_complex_4, dbcsr_type_complex_8, &
dbcsr_conjugate_transpose, dbcsr_type_antihermitian, dbcsr_type_antisymmetric, &
dbcsr_type_complex_4, dbcsr_type_complex_8, &
dbcsr_type_complex_default, dbcsr_type_no_symmetry, dbcsr_type_real_4, dbcsr_type_real_8, &
dbcsr_type_real_default, dbcsr_type_symmetric
dbcsr_type_real_default, dbcsr_type_symmetric, dbcsr_type_hermitian
USE dbcsr_dist_util, ONLY: dbcsr_convert_offsets_to_sizes => convert_offsets_to_sizes, &
dbcsr_convert_sizes_to_offsets => convert_sizes_to_offsets, &
dbcsr_checksum_prv => dbcsr_checksum, &
Expand All @@ -150,7 +151,10 @@ MODULE dbcsr_api
PUBLIC :: dbcsr_type_no_symmetry
PUBLIC :: dbcsr_type_symmetric
PUBLIC :: dbcsr_type_antisymmetric
PUBLIC :: dbcsr_type_hermitian
PUBLIC :: dbcsr_type_antihermitian
PUBLIC :: dbcsr_transpose
PUBLIC :: dbcsr_conjugate_transpose
PUBLIC :: dbcsr_no_transpose
PUBLIC :: dbcsr_type_complex_8
PUBLIC :: dbcsr_type_real_4
Expand Down Expand Up @@ -1401,11 +1405,14 @@ SUBROUTINE dbcsr_trace_${nametype1}$ (matrix_a, trace)
CALL dbcsr_trace_prv(matrix_a%prv, trace)
END SUBROUTINE dbcsr_trace_${nametype1}$

SUBROUTINE dbcsr_dot_${nametype1}$ (matrix_a, matrix_b, result)
SUBROUTINE dbcsr_dot_${nametype1}$ (matrix_a, matrix_b, result, conjugate_a)
!! Computes sum(A(i,j)*B(i,j)) by default. If conjugate_a is true,
!! computes the Frobenius inner product sum(conjg(A(i,j))*B(i,j)).
TYPE(dbcsr_type), INTENT(IN) :: matrix_a, matrix_b
${type1}$, INTENT(INOUT) :: result
LOGICAL, INTENT(IN), OPTIONAL :: conjugate_a

CALL dbcsr_dot_prv(matrix_a%prv, matrix_b%prv, result)
CALL dbcsr_dot_prv(matrix_a%prv, matrix_b%prv, result, conjugate_a)
END SUBROUTINE dbcsr_dot_${nametype1}$

SUBROUTINE dbcsr_multiply_${nametype1}$ (transa, transb, &
Expand Down
111 changes: 62 additions & 49 deletions src/ops/dbcsr_operations.F
Original file line number Diff line number Diff line change
Expand Up @@ -77,7 +77,7 @@ MODULE dbcsr_operations
dbcsr_norm_maxabsnorm, dbcsr_repl_full, dbcsr_repl_none, dbcsr_scalar_type, dbcsr_type, &
dbcsr_type_antihermitian, dbcsr_type_antisymmetric, dbcsr_type_complex_4, &
dbcsr_type_complex_8, dbcsr_type_hermitian, dbcsr_type_no_symmetry, dbcsr_type_real_4, &
dbcsr_type_real_8, dbcsr_type_symmetric
dbcsr_type_real_8, dbcsr_type_symmetric, dbcsr_symmetry_data_type_compatible
USE dbcsr_dist_util, ONLY: find_block_of_element
USE dbcsr_work_operations, ONLY: dbcsr_create, &
dbcsr_finalize, &
Expand Down Expand Up @@ -612,12 +612,17 @@ SUBROUTINE dbcsr_add_anytype(matrix_a, matrix_b, alpha_scalar, beta_scalar, flop
IF (.NOT. dbcsr_valid_index(matrix_a)) &
DBCSR_ABORT("Invalid matrix")

IF ((dbcsr_get_matrix_type(matrix_a) .EQ. dbcsr_type_symmetric .OR. &
dbcsr_get_matrix_type(matrix_a) .EQ. dbcsr_type_antisymmetric) .NEQV. &
(dbcsr_get_matrix_type(matrix_b) .EQ. dbcsr_type_symmetric .OR. &
dbcsr_get_matrix_type(matrix_b) .EQ. dbcsr_type_antisymmetric)) THEN
IF (dbcsr_has_symmetry(matrix_a) .NEQV. dbcsr_has_symmetry(matrix_b)) THEN
DBCSR_ABORT("Summing general with symmetric matrix NYI")
END IF
IF (dbcsr_has_symmetry(matrix_a) .AND. &
dbcsr_get_matrix_type(matrix_a) .NE. dbcsr_get_matrix_type(matrix_b)) &
DBCSR_ABORT("Summing matrices with different symmetry types NYI")
IF (.NOT. dbcsr_symmetry_data_type_compatible(dbcsr_get_matrix_type(matrix_a), &
dbcsr_get_data_type(matrix_a)) .OR. &
.NOT. dbcsr_symmetry_data_type_compatible(dbcsr_get_matrix_type(matrix_b), &
dbcsr_get_data_type(matrix_b))) &
DBCSR_ABORT("Matrix symmetry is incompatible with its scalar data type.")

data_type_a = dbcsr_get_data_type(matrix_a)
data_type_b = dbcsr_get_data_type(matrix_b)
Expand Down Expand Up @@ -1180,22 +1185,7 @@ LOGICAL FUNCTION symmetry_consistent(matrix_type, data_type)
CHARACTER, INTENT(IN) :: matrix_type
INTEGER, INTENT(IN) :: data_type

symmetry_consistent = .FALSE.

SELECT CASE (data_type)
CASE (dbcsr_type_real_4, dbcsr_type_real_8)
SELECT CASE (matrix_type)
CASE (dbcsr_type_no_symmetry, dbcsr_type_symmetric, dbcsr_type_antisymmetric)
symmetry_consistent = .TRUE.
END SELECT
CASE (dbcsr_type_complex_4, dbcsr_type_complex_8)
SELECT CASE (matrix_type)
CASE (dbcsr_type_no_symmetry, dbcsr_type_hermitian, dbcsr_type_antihermitian)
symmetry_consistent = .TRUE.
END SELECT
CASE DEFAULT
DBCSR_ABORT("Invalid data type.")
END SELECT
symmetry_consistent = dbcsr_symmetry_data_type_compatible(matrix_type, data_type)

END FUNCTION symmetry_consistent

Expand Down Expand Up @@ -1246,7 +1236,7 @@ SUBROUTINE dbcsr_copy(matrix_b, matrix_a, name, keep_sparsity, &
!! when copy from complex to real,& the default is to keep only the real part; if this flag is set, the imaginary part is
!! used
CHARACTER, INTENT(IN), OPTIONAL :: matrix_type
!! 'N' for normal, 'T' for transposed, 'S' for symmetric, and 'A' for antisymmetric
!! matrix type compatible with the scalar data type: N/S/A for real, N/H/K for complex

CHARACTER(len=*), PARAMETER :: routineN = 'dbcsr_copy'
CHARACTER :: new_matrix_type, repl_type
Expand Down Expand Up @@ -1358,6 +1348,14 @@ SUBROUTINE dbcsr_copy_into_existing(matrix_b, matrix_a)
IF (dbcsr_get_data_type(matrix_b) .NE. dbcsr_get_data_type(matrix_a)) &
DBCSR_ABORT("Matrices have different data types.")
data_type = dbcsr_get_data_type(matrix_b)
IF (.NOT. symmetry_consistent(dbcsr_get_matrix_type(matrix_a), data_type) .OR. &
.NOT. symmetry_consistent(dbcsr_get_matrix_type(matrix_b), data_type)) &
DBCSR_ABORT("Matrix symmetry is incompatible with its scalar data type.")
IF (dbcsr_has_symmetry(matrix_a) .AND. dbcsr_has_symmetry(matrix_b) .AND. &
dbcsr_get_matrix_type(matrix_a) .NE. dbcsr_get_matrix_type(matrix_b)) &
DBCSR_ABORT("Copying between different symmetric matrix types NYI.")
IF (.NOT. dbcsr_has_symmetry(matrix_b) .AND. dbcsr_has_symmetry(matrix_a)) &
DBCSR_ABORT("Copying a compressed matrix into an uncompressed target NYI.")
neg_real = matrix_b%negate_real
neg_imag = matrix_b%negate_imaginary
making_symmetric = dbcsr_has_symmetry(matrix_b) &
Expand Down Expand Up @@ -2189,8 +2187,7 @@ FUNCTION dbcsr_gershgorin_norm(matrix) RESULT(norm)
nr = dbcsr_nfullrows_total(matrix)
nc = dbcsr_nfullcols_total(matrix)

any_sym = dbcsr_get_matrix_type(matrix) .EQ. dbcsr_type_symmetric .OR. &
dbcsr_get_matrix_type(matrix) .EQ. dbcsr_type_antisymmetric
any_sym = dbcsr_has_symmetry(matrix)

IF (nr .NE. nc) &
DBCSR_ABORT("not a square matrix")
Expand Down Expand Up @@ -2228,8 +2225,7 @@ FUNCTION dbcsr_gershgorin_norm(matrix) RESULT(norm)
DO i = 1, SIZE(data_c, 1)
buff_d(row_offset + i - 1) = buff_d(row_offset + i - 1) + ABS(data_c(i, j))
IF (any_sym .AND. row .NE. col) &
DBCSR_ABORT("Only nonsymmetric matrix so far")
! buff_d(col_offset+j-1) = buff_d(col_offset+j-1) + ABS(data_c(i,j))
buff_d(col_offset + j - 1) = buff_d(col_offset + j - 1) + ABS(data_c(i, j))
END DO
END DO
CASE (dbcsr_type_complex_8)
Expand All @@ -2239,8 +2235,7 @@ FUNCTION dbcsr_gershgorin_norm(matrix) RESULT(norm)
DO i = 1, SIZE(data_z, 1)
buff_d(row_offset + i - 1) = buff_d(row_offset + i - 1) + ABS(data_z(i, j))
IF (any_sym .AND. row .NE. col) &
DBCSR_ABORT("Only nonsymmetric matrix so far")
! buff_d(col_offset+j-1) = buff_d(col_offset+j-1) + ABS(data_z(i,j))
buff_d(col_offset + j - 1) = buff_d(col_offset + j - 1) + ABS(data_z(i, j))
END DO
END DO
CASE DEFAULT
Expand Down Expand Up @@ -2325,8 +2320,7 @@ FUNCTION dbcsr_frobenius_norm(matrix, local) RESULT(norm)
my_local = .FALSE.
IF (PRESENT(local)) my_local = local

any_sym = dbcsr_get_matrix_type(matrix) .EQ. dbcsr_type_symmetric .OR. &
dbcsr_get_matrix_type(matrix) .EQ. dbcsr_type_antisymmetric
any_sym = dbcsr_has_symmetry(matrix)

norm = 0.0_dp
CALL dbcsr_iterator_start(iter, matrix)
Expand All @@ -2345,14 +2339,12 @@ FUNCTION dbcsr_frobenius_norm(matrix, local) RESULT(norm)
CASE (dbcsr_type_complex_4)
CALL dbcsr_iterator_next_block(iter, row, col, data_c, tr, blk)
fac = 1.0_dp
IF (any_sym .AND. row .NE. col) &
DBCSR_ABORT("Only nonsymmetric matrix so far")
IF (any_sym .AND. row .NE. col) fac = 2.0_dp
norm = norm + fac*REAL(SUM(CONJG(data_c)*data_c), KIND=real_8)
CASE (dbcsr_type_complex_8)
CALL dbcsr_iterator_next_block(iter, row, col, data_z, tr, blk)
fac = 1.0_dp
IF (any_sym .AND. row .NE. col) &
DBCSR_ABORT("Only nonsymmetric matrix so far")
IF (any_sym .AND. row .NE. col) fac = 2.0_dp
norm = norm + fac*REAL(SUM(CONJG(data_z)*data_z), KIND=real_8)
CASE DEFAULT
DBCSR_ABORT("Wrong data type")
Expand Down Expand Up @@ -2591,14 +2583,15 @@ SUBROUTINE dbcsr_trace_sd(matrix_a, trace)
CALL timestop(handle)
END SUBROUTINE dbcsr_trace_sd

SUBROUTINE dbcsr_dot_sd(matrix_a, matrix_b, trace)
SUBROUTINE dbcsr_dot_sd(matrix_a, matrix_b, trace, conjugate_a)
!! Dot product of DBCSR matrices
!! \result the dot product of the matrices

TYPE(dbcsr_type), INTENT(IN) :: matrix_a, matrix_b
!! DBCSR matrices
!! DBCSR matrices
REAL(kind=real_8), INTENT(INOUT) :: trace
LOGICAL, INTENT(IN), OPTIONAL :: conjugate_a

CHARACTER(len=*), PARAMETER :: routineN = 'dbcsr_dot_sd'

Expand All @@ -2608,11 +2601,11 @@ SUBROUTINE dbcsr_dot_sd(matrix_a, matrix_b, trace)
CALL timeset(routineN, handle)
IF (dbcsr_get_data_type(matrix_a) .EQ. dbcsr_type_real_8 .AND. &
dbcsr_get_data_type(matrix_b) .EQ. dbcsr_type_real_8) THEN
CALL dbcsr_dot_d(matrix_a, matrix_b, trace)
CALL dbcsr_dot_d(matrix_a, matrix_b, trace, conjugate_a)
ELSEIF (dbcsr_get_data_type(matrix_a) .EQ. dbcsr_type_real_4 .AND. &
dbcsr_get_data_type(matrix_b) .EQ. dbcsr_type_real_4) THEN
trace_4 = 0.0_real_4
CALL dbcsr_dot_s(matrix_a, matrix_b, trace_4)
CALL dbcsr_dot_s(matrix_a, matrix_b, trace_4, conjugate_a)
trace = REAL(trace_4, real_8)
ELSE
DBCSR_ABORT("Invalid combination of data type, NYI")
Expand Down Expand Up @@ -2689,23 +2682,26 @@ SUBROUTINE dbcsr_trace_${nametype1}$ (matrix_a, trace)
CALL timestop(error_handle)
END SUBROUTINE dbcsr_trace_${nametype1}$

SUBROUTINE dbcsr_dot_${nametype1}$ (matrix_a, matrix_b, trace)
SUBROUTINE dbcsr_dot_${nametype1}$ (matrix_a, matrix_b, trace, conjugate_a)
!! Dot product of DBCSR matrices
!! By default this is the bilinear product sum(A(i,j)*B(i,j)). Set
!! conjugate_a to compute the Frobenius inner product sum(conjg(A(i,j))*B(i,j)).

TYPE(dbcsr_type), INTENT(IN) :: matrix_a, matrix_b
!! DBCSR matrices
!! DBCSR matrices
${type1}$, INTENT(INOUT) :: trace
!! the trace of the product of the matrices
LOGICAL, INTENT(IN), OPTIONAL :: conjugate_a

INTEGER :: a_blk, a_col, a_col_size, a_row_size, b_blk, b_col_size, &
b_frst_blk, b_last_blk, b_row_size, nze, row, a_beg, a_end, b_beg, b_end
CHARACTER :: matrix_a_type, matrix_b_type
INTEGER, DIMENSION(:), POINTER :: a_col_blk_size, &
a_row_blk_size, &
b_col_blk_size, b_row_blk_size
${type1}$ :: sym_fac, fac
LOGICAL :: found, matrix_a_symm, matrix_b_symm
${type1}$ :: dot_block, sym_fac
LOGICAL :: found, matrix_a_symm, matrix_b_symm, my_conjugate_a
${type1}$, DIMENSION(:), POINTER :: a_data, b_data

! ---------------------------------------------------------------------------
Expand All @@ -2714,13 +2710,15 @@ SUBROUTINE dbcsr_dot_${nametype1}$ (matrix_a, matrix_b, trace)
.OR. matrix_b%replication_type .NE. dbcsr_repl_none) &
DBCSR_ABORT("Trace of product of replicated matrices not yet possible.")

sym_fac = REAL(1.0, ${kind1}$)
matrix_a_type = dbcsr_get_matrix_type(matrix_a)
matrix_b_type = dbcsr_get_matrix_type(matrix_b)
matrix_a_symm = matrix_a_type == dbcsr_type_symmetric .OR. matrix_a_type == dbcsr_type_antisymmetric
matrix_b_symm = matrix_b_type == dbcsr_type_symmetric .OR. matrix_b_type == dbcsr_type_antisymmetric

IF (matrix_a_symm .AND. matrix_b_symm) sym_fac = REAL(2.0, ${kind1}$)
matrix_a_symm = matrix_a_type /= dbcsr_type_no_symmetry
matrix_b_symm = matrix_b_type /= dbcsr_type_no_symmetry
sym_fac = REAL(1.0, ${kind1}$)
IF (matrix_a_symm .AND. matrix_b_symm .AND. matrix_a_type /= matrix_b_type) &
sym_fac = REAL(-1.0, ${kind1}$)
my_conjugate_a = .FALSE.
IF (PRESENT(conjugate_a)) my_conjugate_a = conjugate_a

! tracing a symmetric with a general matrix is not implemented, as it would require communication of blocks
IF (matrix_a_symm .NEQV. matrix_b_symm) &
Expand Down Expand Up @@ -2766,10 +2764,25 @@ SUBROUTINE dbcsr_dot_${nametype1}$ (matrix_a, matrix_b, trace)
a_end = a_beg + nze - 1
b_beg = ABS(matrix_b%blk_p(b_blk))
b_end = b_beg + nze - 1
fac = REAL(1.0, ${kind1}$)
IF (row .NE. a_col) fac = sym_fac

trace = trace + fac*SUM(a_data(a_beg:a_end)*b_data(b_beg:b_end))
#:if n >= 2
IF (my_conjugate_a) THEN
dot_block = SUM(CONJG(a_data(a_beg:a_end))*b_data(b_beg:b_end))
ELSE
dot_block = SUM(a_data(a_beg:a_end)*b_data(b_beg:b_end))
END IF
IF (matrix_a_symm .AND. row .NE. a_col) THEN
trace = trace + dot_block + sym_fac*CONJG(dot_block)
ELSE
trace = trace + dot_block
END IF
#:else
dot_block = SUM(a_data(a_beg:a_end)*b_data(b_beg:b_end))
IF (matrix_a_symm .AND. row .NE. a_col) THEN
trace = trace + (1.0_${kind1}$+sym_fac)*dot_block
ELSE
trace = trace + dot_block
END IF
#:endif

END IF
END IF
Expand Down
Loading
Loading