aboutsummaryrefslogtreecommitdiff
path: root/tester/test.c
blob: 2c3532f5bb9f6dfacbaaac32e4490b7eecf5e50b (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
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;
}