File size: 2,512 Bytes
3d46076
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
82
83
84
85
86
87
88
89
90
91
92
93
94
95
#version 450

layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in;

layout(std430, binding = 0) readonly buffer RowOffsets {
    int row_offsets[];
};

layout(std430, binding = 1) readonly buffer ColIndices {
    int col_indices[];
};

layout(std430, binding = 2) readonly buffer Weights {
    float weights[];
};

layout(std430, binding = 3) readonly buffer PrevSpikes {
    float prev_spikes[];
};

layout(std430, binding = 4) readonly buffer ExternalInputs {
    float external_inputs[];
};

layout(std430, binding = 5) readonly buffer PotentialsIn {
    float potentials_in[];
};

layout(std430, binding = 6) readonly buffer RefractoryIn {
    int refractory_in[];
};

layout(std430, binding = 7) writeonly buffer PotentialsOut {
    float potentials_out[];
};

layout(std430, binding = 8) writeonly buffer SpikesOut {
    float spikes_out[];
};

layout(std430, binding = 9) writeonly buffer RefractoryOut {
    int refractory_out[];
};

layout(std430, binding = 10) readonly buffer Params {
    int num_neurons;
    float decay;
    float threshold;
    float v_reset;
    float v_rest;
    int t_ref;
} params;

void main() {
    uint i = gl_GlobalInvocationID.x;
    if (i >= uint(params.num_neurons)) {
        return;
    }

    int ref_count = refractory_in[i];
    if (ref_count > 0) {
        // Absolute refractory period: clamp to reset potential, suppress spike
        refractory_out[i] = ref_count - 1;
        potentials_out[i] = params.v_reset;
        spikes_out[i] = 0.0;
        return;
    }

    // Accumulate synaptic current from presynaptic spikes
    int start_idx = row_offsets[i];
    int end_idx = row_offsets[i + 1];

    float synaptic_sum = 0.0;
    for (int k = start_idx; k < end_idx; ++k) {
        int pre_idx = col_indices[k];
        synaptic_sum += weights[k] * prev_spikes[pre_idx];
    }

    float v_old = potentials_in[i];
    // Leaky integration towards resting potential
    float v_cand = params.v_rest + (v_old - params.v_rest) * params.decay + synaptic_sum + external_inputs[i];

    if (v_cand >= params.threshold) {
        // Threshold crossed: fire action potential (spike), reset membrane, initiate refractory period
        spikes_out[i] = 1.0;
        potentials_out[i] = params.v_reset;
        refractory_out[i] = params.t_ref;
    } else {
        // Subthreshold: decay potential, no spike, clamp to floor
        spikes_out[i] = 0.0;
        potentials_out[i] = max(v_cand, params.v_reset - 1.0);
        refractory_out[i] = 0;
    }
}