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
4 changes: 2 additions & 2 deletions tests/03_NAO_multik/CASES_GPU.txt
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@ scf_out_hsk
scf_out_hsk_binary
scf_out_hsr
scf_out_hsr_binary_spin2
#scf_out_hsr_spin4
scf_out_hsr_spin4
scf_out_dh_t
scf_out_dos_spin4
scf_out_mul
Expand All @@ -46,7 +46,7 @@ nscf_out_dos
nscf_out_band_pband
nscf_out_pot1
nscf_out_mul
#nscf_out_hsr_tr_rr
nscf_out_hsr_tr_rr
relax_bfgs2
relax_old_cg
relax_cell
Expand Down
7 changes: 4 additions & 3 deletions tests/integrate/tools/catch_properties.sh
Original file line number Diff line number Diff line change
Expand Up @@ -445,13 +445,14 @@ fi
#-----------------------------------
#echo $has_hs2
if ! test -z "$has_hs2" && [ $has_hs2 == 1 ]; then
python3 $COMPARE_SCRIPT hrs1_nao.csr.ref OUT.autotest/hrs1_nao.csr 8
HSR_CSR_COMPARE="../../integrate/tools/compare_hsr_csr.py"
python3 $HSR_CSR_COMPARE hrs1_nao.csr.ref OUT.autotest/hrs1_nao.csr 8
echo "CompareHR_pass $?" >>$1
if ! test -z "$nspin" && [ "$nspin" -eq 2 ]; then
python3 $COMPARE_SCRIPT hrs2_nao.csr.ref OUT.autotest/hrs2_nao.csr 8
python3 $HSR_CSR_COMPARE hrs2_nao.csr.ref OUT.autotest/hrs2_nao.csr 8
echo "CompareHR2_pass $?" >>$1
fi
python3 $COMPARE_SCRIPT sr_nao.csr.ref OUT.autotest/sr_nao.csr 8
python3 $HSR_CSR_COMPARE sr_nao.csr.ref OUT.autotest/sr_nao.csr 8
echo "CompareSR_pass $?" >>$1
elif ! test -z "$has_hs2" && [ "$has_hs2" == 2 ]; then
HSR_BINARY_COMPARE="../../integrate/tools/compare_hsr_binary.py"
Expand Down
222 changes: 222 additions & 0 deletions tests/integrate/tools/compare_hsr_csr.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,222 @@
#!/usr/bin/env python3

import math
import sys

from compare_hsr_binary import read_text_reference


def csr_to_entries(nbasis, values, columns, row_ptr, rvec):
if len(columns) != len(values):
raise ValueError(
"CSR column/value count differs for R = {}: columns={} values={}".format(
rvec, len(columns), len(values)
)
)

if len(row_ptr) != nbasis + 1:
raise ValueError(
"CSR row pointer count differs for R = {}: expected={} actual={}".format(
rvec, nbasis + 1, len(row_ptr)
)
)

if row_ptr[0] != 0 or row_ptr[-1] != len(values):
raise ValueError(
"invalid CSR row pointers for R = {}".format(rvec)
)

if any(ptr < 0 or ptr > len(values) for ptr in row_ptr):
raise ValueError(
"CSR row pointer out of range for R = {}".format(rvec)
)

if any(lhs > rhs for lhs, rhs in zip(row_ptr, row_ptr[1:])):
raise ValueError(
"non-monotonic CSR row pointers for R = {}".format(rvec)
)

entries = {}

for row in range(nbasis):
begin = row_ptr[row]
end = row_ptr[row + 1]

for index in range(begin, end):
col = columns[index]

if col < 0 or col >= nbasis:
raise ValueError(
"CSR column {} out of range for R = {}".format(col, rvec)
)

key = (row, col)

if key in entries:
raise ValueError(
"duplicate CSR entry {} for R = {}".format(key, rvec)
)

entries[key] = values[index]

return entries


def blocks_to_map(nbasis, blocks):
result = {}

for rvec, values, columns, row_ptr in blocks:
if rvec in result:
raise ValueError("duplicate R block: {}".format(rvec))

result[rvec] = csr_to_entries(
nbasis, values, columns, row_ptr, rvec
)

return result


def read_pair(reference_filename, actual_filename):
errors = []

# Prefer real because the common nspin=1/2 H(R), S(R) outputs are real.
# Fall back to complex for spinor matrices.
for value_type in ("real", "complex"):
try:
reference_nbasis, reference_blocks = read_text_reference(
reference_filename, value_type
)
actual_nbasis, actual_blocks = read_text_reference(
actual_filename, value_type
)

return (
value_type,
reference_nbasis,
reference_blocks,
actual_nbasis,
actual_blocks,
)
except (OSError, ValueError) as error:
errors.append("{}: {}".format(value_type, error))

raise ValueError(
"failed to parse files as real or complex CSR: {}".format(
"; ".join(errors)
)
)


def compare(reference_filename, actual_filename, tolerance):
(
value_type,
reference_nbasis,
reference_blocks,
actual_nbasis,
actual_blocks,
) = read_pair(reference_filename, actual_filename)

if reference_nbasis != actual_nbasis:
raise ValueError(
"matrix dimension differs: reference={} actual={}".format(
reference_nbasis, actual_nbasis
)
)

reference = blocks_to_map(reference_nbasis, reference_blocks)
actual = blocks_to_map(actual_nbasis, actual_blocks)

max_difference = 0.0
max_location = None

# Missing R blocks and missing CSR entries are mathematically zero.
for rvec in sorted(set(reference) | set(actual)):
reference_entries = reference.get(rvec, {})
actual_entries = actual.get(rvec, {})

for row_col in sorted(set(reference_entries) | set(actual_entries)):
expected = reference_entries.get(row_col, 0.0)
calculated = actual_entries.get(row_col, 0.0)

if not math.isfinite(abs(expected)) or not math.isfinite(abs(calculated)):
row, col = row_col
raise ValueError(
"non-finite matrix value at R={}, row={}, col={}: "
"reference={} actual={}".format(
rvec, row, col, expected, calculated
)
)

difference = abs(calculated - expected)

if difference > max_difference:
max_difference = difference
max_location = (rvec, row_col, expected, calculated)

if difference > tolerance:
row, col = row_col
raise ValueError(
"matrix value differs at R={}, row={}, col={}: "
"reference={} actual={} difference={} tolerance={}".format(
rvec,
row,
col,
expected,
calculated,
difference,
tolerance,
)
)

if max_location is None:
print(
"H(R)/S(R) CSR comparison passed: "
"type={} max_abs_difference=0".format(value_type)
)
else:
rvec, (row, col), expected, calculated = max_location
print(
"H(R)/S(R) CSR comparison passed: "
"type={} max_abs_difference={} "
"at R={}, row={}, col={}, reference={}, actual={}".format(
value_type,
max_difference,
rvec,
row,
col,
expected,
calculated,
)
)


def main():
if len(sys.argv) != 4:
print(
"usage: compare_hsr_csr.py "
"REFERENCE_CSR ACTUAL_CSR ACCURACY"
)
return 2

reference_filename = sys.argv[1]
actual_filename = sys.argv[2]

try:
tolerance = 10.0 ** (-int(sys.argv[3]))

compare(
reference_filename,
actual_filename,
tolerance,
)
except (OSError, ValueError) as error:
print(
"failed to compare H(R)/S(R) CSR output: {}".format(error)
)
return 1

return 0


if __name__ == "__main__":
sys.exit(main())
Loading