diff options
Diffstat (limited to 'tester/test.c')
| -rw-r--r-- | tester/test.c | 81 |
1 files changed, 81 insertions, 0 deletions
diff --git a/tester/test.c b/tester/test.c new file mode 100644 index 0000000..9e11d8c --- /dev/null +++ b/tester/test.c @@ -0,0 +1,81 @@ +#include "tests.h" + +#include <math.h> +#include <stdio.h> + +static int tests_run = 0; +static int tests_passed = 0; + +void psi_test_section(const char* name) +{ + printf("\nāā %s āā\n", name); +} + +void psi_test_check(bool ok, const char* name) +{ + tests_run++; + if (ok) + tests_passed++; + + printf(" [%s] %s\n", ok ? "PASS" : "FAIL", name); +} + +bool psi_amps_match(const struct PsiComplex* actual, const struct PsiComplex* expected, size_t n) +{ + for (size_t i = 0; i < n; i++) + { + if (fabs(actual[i].real - expected[i].real) > 1e-9) + return false; + if (fabs(actual[i].imaginary - expected[i].imaginary) > 1e-9) + return false; + } + + return true; +} + +void psi_check_circuit(const char* name, struct PsiQuantumCircuit* circuit, + const struct PsiComplex* expected, size_t n) +{ + const struct PsiVector* state = psi_compute_circuit(circuit); + bool ok = state->size == n && psi_amps_match(state->data, expected, n); + psi_test_check(ok, name); +} + +void psi_check_runtimes_agree(const char* name, size_t num_qubits, + void (*build)(struct PsiQuantumCircuit*)) +{ + enum PsiRuntime runtimes[] = { + PSI_RUNTIME_BASIC, + PSI_RUNTIME_BATCHED, + PSI_RUNTIME_SIMD, + PSI_RUNTIME_STRUCTURE_AWARE, + }; + + struct PsiQuantumCircuit base = psi_new_quantum_circuit(num_qubits); + build(&base); + const struct PsiVector* base_state = psi_compute_circuit_with(&base, PSI_RUNTIME_BASIC); + struct PsiVector reference = psi_clone_vector(*base_state); + + bool ok = true; + for (size_t i = 1; i < sizeof runtimes / sizeof runtimes[0]; i++) + { + struct PsiQuantumCircuit circuit = psi_new_quantum_circuit(num_qubits); + build(&circuit); + const struct PsiVector* state = psi_compute_circuit_with(&circuit, runtimes[i]); + if (state->size != reference.size || + !psi_amps_match(state->data, reference.data, reference.size)) + ok = false; + + psi_free_quantum_circuit(&circuit); + } + + psi_free_vector(&reference); + psi_free_quantum_circuit(&base); + psi_test_check(ok, name); +} + +int psi_test_summary(void) +{ + printf("\n%d/%d checks passed\n", tests_passed, tests_run); + return tests_passed == tests_run ? 0 : 1; +} |
