#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; }