aboutsummaryrefslogtreecommitdiff
path: root/tester/threading.c
blob: 78f788ab64c2de1ec8759ba9c51e738ed163bc72 (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
#include <stdio.h>

#include "tests.h"

static void build_wide(struct PsiQuantumCircuit* c)
{
    for (size_t i = 0; i < 8; i++)
        psi_apply_h(c, i);
    for (size_t i = 0; i + 1 < 8; i++)
        psi_apply_cnot(c, i, i + 1);
    psi_apply_rx(c, 3, 0.7);
    psi_apply_t(c, 5);
    psi_apply_cz(c, 2, 6);
}

void run_threading_tests(void)
{
    psi_test_section("Threading (parallel runtimes)");

    printf("  threads: %zu\n", psi_thread_count());

    struct PsiQuantumCircuit base = psi_new_quantum_circuit(8);
    build_wide(&base);
    struct PsiVector reference =
            psi_clone_vector(*psi_compute_circuit_with(&base, PSI_RUNTIME_BASIC));

    enum PsiRuntime runtimes[] = {
        PSI_RUNTIME_BASIC_MT,
        PSI_RUNTIME_BATCHED_MT,
        PSI_RUNTIME_SIMD_MT,
        PSI_RUNTIME_STRUCTURE_AWARE_MT,
    };
    const char* names[] = {
        "BASIC_MT matches BASIC (8 qubits)",
        "BATCHED_MT matches BASIC (8 qubits)",
        "SIMD_MT matches BASIC (8 qubits)",
        "STRUCTURE_AWARE_MT matches BASIC (8 qubits)",
    };

    for (size_t i = 0; i < sizeof runtimes / sizeof runtimes[0]; i++)
    {
        struct PsiQuantumCircuit c = psi_new_quantum_circuit(8);
        build_wide(&c);
        const struct PsiVector* state = psi_compute_circuit_with(&c, runtimes[i]);
        bool ok = state->size == reference.size &&
                psi_amps_match(state->data, reference.data, reference.size);
        psi_test_check(ok, names[i]);
        psi_free_quantum_circuit(&c);
    }

    psi_free_vector(&reference);
    psi_free_quantum_circuit(&base);
}