#version 450
// Save the trailing conv window after a full-sequence causal conv (prefill),
// mirroring cuda/gdn.cu save_conv_state_kernel. One invocation per (channel,
// batch): copy the last kernel_size input samples into conv_state, zero-padding
// on the left when seq_len < kernel_size.
//
// x: [B, conv_dim, S] conv_state_out: [B, conv_dim, kernel_size]
layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in;
layout(set = 0, binding = 0) readonly buffer X { float x[]; };
layout(set = 0, binding = 1) writeonly buffer Cs { float cs[]; };
layout(push_constant) uniform Pc { uint batch_size; uint conv_dim; uint seq_len; uint kernel_size; };
void main() {
uint ch = gl_GlobalInvocationID.x;
uint b = gl_GlobalInvocationID.y;
if (ch >= conv_dim || b >= batch_size) { return; }
uint x_base = (b * conv_dim + ch) * seq_len;
uint cs_base = (b * conv_dim + ch) * kernel_size;
int pad = int(kernel_size) - int(seq_len);
for (uint i = 0u; i < kernel_size; i++) {
if (int(i) < pad) {
cs[cs_base + i] = 0.0;
} else {
cs[cs_base + i] = x[x_base + (seq_len - kernel_size + i)];
}
}
}