* vulkan: fuse snake activation (mul, sin, sqr, mul, add) Add snake.comp shader with F32 / F16 / BF16 pipelines and ggml_vk_snake_dispatch_fused. The matcher recognizes the naive 5 op decomposition emitted by audio decoders (BigVGAN, Vocos) for snake activation y = x + sin(a*x)^2 * inv_b and rewrites it to a single elementwise kernel. test_snake_fuse from the CUDA PR now also compares CPU naive vs Vulkan fused across F32 / F16 / BF16. * vulkan: address jeffbolznv review for fused snake activation Rename T / C to ne0 / ne1 in the shader and push constants to match the standard naming convention used across the Vulkan backend. Tighten ggml_vk_can_fuse_snake: require x and dst to be contiguous (the shader uses idx = i0 + i1 * ne0) and require a / inv_b to be tightly packed on the broadcast dim (the shader reads data_a[i1]). * vulkan: tighten snake fusion type checks for all operands (address jeffbolznv review) * vulkan: reject snake fusion when ne[2] or ne[3] > 1 (address jeffbolznv review) * vulkan: address 0cc4m review for fused snake activation snake.comp is renamed to follow the ggml DATA_A_* / A_TYPE convention. A_TYPE now applies to the activation tensor data_a instead of the broadcast multiplier, and the bindings become data_a (A_TYPE), data_b (float), data_c (float) and data_d (D_TYPE). A header at the top of the shader maps each buffer to its role in y = x + sin(b * x)^2 * c. On the C++ side, ggml_vk_can_fuse_snake reuses the existing snake_pattern constant instead of duplicating the op list, sin_node is extracted as a named local alongside the other chain nodes, and the broadcast operands a and inv_b are now required to be GGML_TYPE_F32 to match the hardcoded float bindings on data_b and data_c (the previous a->type == x->type would silently reject any future BF16 or F16 chain once the supports_op gate for SIN / SQR is lifted). ggml_vk_snake_dispatch_fused gets an explicit GGML_TYPE_F32 case and GGML_ABORT on default in place of the silent f32 fallback, and a stale comment about data_a[i1] / data_inv_b[i1] is refreshed to match the new binding names.
50 lines
1.4 KiB
Plaintext
50 lines
1.4 KiB
Plaintext
#version 450
|
|
|
|
#include "types.glsl"
|
|
|
|
// Fused snake activation: y = x + sin(b * x)^2 * c
|
|
// data_a [ne0, ne1] per element activation x (A_TYPE)
|
|
// data_b [1, ne1] per channel multiplier (float)
|
|
// data_c [1, ne1] per channel inverse scale (float, precomputed as 1 / freq)
|
|
// data_d [ne0, ne1] output y (D_TYPE)
|
|
layout (binding = 0) readonly buffer A {A_TYPE data_a[];};
|
|
layout (binding = 1) readonly buffer B {float data_b[];};
|
|
layout (binding = 2) readonly buffer C {float data_c[];};
|
|
layout (binding = 3) writeonly buffer D {D_TYPE data_d[];};
|
|
|
|
layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in;
|
|
|
|
layout (push_constant) uniform parameter {
|
|
uint32_t ne0;
|
|
uint32_t ne1;
|
|
} p;
|
|
|
|
// Load A_TYPE to float
|
|
float load_val(uint32_t idx) {
|
|
#if defined(DATA_A_BF16)
|
|
return bf16_to_fp32(uint32_t(data_a[idx]));
|
|
#else
|
|
return float(data_a[idx]);
|
|
#endif
|
|
}
|
|
|
|
// Store float as D_TYPE
|
|
void store_val(uint32_t idx, float v) {
|
|
#if defined(DATA_D_BF16)
|
|
data_d[idx] = D_TYPE(fp32_to_bf16(v));
|
|
#else
|
|
data_d[idx] = D_TYPE(v);
|
|
#endif
|
|
}
|
|
|
|
void main() {
|
|
const uint32_t i0 = gl_GlobalInvocationID.x;
|
|
const uint32_t i1 = gl_GlobalInvocationID.y;
|
|
if (i0 >= p.ne0 || i1 >= p.ne1) return;
|
|
|
|
const uint32_t idx = i0 + i1 * p.ne0;
|
|
const float xi = load_val(idx);
|
|
const float s = sin(data_b[i1] * xi);
|
|
store_val(idx, xi + s * s * data_c[i1]);
|
|
}
|