Spaces:
Running
Running
Download shaders/plasticity.comp from timfromhcs/FlyBrain-Lab: direct link, hf CLI and curl.
- Browser
- Download file 1.59 kB
-
https://huggingface.co/spaces/timfromhcs/FlyBrain-Lab/resolve/main/shaders/plasticity.comp
- Command line
-
hf download hf://spaces/timfromhcs/FlyBrain-Lab/shaders/plasticity.comp
-
curl -L -o plasticity.comp https://huggingface.co/spaces/timfromhcs/FlyBrain-Lab/resolve/main/shaders/plasticity.comp
1.59 kB
| 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; | |
| } | |