File size: 1,588 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
#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) buffer Weights {
    float weights[];
};

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

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

layout(std430, binding = 5) readonly buffer PlasticityParams {
    int num_synapses;
    int num_neurons;
    float learning_rate;
    float reward;
    float weight_decay;
    float min_weight;
    float max_weight;
} params;

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

    int pre_idx = col_indices[k];
    float a_pre = pre_activations[pre_idx];

    // Postsynaptic neuron = CSR row owning synapse k (binary search).
    int lo = 0;
    int hi = params.num_neurons;
    while (lo < hi) {
        int mid = (lo + hi) / 2;
        if (row_offsets[mid + 1] <= int(k)) {
            lo = mid + 1;
        } else {
            hi = mid;
        }
    }
    float a_post = post_activations[lo];

    // Three-factor rule: pre * post * reward - decay (matches CPU reference).
    float w = weights[k];
    float delta_w = params.learning_rate * params.reward * (a_pre * a_post - params.weight_decay * w);

    float new_w = clamp(w + delta_w, params.min_weight, params.max_weight);
    weights[k] = new_w;
}