FlyBrain-Lab / shaders /plasticity.comp
timfromhcs's picture
FlyBrain v4.1.0 Space build (REAL_SUBGRAPH, CPU-only, honest backend)
3d46076 verified
Raw History Blame Contribute Delete
1.59 kB
#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;
}