tools: add iris.c
This commit is contained in:
@@ -0,0 +1,29 @@
|
||||
#version 450
|
||||
/*
|
||||
* Split a fused projection row into 2 or 3 equal-width contiguous streams.
|
||||
* fused: [seq, fused_dim], row laid out as [stream0 | stream1 | (stream2)]
|
||||
* a[s,e] = row[0*width + e], b[s,e] = row[1*width + e], c = row[2*width + e]
|
||||
* Used for fused QKV (n_streams=3, width=hidden) and fused gate/up
|
||||
* (n_streams=2, width=mlp_hidden). When n_streams==2, binding `c` is unused.
|
||||
* One invocation per (s, e).
|
||||
*/
|
||||
layout(local_size_x = 16, local_size_y = 16) in;
|
||||
|
||||
layout(std430, binding = 0) readonly buffer F { float fused[]; };
|
||||
layout(std430, binding = 1) writeonly buffer A { float a[]; };
|
||||
layout(std430, binding = 2) writeonly buffer B { float b[]; };
|
||||
layout(std430, binding = 3) buffer C { float c[]; };
|
||||
|
||||
layout(push_constant) uniform PC { uint seq, width, fused_dim, n_streams; } pc;
|
||||
|
||||
void main() {
|
||||
uint e = gl_GlobalInvocationID.x;
|
||||
uint s = gl_GlobalInvocationID.y;
|
||||
if (s >= pc.seq || e >= pc.width) return;
|
||||
|
||||
uint row = s * pc.fused_dim;
|
||||
uint dst = s * pc.width + e;
|
||||
a[dst] = fused[row + e];
|
||||
b[dst] = fused[row + pc.width + e];
|
||||
if (pc.n_streams == 3u) c[dst] = fused[row + 2u * pc.width + e];
|
||||
}
|
||||
Reference in New Issue
Block a user