Spaces:
Running
Running
| 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; | |
| } | |
| } | |