/** * WGSL compute shaders for PP-OCRv6 small text recognition. * * The PP-OCRv6_small_rec model (5.2M params) has these major shader requirements: * 1. Conv2D (regular + depthwise) — LCNetV4 backbone * 2. HardSwish / HardSigmoid / Sigmoid activations * 3. BatchNorm (folded into preceding conv at load time) * 4. MaxPool, AveragePool * 5. Reshape / Permute / Slice — layout transforms between 4D feature map and 3D sequence * 6. LayerNorm — SVTR encoder pre-norm * 7. MatMul + bias — linear projections * 8. GELU activation * 9. Multi-head self-attention (Q @ K^T → softmax → @ V) * 10. Softmax over arbitrary axis (vocab for CTC) * 11. ArgMax — CTC decoding * * Layouts used throughout: * - 4D NCHW backbone: [N, C, H, W] flat: n*C*H*W + c*H*W + h*W + w * - 3D sequence: [N, T, D] flat: n*T*D + t*D + d * - 5D QKV split: [N, T, 3, H, D] flat: n*T*3*H*D + t*3*H*D + s*H*D + h*D + d * - 4D attention score map: [N, H, Tq, Tk] flat: n*H*Tq*Tk + h*Tq*Tk + tq*Tk + tk * * Memory layout for weights: * - Conv2D weight: [Cout, Cin/groups, kH, kW] (PyTorch / ONNX contiguous) * - Depthwise: [Cout=groups, 1, kH, kW] when groups = in_channels * - Linear weight: [out_features, in_features] (PyTorch / ONNX convention) * - LayerNorm γ,β: [D] */ export declare const conv2dShader = "\n@group(0) @binding(0) var input: array; // [N, Cin, H, W]\n@group(0) @binding(1) var weight: array; // [Cout, Cin, kH, kW]\n@group(0) @binding(2) var bias: array; // [Cout] (always length Cout or unused)\n@group(0) @binding(3) var output: array; // [N, Cout, Hout, Wout]\n\nstruct Params {\n N: u32,\n Cin: u32,\n Cout: u32,\n H: u32,\n W: u32,\n Hout: u32,\n Wout: u32,\n kH: u32,\n kW: u32,\n strideH: u32,\n strideW: u32,\n padTop: u32,\n padLeft: u32,\n dilH: u32,\n dilW: u32,\n use_bias: u32,\n}\n@group(0) @binding(4) var params: Params;\n\n@compute @workgroup_size(64, 4)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx_n = gid.z; // batch index\n if (idx_n >= params.N) { return; }\n\n let idx_co = gid.y; // output channel pick (within block)\n let idx_xy = gid.x; // flat output position within the channel\n\n let Cout = params.Cout;\n let Hout = params.Hout;\n let Wout = params.Wout;\n let per_ch = Hout * Wout;\n if (idx_co >= Cout || idx_xy >= per_ch) { return; }\n\n let h_out = idx_xy / Wout;\n let w_out = idx_xy % Wout;\n\n let Cin = params.Cin;\n let kH = params.kH;\n let kW = params.kW;\n let strideH = params.strideH;\n let strideW = params.strideW;\n let padTop = params.padTop;\n let padLeft = params.padLeft;\n let dilH = params.dilH;\n let dilW = params.dilW;\n\n let co = idx_co;\n let n = idx_n;\n let n_offset = u32(n) * Cin * params.H * params.W;\n\n var sum = 0.0;\n for (var ci = 0u; ci < Cin; ci++) {\n let ci_offset = n_offset + ci * params.H * params.W;\n for (var kh = 0u; kh < kH; kh++) {\n let h_in = i32(h_out * strideH) + i32(kh * dilH) - i32(padTop);\n if (h_in < 0 || u32(h_in) >= params.H) { continue; }\n for (var kw = 0u; kw < kW; kw++) {\n let w_in = i32(w_out * strideW) + i32(kw * dilW) - i32(padLeft);\n if (w_in < 0 || u32(w_in) >= params.W) { continue; }\n // Weight layout: [Cout, Cin, kH, kW]\n let w_idx = co * Cin * kH * kW + ci * kH * kW + kh * kW + kw;\n let in_idx = ci_offset + u32(h_in) * params.W + u32(w_in);\n sum += input[in_idx] * weight[w_idx];\n }\n }\n }\n\n if (params.use_bias != 0u) {\n sum += bias[co];\n }\n\n // Output layout: [N, Cout, Hout, Wout]\n let out_idx = u32(n) * Cout * per_ch + co * per_ch + idx_xy;\n output[out_idx] = sum;\n}\n"; export declare const conv2dVec4CoutTiledShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var weight: array;\n@group(0) @binding(2) var bias: array;\n@group(0) @binding(3) var output: array;\n\nstruct Params {\n N: u32, Cin: u32, Cout: u32,\n H: u32, W: u32,\n Hout: u32, Wout: u32,\n kH: u32, kW: u32,\n strideH: u32, strideW: u32,\n padTop: u32, padLeft: u32,\n dilH: u32, dilW: u32,\n use_bias: u32,\n}\n@group(0) @binding(4) var params: Params;\n\nconst CTILE: u32 = 16u; // Cin tile width (fixed at compile time)\nconst COUT_PER_WG: u32 = 16u; // 4 y-threads \u00D7 4 Cout/thread = 16 Cout per workgroup\n\n// 16 Cout \u00D7 16 Cin tile \u00D7 9 kernel positions = 2304 f32 for up to 3\u00D73 kernels\nvar sharedW: array;\n\n@compute @workgroup_size(64, 4)\nfn main(\n @builtin(global_invocation_id) gid: vec3,\n @builtin(local_invocation_id) lid: vec3,\n @builtin(local_invocation_index) lii: u32,\n @builtin(workgroup_id) wid: vec3,\n) {\n // No early returns before workgroupBarrier() \u2014 WGSL requires barriers to be in\n // uniform control flow. The dispatch uses wgZ=N and wgY=ceil(Cout/16) exactly,\n // so idx_n < params.N and co_start < params.Cout are always satisfied.\n let idx_n = gid.z;\n let idx_xy = gid.x;\n let per_ch = params.Hout * params.Wout;\n let Cin = params.Cin;\n let kH = params.kH;\n let kW = params.kW;\n let kHW = kH * kW;\n let H = params.H;\n let W = params.W;\n\n // Cout range for this workgroup: [co_start, co_start + 16)\n let co_start = wid.y * COUT_PER_WG;\n // First Cout for this thread: co0 = wid.y*16 + lid.y*4\n let co0 = gid.y * 4u;\n // Local Cout offset: 0, 4, 8, or 12\n let local_co0 = lid.y * 4u;\n\n // Number of valid output channels for this thread (0..4).\n // select(false_val, true_val, cond): u32 subtraction is wrapped but result discarded when cond is false.\n let nco = select(0u, min(4u, params.Cout - co0), co0 < params.Cout);\n\n let w_per_co = Cin * kHW;\n\n var sum0 = 0.0; var sum1 = 0.0; var sum2 = 0.0; var sum3 = 0.0;\n\n // Spatial position for this thread (may be out of range for threads beyond Hout*Wout).\n let h_out = idx_xy / params.Wout;\n let w_out = idx_xy % params.Wout;\n\n let n_offset = idx_n * Cin * H * W;\n let numCTiles = (Cin + CTILE - 1u) / CTILE;\n\n for (var ct = 0u; ct < numCTiles; ct++) {\n let ci_start = ct * CTILE;\n let ci_end = min(ci_start + CTILE, Cin);\n let this_ctile = ci_end - ci_start;\n let stride_co = CTILE * kHW;\n\n // \u2500\u2500 Cooperative weight load (all 256 threads) \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\n // Fixed 9 iterations (ceil(COUT_PER_WG * CTILE * kHW_MAX / 256) = ceil(2304/256))\n // keeps iteration count uniform across all threads so workgroupBarrier is valid.\n let n_entries = COUT_PER_WG * this_ctile * kHW;\n\n for (var li = 0u; li < 9u; li++) {\n let si = lii + li * 256u;\n if (si < n_entries) {\n let flat_kk = si % kHW;\n let local_ci = (si / kHW) % this_ctile;\n let local_co = si / (kHW * this_ctile);\n\n let global_co = co_start + local_co;\n let global_ci = ci_start + local_ci;\n\n var w_val = 0.0;\n if (global_co < params.Cout) {\n w_val = weight[global_co * w_per_co + global_ci * kHW + flat_kk];\n }\n sharedW[local_co * stride_co + local_ci * kHW + flat_kk] = w_val;\n }\n }\n\n workgroupBarrier();\n\n // \u2500\u2500 Accumulate (spatial guard inside; barriers always reached) \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\n for (var lci = 0u; lci < this_ctile; lci++) {\n let ci = ci_start + lci;\n let ci_h_base = n_offset + ci * H * W;\n\n for (var kh = 0u; kh < kH; kh++) {\n let h_in = i32(h_out * params.strideH) + i32(kh * params.dilH) - i32(params.padTop);\n if (h_in < 0 || u32(h_in) >= H || idx_xy >= per_ch) { continue; }\n let in_hw_base = ci_h_base + u32(h_in) * W;\n\n for (var kw = 0u; kw < kW; kw++) {\n let w_in = i32(w_out * params.strideW) + i32(kw * params.dilW) - i32(params.padLeft);\n if (w_in < 0 || u32(w_in) >= W) { continue; }\n let x_val = input[in_hw_base + u32(w_in)];\n\n let w_off = lci * kHW + kh * kW + kw;\n sum0 += x_val * sharedW[ local_co0 * stride_co + w_off];\n if (nco > 1u) { sum1 += x_val * sharedW[(local_co0 + 1u) * stride_co + w_off]; }\n if (nco > 2u) { sum2 += x_val * sharedW[(local_co0 + 2u) * stride_co + w_off]; }\n if (nco > 3u) { sum3 += x_val * sharedW[(local_co0 + 3u) * stride_co + w_off]; }\n }\n }\n }\n\n workgroupBarrier();\n }\n\n if (idx_xy >= per_ch || nco == 0u) { return; }\n\n let out_base = idx_n * params.Cout * per_ch + idx_xy;\n if (params.use_bias != 0u) {\n output[out_base + co0 * per_ch] = sum0 + bias[co0];\n if (nco > 1u) { output[out_base + (co0 + 1u) * per_ch] = sum1 + bias[co0 + 1u]; }\n if (nco > 2u) { output[out_base + (co0 + 2u) * per_ch] = sum2 + bias[co0 + 2u]; }\n if (nco > 3u) { output[out_base + (co0 + 3u) * per_ch] = sum3 + bias[co0 + 3u]; }\n } else {\n output[out_base + co0 * per_ch] = sum0;\n if (nco > 1u) { output[out_base + (co0 + 1u) * per_ch] = sum1; }\n if (nco > 2u) { output[out_base + (co0 + 2u) * per_ch] = sum2; }\n if (nco > 3u) { output[out_base + (co0 + 3u) * per_ch] = sum3; }\n }\n}\n"; export declare const conv2dVec4CoutShader = "\n@group(0) @binding(0) var input: array; // [N, Cin, H, W]\n@group(0) @binding(1) var weight: array; // [Cout, Cin, kH, kW]\n@group(0) @binding(2) var bias: array; // [Cout]\n@group(0) @binding(3) var output: array; // [N, Cout, Hout, Wout]\n\nstruct Params {\n N: u32,\n Cin: u32,\n Cout: u32,\n H: u32,\n W: u32,\n Hout: u32,\n Wout: u32,\n kH: u32,\n kW: u32,\n strideH: u32,\n strideW: u32,\n padTop: u32,\n padLeft: u32,\n dilH: u32,\n dilW: u32,\n use_bias: u32,\n}\n@group(0) @binding(4) var params: Params;\n\n@compute @workgroup_size(64, 4)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx_n = gid.z;\n if (idx_n >= params.N) { return; }\n\n let idx_xy = gid.x;\n let Hout = params.Hout;\n let Wout = params.Wout;\n let per_ch = Hout * Wout;\n if (idx_xy >= per_ch) { return; }\n\n let h_out = idx_xy / Wout;\n let w_out = idx_xy % Wout;\n\n let Cin = params.Cin;\n let kH = params.kH;\n let kW = params.kW;\n let strideH = params.strideH;\n let strideW = params.strideW;\n let padTop = params.padTop;\n let padLeft = params.padLeft;\n let dilH = params.dilH;\n let dilW = params.dilW;\n let H = params.H;\n let W = params.W;\n\n // 4 output channels per thread: co0, co0+1, co0+2, co0+3\n let co0 = gid.y * 4u;\n if (co0 >= params.Cout) { return; }\n // Active channel count (may be < 4 for the last group).\n let nco = min(4u, params.Cout - co0);\n\n let n = idx_n;\n let n_offset = u32(n) * Cin * H * W;\n let w_per_co = Cin * kH * kW;\n\n var sum0 = 0.0;\n var sum1 = 0.0;\n var sum2 = 0.0;\n var sum3 = 0.0;\n\n for (var ci = 0u; ci < Cin; ci++) {\n let ci_offset = n_offset + ci * H * W;\n for (var kh = 0u; kh < kH; kh++) {\n let h_in = i32(h_out * strideH) + i32(kh * dilH) - i32(padTop);\n if (h_in < 0 || u32(h_in) >= H) { continue; }\n let in_h_base = ci_offset + u32(h_in) * W;\n for (var kw = 0u; kw < kW; kw++) {\n let w_in = i32(w_out * strideW) + i32(kw * dilW) - i32(padLeft);\n if (w_in < 0 || u32(w_in) >= W) { continue; }\n let x = input[in_h_base + u32(w_in)];\n let w_base = ci * kH * kW + kh * kW + kw;\n sum0 += x * weight[co0 * w_per_co + w_base];\n if (nco > 1u) { sum1 += x * weight[(co0 + 1u) * w_per_co + w_base]; }\n if (nco > 2u) { sum2 += x * weight[(co0 + 2u) * w_per_co + w_base]; }\n if (nco > 3u) { sum3 += x * weight[(co0 + 3u) * w_per_co + w_base]; }\n }\n }\n }\n\n let out_base = u32(n) * params.Cout * per_ch + idx_xy;\n if (params.use_bias != 0u) {\n if (nco > 0u) { output[out_base + co0 * per_ch] = sum0 + bias[co0]; }\n if (nco > 1u) { output[out_base + (co0 + 1u) * per_ch] = sum1 + bias[co0 + 1u]; }\n if (nco > 2u) { output[out_base + (co0 + 2u) * per_ch] = sum2 + bias[co0 + 2u]; }\n if (nco > 3u) { output[out_base + (co0 + 3u) * per_ch] = sum3 + bias[co0 + 3u]; }\n } else {\n if (nco > 0u) { output[out_base + co0 * per_ch] = sum0; }\n if (nco > 1u) { output[out_base + (co0 + 1u) * per_ch] = sum1; }\n if (nco > 2u) { output[out_base + (co0 + 2u) * per_ch] = sum2; }\n if (nco > 3u) { output[out_base + (co0 + 3u) * per_ch] = sum3; }\n }\n}\n"; export declare const depthwiseConv2dShader = "\n@group(0) @binding(0) var input: array; // [N, C, H, W] C = groups\n@group(0) @binding(1) var weight: array; // [C, 1, kH, kW]\n@group(0) @binding(2) var bias: array; // [C]\n@group(0) @binding(3) var output: array; // [N, C, Hout, Wout]\n\nstruct Params {\n N: u32,\n Cin: u32,\n Cout: u32,\n H: u32,\n W: u32,\n Hout: u32,\n Wout: u32,\n kH: u32,\n kW: u32,\n strideH: u32,\n strideW: u32,\n padTop: u32,\n padLeft: u32,\n dilH: u32,\n dilW: u32,\n use_bias: u32,\n}\n@group(0) @binding(4) var params: Params;\n\n@compute @workgroup_size(64, 4)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let n = gid.z;\n if (n >= params.N) { return; }\n\n let c = gid.y;\n let idx_xy = gid.x;\n let per_ch = params.Hout * params.Wout;\n if (c >= params.Cout || idx_xy >= per_ch) { return; }\n\n let h_out = idx_xy / params.Wout;\n let w_out = idx_xy % params.Wout;\n\n let H = params.H;\n let W = params.W;\n let kH = params.kH;\n let kW = params.kW;\n let strideH = params.strideH;\n let strideW = params.strideW;\n let padTop = params.padTop;\n let padLeft = params.padLeft;\n\n let ci = c * params.Cin / params.Cout;\n let ch_offset = u32(n) * params.Cin * H * W + ci * H * W;\n\n var sum = 0.0;\n for (var kh = 0u; kh < kH; kh++) {\n let h_in = i32(h_out * strideH) + i32(kh) - i32(padTop);\n if (h_in < 0 || u32(h_in) >= H) { continue; }\n for (var kw = 0u; kw < kW; kw++) {\n let w_in = i32(w_out * strideW) + i32(kw) - i32(padLeft);\n if (w_in < 0 || u32(w_in) >= W) { continue; }\n let w_idx = c * kH * kW + kh * kW + kw;\n let in_idx = ch_offset + u32(h_in) * W + u32(w_in);\n sum += input[in_idx] * weight[w_idx];\n }\n }\n\n if (params.use_bias != 0u) {\n sum += bias[c];\n }\n\n let out_idx = u32(n) * params.Cout * per_ch + c * per_ch + idx_xy;\n output[out_idx] = sum;\n}\n"; export declare const hardSigmoidShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n let x = input[idx];\n // clamp((x + 3) / 6, 0, 1) using relu6\n output[idx] = clamp((x + 3.0) * 0.16666666, 0.0, 1.0);\n}\n"; export declare const hardSwishShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n let x = input[idx];\n output[idx] = x * clamp((x + 3.0) * 0.16666666, 0.0, 1.0);\n}\n"; export declare const addShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var b: array;\n@group(0) @binding(2) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(3) var params: Params;\n\n// 64 threads \u00D7 4 elements = 256 elements/workgroup \u2192 same dispatch count as before.\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let base = gid.x * 4u;\n if (base >= params.size) { return; }\n output[base] = a[base] + b[base];\n if (base + 1u < params.size) { output[base + 1u] = a[base + 1u] + b[base + 1u]; }\n if (base + 2u < params.size) { output[base + 2u] = a[base + 2u] + b[base + 2u]; }\n if (base + 3u < params.size) { output[base + 3u] = a[base + 3u] + b[base + 3u]; }\n}\n"; export declare const mulShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var b: array;\n@group(0) @binding(2) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(3) var params: Params;\n\n// 64 threads \u00D7 4 elements = 256 elements/workgroup \u2192 same dispatch count as before.\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let base = gid.x * 4u;\n if (base >= params.size) { return; }\n output[base] = a[base] * b[base];\n if (base + 1u < params.size) { output[base + 1u] = a[base + 1u] * b[base + 1u]; }\n if (base + 2u < params.size) { output[base + 2u] = a[base + 2u] * b[base + 2u]; }\n if (base + 3u < params.size) { output[base + 3u] = a[base + 3u] * b[base + 3u]; }\n}\n"; export declare const addMulShader = "\n@group(0) @binding(0) var x: array;\n@group(0) @binding(1) var scale: array; // per-channel scale\n@group(0) @binding(2) var biasT: array; // per-channel bias\n@group(0) @binding(3) var output: array;\n\nstruct Params {\n channels: u32,\n spatial: u32, // H*W\n use_biasT: u32,\n}\n@group(0) @binding(4) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n let total = params.channels * params.spatial;\n if (idx >= total) { return; }\n let c = idx / params.spatial;\n var v = x[idx] * scale[c];\n if (params.use_biasT != 0u) {\n v = v + biasT[c];\n }\n output[idx] = v;\n}\n"; export declare const maxPool2dShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n N: u32,\n C: u32,\n H: u32,\n W: u32,\n Hout: u32,\n Wout: u32,\n kH: u32,\n kW: u32,\n strideH: u32,\n strideW: u32,\n padTop: u32,\n padLeft: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(64, 4)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let n = gid.z; if (n >= params.N) { return; }\n let c = gid.y; if (c >= params.C) { return; }\n let idx = gid.x;\n let per_ch = params.Hout * params.Wout;\n if (idx >= per_ch) { return; }\n let h_out = idx / params.Wout;\n let w_out = idx % params.Wout;\n\n let ch_off = u32(n) * params.C * params.H * params.W + c * params.H * params.W;\n\n var maxv = -1e30;\n for (var kh = 0u; kh < params.kH; kh++) {\n let h_in = i32(h_out * params.strideH) + i32(kh) - i32(params.padTop);\n if (h_in < 0 || u32(h_in) >= params.H) { continue; }\n for (var kw = 0u; kw < params.kW; kw++) {\n let w_in = i32(w_out * params.strideW) + i32(kw) - i32(params.padLeft);\n if (w_in < 0 || u32(w_in) >= params.W) { continue; }\n let v = input[ch_off + u32(h_in) * params.W + u32(w_in)];\n maxv = max(maxv, v);\n }\n }\n\n let out_idx = u32(n) * params.C * per_ch + c * per_ch + idx;\n output[out_idx] = maxv;\n}\n"; export declare const avgPool2dShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n N: u32,\n C: u32,\n H: u32,\n W: u32,\n Hout: u32,\n Wout: u32,\n kH: u32,\n kW: u32,\n strideH: u32,\n strideW: u32,\n padTop: u32,\n padLeft: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(64, 4)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let n = gid.z; if (n >= params.N) { return; }\n let c = gid.y; if (c >= params.C) { return; }\n let idx = gid.x;\n let per_ch = params.Hout * params.Wout;\n if (idx >= per_ch) { return; }\n let h_out = idx / params.Wout;\n let w_out = idx % params.Wout;\n\n let ch_off = u32(n) * params.C * params.H * params.W + c * params.H * params.W;\n\n var sum = 0.0;\n var count = 0u;\n for (var kh = 0u; kh < params.kH; kh++) {\n let h_in = i32(h_out * params.strideH) + i32(kh) - i32(params.padTop);\n if (h_in < 0 || u32(h_in) >= params.H) { continue; }\n for (var kw = 0u; kw < params.kW; kw++) {\n let w_in = i32(w_out * params.strideW) + i32(kw) - i32(params.padLeft);\n if (w_in < 0 || u32(w_in) >= params.W) { continue; }\n sum += input[ch_off + u32(h_in) * params.W + u32(w_in)];\n count++;\n }\n }\n let avg = select(0.0, sum / f32(count), count > 0u);\n\n let out_idx = u32(n) * params.C * per_ch + c * per_ch + idx;\n output[out_idx] = avg;\n}\n"; export declare const globalAvgPool2dShader = "\n@group(0) @binding(0) var input: array; // [N, C, H, W]\n@group(0) @binding(1) var output: array; // [N, C]\n\nstruct Params {\n N: u32,\n C: u32,\n H: u32,\n W: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let n = gid.y;\n let c = gid.x;\n if (n >= params.N || c >= params.C) { return; }\n\n let total = params.H * params.W;\n let base = (u32(n) * params.C + c) * total;\n\n var sum = 0.0;\n for (var i = 0u; i < total; i++) {\n sum += input[base + i];\n }\n output[u32(n) * params.C + c] = sum / f32(total);\n}\n"; export declare const reluShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = max(0.0, input[idx]);\n}\n"; export declare const relu6Shader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = clamp(input[idx], 0.0, 6.0);\n}\n"; export declare const layerNormShader = "\n@group(0) @binding(0) var input: array; // [Batch, D]\n@group(0) @binding(1) var gamma: array; // [D]\n@group(0) @binding(2) var beta: array; // [D]\n@group(0) @binding(3) var output: array; // [Batch, D]\n\nstruct Params {\n batch: u32,\n D: u32,\n eps: f32,\n}\n@group(0) @binding(4) var params: Params;\n\n// One workgroup per token; 128 threads share the reduction across D elements.\nvar wg: array;\n\n@compute @workgroup_size(128)\nfn main(\n @builtin(workgroup_id) wid: vec3,\n @builtin(local_invocation_index) lii: u32,\n) {\n let b = wid.x;\n if (b >= params.batch) { return; }\n let base = b * params.D;\n let D = params.D;\n\n // Phase 1: parallel sum \u2192 mean\n var s = 0.0;\n var i = lii;\n loop { if (i >= D) { break; } s += input[base + i]; i += 128u; }\n wg[lii] = s;\n workgroupBarrier();\n for (var stride = 64u; stride > 0u; stride >>= 1u) {\n if (lii < stride) { wg[lii] += wg[lii + stride]; }\n workgroupBarrier();\n }\n let mean = wg[0] / f32(D);\n\n // Phase 2: parallel sum of (x - mean)\u00B2 \u2192 variance\n var v = 0.0;\n i = lii;\n loop {\n if (i >= D) { break; }\n let diff = input[base + i] - mean;\n v += diff * diff;\n i += 128u;\n }\n wg[lii] = v;\n workgroupBarrier();\n for (var stride = 64u; stride > 0u; stride >>= 1u) {\n if (lii < stride) { wg[lii] += wg[lii + stride]; }\n workgroupBarrier();\n }\n let inv_std = 1.0 / sqrt(wg[0] / f32(D) + params.eps);\n\n // Phase 3: normalize + affine\n i = lii;\n loop {\n if (i >= D) { break; }\n output[base + i] = (input[base + i] - mean) * inv_std * gamma[i] + beta[i];\n i += 128u;\n }\n}\n"; export declare const geluShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n let x = input[idx];\n let c: f32 = 0.7978845608028654; // sqrt(2 / pi)\n let inner = clamp(c * (x + 0.044715 * x * x * x), -15.0, 15.0);\n output[idx] = 0.5 * x * (1.0 + tanh(inner));\n}\n"; export declare const linearShader = "\nconst TILE: u32 = 16u;\n\n@group(0) @binding(0) var A: array; // [M, K]\n@group(0) @binding(1) var Wt: array; // [N, K] (linear weight, stored as [out, in])\n@group(0) @binding(2) var bias: array; // [N] (or empty if use_bias=0)\n@group(0) @binding(3) var Y: array; // [M, N]\n\nstruct Params {\n M: u32,\n K: u32,\n N: u32,\n use_bias: u32,\n}\n@group(0) @binding(4) var params: Params;\n\nvar tileA: array;\nvar tileW: array;\n\n@compute @workgroup_size(16, 16)\nfn main(\n @builtin(global_invocation_id) gid: vec3,\n @builtin(local_invocation_id) lid: vec3,\n) {\n let row = gid.x;\n let col = gid.y;\n let lr = lid.x;\n let lc = lid.y;\n\n var acc = 0.0;\n let numTiles = (params.K + TILE - 1u) / TILE;\n\n for (var t = 0u; t < numTiles; t++) {\n let aCol = t * TILE + lc;\n if (row < params.M && aCol < params.K) {\n tileA[lr * TILE + lc] = A[row * params.K + aCol];\n } else {\n tileA[lr * TILE + lc] = 0.0;\n }\n\n let wCol = t * TILE + lr;\n if (col < params.N && wCol < params.K) {\n // W is stored [K, N] (in_features, out_features) as in ONNX MatMul.\n // Y[m,n] = \u03A3_k A[m,k] * W[k,n]. W[k,n] = W[k*N + n].\n tileW[lr * TILE + lc] = Wt[wCol * params.N + col];\n } else {\n tileW[lr * TILE + lc] = 0.0;\n }\n\n workgroupBarrier();\n\n for (var k = 0u; k < TILE; k++) {\n acc += tileA[lr * TILE + k] * tileW[k * TILE + lc];\n }\n\n workgroupBarrier();\n }\n\n if (row < params.M && col < params.N) {\n if (params.use_bias != 0u) {\n Y[row * params.N + col] = acc + bias[col];\n } else {\n Y[row * params.N + col] = acc;\n }\n }\n}\n"; export declare const linearWideShader = "\nconst TILE_M: u32 = 16u;\nconst TILE_K: u32 = 16u;\nconst TILE_N: u32 = 64u;\n\n@group(0) @binding(0) var A: array;\n@group(0) @binding(1) var Wt: array; // [N, K] (transposed)\n@group(0) @binding(2) var bias: array;\n@group(0) @binding(3) var Y: array;\n\nstruct Params { M: u32, K: u32, N: u32, use_bias: u32 }\n@group(0) @binding(4) var params: Params;\n\nvar tileA: array; // 16 \u00D7 16\nvar tileW: array; // 16 \u00D7 64\n\n@compute @workgroup_size(TILE_M, TILE_M) // 256 threads\nfn main(\n @builtin(workgroup_id) wid: vec3,\n @builtin(local_invocation_id) lid: vec3,\n) {\n let r = lid.x;\n let c = lid.y;\n let row = wid.x * TILE_M + r;\n let colBase = wid.y * TILE_N + c * 4u;\n let li = r * TILE_M + c; // 0..255\n\n var acc = vec4(0.0);\n let numTiles = (params.K + TILE_K - 1u) / TILE_K;\n\n for (var t = 0u; t < numTiles; t++) {\n // Load A tile (1 element/thread, linearly indexed)\n let aRow = wid.x * TILE_M + li / TILE_K;\n let aCol = t * TILE_K + li % TILE_K;\n tileA[li] = select(0.0, A[aRow * params.K + aCol],\n aRow < params.M && aCol < params.K);\n\n // Load W tile (4 elements/thread via linearized index \u2014 coalesced).\n // Consecutive threads (same SIMD group) get consecutive li values \u2192\n // consecutive w_n values \u2192 consecutive N addresses in the same K row.\n for (var jj = 0u; jj < 4u; jj++) {\n let w_li = li + jj * 256u; // 0..1023 = TILE_K \u00D7 TILE_N\n let w_k = w_li / TILE_N; // 0..15\n let w_n = w_li % TILE_N; // 0..63\n let wRow = t * TILE_K + w_k;\n let wCol = wid.y * TILE_N + w_n;\n tileW[w_li] = select(0.0, Wt[wRow * params.N + wCol],\n wRow < params.K && wCol < params.N);\n }\n workgroupBarrier();\n\n for (var k = 0u; k < TILE_K; k++) {\n let a = tileA[r * TILE_K + k];\n acc += a * vec4(\n tileW[k * TILE_N + c * 4u],\n tileW[k * TILE_N + c * 4u + 1u],\n tileW[k * TILE_N + c * 4u + 2u],\n tileW[k * TILE_N + c * 4u + 3u],\n );\n }\n workgroupBarrier();\n }\n\n if (row >= params.M) { return; }\n let bv = select(vec4(0.0),\n vec4(bias[colBase], bias[colBase+1u], bias[colBase+2u], bias[colBase+3u]),\n params.use_bias != 0u);\n if (colBase < params.N) { Y[row * params.N + colBase ] = acc.x + bv.x; }\n if (colBase + 1u < params.N) { Y[row * params.N + colBase + 1u ] = acc.y + bv.y; }\n if (colBase + 2u < params.N) { Y[row * params.N + colBase + 2u ] = acc.z + bv.z; }\n if (colBase + 3u < params.N) { Y[row * params.N + colBase + 3u ] = acc.w + bv.w; }\n}\n"; export declare const linearBatchedShader = "\nconst TILE: u32 = 16u;\n\n@group(0) @binding(0) var A: array;\n@group(0) @binding(1) var W: array;\n@group(0) @binding(2) var bias: array;\n@group(0) @binding(3) var Y: array;\n\nstruct Params {\n B: u32,\n M: u32,\n K: u32,\n N: u32,\n use_bias: u32,\n W_batched: u32, // 1 = W also varies per batch (e.g. Q@K^T attention)\n}\n@group(0) @binding(4) var params: Params;\n\nvar tileA: array;\nvar tileW: array;\n\n@compute @workgroup_size(16, 16, 1)\nfn main(\n @builtin(global_invocation_id) gid: vec3,\n @builtin(workgroup_id) wid: vec3,\n @builtin(local_invocation_id) lid: vec3,\n) {\n let batch = wid.z;\n let row = gid.x;\n let col = gid.y;\n let lr = lid.x;\n let lc = lid.y;\n\n let aBase = batch * params.M * params.K;\n let wBase = select(0u, batch * params.K * params.N, params.W_batched != 0u);\n let yBase = batch * params.M * params.N;\n var acc = 0.0;\n let numTiles = (params.K + TILE - 1u) / TILE;\n for (var t = 0u; t < numTiles; t++) {\n let aCol = t * TILE + lc;\n if (row < params.M && aCol < params.K) {\n tileA[lr * TILE + lc] = A[aBase + row * params.K + aCol];\n } else {\n tileA[lr * TILE + lc] = 0.0;\n }\n let wCol = t * TILE + lr;\n if (col < params.N && wCol < params.K) {\n // W is stored [K, N] (in_features, out_features) as in ONNX MatMul.\n // Y[b,m,n] = \u03A3_k A[b,m,k] * W[k,n]. W[k,n] = W[k*N + n].\n tileW[lr * TILE + lc] = W[wBase + wCol * params.N + col];\n } else {\n tileW[lr * TILE + lc] = 0.0;\n }\n workgroupBarrier();\n for (var k = 0u; k < TILE; k++) {\n acc += tileA[lr * TILE + k] * tileW[k * TILE + lc];\n }\n workgroupBarrier();\n }\n if (row < params.M && col < params.N) {\n if (params.use_bias != 0u) {\n Y[yBase + row * params.N + col] = acc + bias[col];\n } else {\n Y[yBase + row * params.N + col] = acc;\n }\n }\n}\n"; export declare const linearGeluShader = "\nconst TILE: u32 = 16u;\n\n@group(0) @binding(0) var A: array;\n@group(0) @binding(1) var Wt: array;\n@group(0) @binding(2) var bias: array;\n@group(0) @binding(3) var Y: array;\n\nstruct Params {\n M: u32,\n K: u32,\n N: u32,\n}\n@group(0) @binding(4) var params: Params;\n\nvar tileA: array;\nvar tileW: array;\n\nconst SQRT_PI_INV: f32 = 0.7978845608028654; // sqrt(2 / pi)\n\n@compute @workgroup_size(16, 16)\nfn main(\n @builtin(global_invocation_id) gid: vec3,\n @builtin(local_invocation_id) lid: vec3,\n) {\n let row = gid.x;\n let col = gid.y;\n let lr = lid.x;\n let lc = lid.y;\n\n var acc = 0.0;\n let numTiles = (params.K + TILE - 1u) / TILE;\n\n for (var t = 0u; t < numTiles; t++) {\n let aCol = t * TILE + lc;\n if (row < params.M && aCol < params.K) {\n tileA[lr * TILE + lc] = A[row * params.K + aCol];\n } else {\n tileA[lr * TILE + lc] = 0.0;\n }\n let wCol = t * TILE + lr;\n if (col < params.N && wCol < params.K) {\n // W is stored [K, N] (in_features, out_features) as in ONNX MatMul.\n tileW[lr * TILE + lc] = Wt[wCol * params.N + col];\n } else {\n tileW[lr * TILE + lc] = 0.0;\n }\n\n workgroupBarrier();\n for (var k = 0u; k < TILE; k++) {\n acc += tileA[lr * TILE + k] * tileW[k * TILE + lc];\n }\n workgroupBarrier();\n }\n\n if (row < params.M && col < params.N) {\n var v = acc + bias[col];\n // Fused tanh-based GELU on the pre-activation\n let c: f32 = 0.7978845608028654; // sqrt(2 / pi)\n let inner = clamp(c * (v + 0.044715 * v * v * v), -15.0, 15.0);\n v = 0.5 * v * (1.0 + tanh(inner));\n Y[row * params.N + col] = v;\n }\n}\n"; export declare const softmaxShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n batch_size: u32,\n dim_size: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let b = gid.x;\n if (b >= params.batch_size) { return; }\n let off = b * params.dim_size;\n\n var maxv = input[off];\n for (var i = 1u; i < params.dim_size; i++) {\n maxv = max(maxv, input[off + i]);\n }\n\n var esum = 0.0;\n for (var i = 0u; i < params.dim_size; i++) {\n let e = exp(input[off + i] - maxv);\n output[off + i] = e;\n esum += e;\n }\n\n for (var i = 0u; i < params.dim_size; i++) {\n output[off + i] = output[off + i] / esum;\n }\n}\n"; export declare const softmaxParallelShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { batch_size: u32, dim_size: u32 }\n@group(0) @binding(2) var params: Params;\n\nvar wg_buf: array;\n\n@compute @workgroup_size(256)\nfn main(\n @builtin(workgroup_id) wid: vec3,\n @builtin(local_invocation_index) lii: u32,\n) {\n let b = wid.x;\n if (b >= params.batch_size) { return; }\n\n let off = b * params.dim_size;\n let D = params.dim_size;\n\n // \u2500\u2500 Phase 1: parallel max \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\n var local_max = -1e30;\n var i = lii;\n loop {\n if (i >= D) { break; }\n local_max = max(local_max, input[off + i]);\n i += 256u;\n }\n wg_buf[lii] = local_max;\n workgroupBarrier();\n for (var s = 128u; s > 0u; s = s >> 1u) {\n if (lii < s) { wg_buf[lii] = max(wg_buf[lii], wg_buf[lii + s]); }\n workgroupBarrier();\n }\n let gmax = wg_buf[0];\n\n // \u2500\u2500 Phase 2: exp + partial sum \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\n var local_sum = 0.0;\n i = lii;\n loop {\n if (i >= D) { break; }\n let e = exp(input[off + i] - gmax);\n output[off + i] = e;\n local_sum += e;\n i += 256u;\n }\n wg_buf[lii] = local_sum;\n workgroupBarrier();\n for (var s = 128u; s > 0u; s = s >> 1u) {\n if (lii < s) { wg_buf[lii] = wg_buf[lii] + wg_buf[lii + s]; }\n workgroupBarrier();\n }\n let gsum = wg_buf[0];\n\n // \u2500\u2500 Phase 3: normalize \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\n i = lii;\n loop {\n if (i >= D) { break; }\n output[off + i] = output[off + i] / gsum;\n i += 256u;\n }\n}\n"; export declare const softmaxScaleShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n outer: u32,\n inner: u32,\n scale: f32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let o = gid.x;\n if (o >= params.outer) { return; }\n let off = o * params.inner;\n\n var maxv = -1e30;\n for (var i = 0u; i < params.inner; i++) {\n let v = input[off + i] * params.scale;\n maxv = max(maxv, v);\n }\n\n var esum = 0.0;\n for (var i = 0u; i < params.inner; i++) {\n let v = input[off + i] * params.scale;\n let e = exp(v - maxv);\n output[off + i] = e;\n esum += e;\n }\n\n for (var i = 0u; i < params.inner; i++) {\n output[off + i] = output[off + i] / esum;\n }\n}\n"; export declare const argMaxShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n batch_size: u32,\n dim_size: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let b = gid.x;\n if (b >= params.batch_size) { return; }\n let off = b * params.dim_size;\n\n var best = 0u;\n var bestv = input[off];\n for (var i = 1u; i < params.dim_size; i++) {\n let v = input[off + i];\n if (v > bestv) {\n bestv = v;\n best = i;\n }\n }\n output[b] = best;\n}\n"; export declare const argMaxCtcShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n batch_size: u32,\n dim_size: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n// One workgroup per timestep; 256 threads split the 18710-vocab scan in parallel.\nvar wg_val: array;\nvar wg_idx: array;\n\n@compute @workgroup_size(256)\nfn main(\n @builtin(workgroup_id) wid: vec3,\n @builtin(local_invocation_index) lii: u32,\n) {\n let b = wid.x;\n if (b >= params.batch_size) { return; }\n let off = b * params.dim_size;\n let D = params.dim_size;\n\n // Each thread scans its strided slice of the vocab dimension.\n var best_idx = lii;\n var best_val = select(-1e30, input[off + lii], lii < D);\n var i = lii + 256u;\n loop {\n if (i >= D) { break; }\n let v = input[off + i];\n if (v > best_val) { best_val = v; best_idx = i; }\n i += 256u;\n }\n wg_val[lii] = best_val;\n wg_idx[lii] = best_idx;\n workgroupBarrier();\n\n // Tree reduction: keep the max value (ties broken by lower index).\n for (var stride = 128u; stride > 0u; stride >>= 1u) {\n if (lii < stride) {\n if (wg_val[lii + stride] > wg_val[lii]) {\n wg_val[lii] = wg_val[lii + stride];\n wg_idx[lii] = wg_idx[lii + stride];\n }\n }\n workgroupBarrier();\n }\n\n if (lii == 0u) {\n // output[b*2] = index (as f32), output[b*2+1] = max value\n output[b * 2u] = f32(wg_idx[0]);\n output[b * 2u + 1u] = wg_val[0];\n }\n}\n"; export declare const permuteShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n D0: u32,\n D1: u32,\n D2: u32,\n D3: u32,\n S0: u32,\n S1: u32,\n S2: u32,\n S3: u32,\n perm0: u32,\n perm1: u32,\n perm2: u32,\n perm3: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let total = params.D0 * params.D1 * params.D2 * params.D3;\n let idx = gid.x;\n if (idx >= total) { return; }\n\n // Decompose flat output index into 4D output coords using OUTPUT shape.\n let q0 = params.D1 * params.D2 * params.D3;\n let q1 = params.D2 * params.D3;\n let q2 = params.D3;\n\n let a = idx / q0;\n let r1 = idx % q0;\n let b = r1 / q1;\n let r2 = r1 % q1;\n let c = r2 / q2;\n let d = r2 % q2;\n\n // Output axis k samples input axis perm_k. Use INPUT strides (S0..S3)\n // \u2014 NOT output strides \u2014 to compute the source flat index.\n var src_idx = 0u;\n if (params.perm0 == 0u) { src_idx += a * params.S0; }\n else if (params.perm0 == 1u) { src_idx += a * params.S1; }\n else if (params.perm0 == 2u) { src_idx += a * params.S2; }\n else { src_idx += a * params.S3; }\n if (params.perm1 == 0u) { src_idx += b * params.S0; }\n else if (params.perm1 == 1u) { src_idx += b * params.S1; }\n else if (params.perm1 == 2u) { src_idx += b * params.S2; }\n else { src_idx += b * params.S3; }\n if (params.perm2 == 0u) { src_idx += c * params.S0; }\n else if (params.perm2 == 1u) { src_idx += c * params.S1; }\n else if (params.perm2 == 2u) { src_idx += c * params.S2; }\n else { src_idx += c * params.S3; }\n if (params.perm3 == 0u) { src_idx += d * params.S0; }\n else if (params.perm3 == 1u) { src_idx += d * params.S1; }\n else if (params.perm3 == 2u) { src_idx += d * params.S2; }\n else { src_idx += d * params.S3; }\n\n output[idx] = input[src_idx];\n}\n"; export declare const permute5dShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n O0: u32, O1: u32, O2: u32, O3: u32, O4: u32,\n IS0: u32, IS1: u32, IS2: u32, IS3: u32, IS4: u32,\n perm0: u32, perm1: u32, perm2: u32, perm3: u32, perm4: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let total = params.O0 * params.O1 * params.O2 * params.O3 * params.O4;\n let idx = gid.x;\n if (idx >= total) { return; }\n\n // Decompose flat output index \u2192 5D output coordinates.\n let q01 = params.O1 * params.O2 * params.O3 * params.O4;\n let q12 = params.O2 * params.O3 * params.O4;\n let q23 = params.O3 * params.O4;\n let q34 = params.O4;\n\n let a = idx / q01;\n let r1 = idx % q01;\n let b = r1 / q12;\n let r2 = r1 % q12;\n let c = r2 / q23;\n let r3 = r2 % q23;\n let d = r3 / q34;\n let e = r3 % q34;\n\n // For each output axis k, perm_k names the input axis; multiply by IS_perm_k.\n var src_idx = 0u;\n if (params.perm0 == 0u) { src_idx += a * params.IS0; }\n else if (params.perm0 == 1u) { src_idx += a * params.IS1; }\n else if (params.perm0 == 2u) { src_idx += a * params.IS2; }\n else if (params.perm0 == 3u) { src_idx += a * params.IS3; }\n else { src_idx += a * params.IS4; }\n\n if (params.perm1 == 0u) { src_idx += b * params.IS0; }\n else if (params.perm1 == 1u) { src_idx += b * params.IS1; }\n else if (params.perm1 == 2u) { src_idx += b * params.IS2; }\n else if (params.perm1 == 3u) { src_idx += b * params.IS3; }\n else { src_idx += b * params.IS4; }\n\n if (params.perm2 == 0u) { src_idx += c * params.IS0; }\n else if (params.perm2 == 1u) { src_idx += c * params.IS1; }\n else if (params.perm2 == 2u) { src_idx += c * params.IS2; }\n else if (params.perm2 == 3u) { src_idx += c * params.IS3; }\n else { src_idx += c * params.IS4; }\n\n if (params.perm3 == 0u) { src_idx += d * params.IS0; }\n else if (params.perm3 == 1u) { src_idx += d * params.IS1; }\n else if (params.perm3 == 2u) { src_idx += d * params.IS2; }\n else if (params.perm3 == 3u) { src_idx += d * params.IS3; }\n else { src_idx += d * params.IS4; }\n\n if (params.perm4 == 0u) { src_idx += e * params.IS0; }\n else if (params.perm4 == 1u) { src_idx += e * params.IS1; }\n else if (params.perm4 == 2u) { src_idx += e * params.IS2; }\n else if (params.perm4 == 3u) { src_idx += e * params.IS3; }\n else { src_idx += e * params.IS4; }\n\n output[idx] = input[src_idx];\n}\n"; export declare const reshapePassthroughShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = input[idx];\n}\n"; export declare const squeezeShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n // Input total size = product of dims\n in_total: u32,\n out_total: u32,\n // Strides of each input axis (excluding the squeezed dim of size 1)\n s0: u32, s1: u32, s2: u32, s3: u32,\n // Sizes of each input axis (excluding the squeezed one)\n d0: u32, d1: u32, d2: u32, d3: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.out_total) { return; }\n let rem = idx;\n let a = rem / (params.d1 * params.d2 * params.d3);\n let r = rem % (params.d1 * params.d2 * params.d3);\n let b = r / (params.d2 * params.d3);\n let r2 = r % (params.d2 * params.d3);\n let c = r2 / params.d3;\n let d = r2 % params.d3;\n output[idx] = input[a * params.s0 + b * params.s1 + c * params.s2 + d * params.s3];\n}\n"; export declare const sliceShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n // Source shape and \"start\" offset\n start0: u32, start1: u32, start2: u32, start3: u32,\n // Destination/output shape (= sliced portion)\n D0: u32, D1: u32, D2: u32, D3: u32,\n // Source strides (input shape with given slice dim range collapsed)\n s0: u32, s1: u32, s2: u32, s3: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let total = params.D0 * params.D1 * params.D2 * params.D3;\n let idx = gid.x;\n if (idx >= total) { return; }\n let rem = idx;\n let a = rem / (params.D1 * params.D2 * params.D3);\n let r = rem % (params.D1 * params.D2 * params.D3);\n let b = r / (params.D2 * params.D3);\n let r2 = r % (params.D2 * params.D3);\n let c = r2 / params.D3;\n let d = r2 % params.D3;\n\n output[idx] = input[\n (a + params.start0) * params.s0 +\n (b + params.start1) * params.s1 +\n (c + params.start2) * params.s2 +\n (d + params.start3) * params.s3\n ];\n}\n"; export declare const concatChannelsShader = "\n@group(0) @binding(0) var a: array; // [N, Ca, W]\n@group(0) @binding(1) var b: array; // [N, Cb, W]\n@group(0) @binding(2) var output: array; // [N, Ca+Cb, W]\n\nstruct Params {\n N: u32,\n Ca: u32,\n Cb: u32,\n W: u32,\n}\n@group(0) @binding(3) var params: Params;\n\n@compute @workgroup_size(64, 4)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let n = gid.z;\n let c = gid.y;\n let w = gid.x;\n if (n >= params.N || w >= params.W) { return; }\n let c_total = params.Ca + params.Cb;\n if (c >= c_total) { return; }\n if (c < params.Ca) {\n let idx = n * c_total * params.W + c * params.W + w;\n output[idx] = a[n * params.Ca * params.W + c * params.W + w];\n } else {\n let cb = c - params.Ca;\n let idx = n * c_total * params.W + c * params.W + w;\n output[idx] = b[n * params.Cb * params.W + cb * params.W + w];\n }\n}\n"; export declare const subShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var b: array;\n@group(0) @binding(2) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(3) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = a[idx] - b[idx];\n}\n"; export declare const divShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var b: array;\n@group(0) @binding(2) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(3) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = a[idx] / b[idx];\n}\n"; export declare const powShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n size: u32,\n _pad: u32, _pad2: u32,\n exp: f32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n let v = input[idx];\n // WGSL pow() is defined as exp(log(x)*exp) \u2192 NaN for x<0 even with integer\n // exponents. LayerNorm uses Pow(x, 2) for variance = (x-mean)\u00B2 where x-mean\n // can be negative. Use x*x for exp==2; for other integer exponents fall back\n // to sign-corrected pow(abs(x), exp).\n if (params.exp == 2.0) {\n output[idx] = v * v;\n } else {\n let a = abs(v);\n let r = pow(a, params.exp);\n output[idx] = select(r, -r, (v < 0.0) && (fract(params.exp) == 0.0));\n }\n}\n"; export declare const scaleShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var scale: array;\n@group(0) @binding(2) var output: array;\n\nstruct Params {\n channels: u32,\n spatial: u32,\n}\n@group(0) @binding(3) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n let total = params.channels * params.spatial;\n if (idx >= total) { return; }\n let c = idx / params.spatial;\n output[idx] = input[idx] * scale[c];\n}\n"; export declare const broadcastBinaryShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var b: array;\n@group(0) @binding(2) var output: array;\n\nstruct Params {\n out_size: u32,\n // Output dims (4D)\n O0: u32, O1: u32, O2: u32, O3: u32,\n // Source A dims (right-aligned to 4D)\n a0: u32, a1: u32, a2: u32, a3: u32,\n // Source A strides (number of elements per axis)\n sa0: u32, sa1: u32, sa2: u32, sa3: u32,\n // Source B dims (right-aligned to 4D)\n b0: u32, b1: u32, b2: u32, b3: u32,\n // Source B strides\n sb0: u32, sb1: u32, sb2: u32, sb3: u32,\n op: u32, // 0=add, 1=sub, 2=mul, 3=div\n}\n@group(0) @binding(3) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.out_size) { return; }\n\n let o3 = idx % params.O3;\n let o2 = (idx / params.O3) % params.O2;\n let o1 = (idx / (params.O3 * params.O2)) % params.O1;\n let o0 = idx / (params.O3 * params.O2 * params.O1);\n\n // Broadcast rule: if dim size == 1 the axis is broadcast \u2192 always read index 0\n // on that axis (add 0). If dim size > 1 advance by the output coordinate.\n let sa = select(0u, o0 * params.sa0, params.a0 != 1u)\n + select(0u, o1 * params.sa1, params.a1 != 1u)\n + select(0u, o2 * params.sa2, params.a2 != 1u)\n + select(0u, o3 * params.sa3, params.a3 != 1u);\n let sb = select(0u, o0 * params.sb0, params.b0 != 1u)\n + select(0u, o1 * params.sb1, params.b1 != 1u)\n + select(0u, o2 * params.sb2, params.b2 != 1u)\n + select(0u, o3 * params.sb3, params.b3 != 1u);\n\n let av = a[sa];\n let bv = b[sb];\n if (params.op == 0u) {\n output[idx] = av + bv;\n } else if (params.op == 1u) {\n output[idx] = av - bv;\n } else if (params.op == 2u) {\n output[idx] = av * bv;\n } else {\n output[idx] = av / bv;\n }\n}\n"; export declare const mhaQKTScaledShader = "\n@group(0) @binding(0) var Q: array; // [N, H, Tq, D]\n@group(0) @binding(1) var K: array; // [N, H, Tk, D]\n@group(0) @binding(2) var V: array; // [N, H, Tk, D]\n@group(0) @binding(3) var output: array; // [N, H, Tq, D]\n\nstruct Params {\n N: u32,\n H: u32,\n Tq: u32,\n Tk: u32,\n D: u32,\n}\n@group(0) @binding(4) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let d_idx = gid.x;\n let nh_t = gid.y;\n let n = nh_t / (params.H * params.Tq);\n let rem = nh_t % (params.H * params.Tq);\n let h = rem / params.Tq;\n let tq = rem % params.Tq;\n if (d_idx >= params.D || n >= params.N || h >= params.H || tq >= params.Tq) { return; }\n\n let scale = 1.0 / sqrt(f32(params.D));\n let qoff = ((n * params.H + h) * params.Tq + tq) * params.D;\n\n // Online softmax (Milakov & Gimelshein 2018): single pass over Tk,\n // eliminating the second QK^T scan that the two-pass approach required.\n var running_max = -1e30;\n var running_sum = 0.0;\n var vd = 0.0;\n\n for (var tk = 0u; tk < params.Tk; tk++) {\n let koff = ((n * params.H + h) * params.Tk + tk) * params.D;\n var s = 0.0;\n for (var d = 0u; d < params.D; d++) {\n s += Q[qoff + d] * K[koff + d];\n }\n s *= scale;\n\n let new_max = max(running_max, s);\n let correction = exp(running_max - new_max);\n let e_s = exp(s - new_max);\n let voff = ((n * params.H + h) * params.Tk + tk) * params.D;\n vd = vd * correction + e_s * V[voff + d_idx];\n running_sum = running_sum * correction + e_s;\n running_max = new_max;\n }\n\n let o_off = ((n * params.H + h) * params.Tq + tq) * params.D;\n output[o_off + d_idx] = vd / running_sum;\n}\n"; export declare const collapseHtoSeqShader = "\n@group(0) @binding(0) var input: array; // [N, C, 1, W]\n@group(0) @binding(1) var output: array; // [N, W, C]\n\nstruct Params {\n N: u32,\n C: u32,\n W: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(64, 4)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let n = gid.z;\n if (n >= params.N) { return; }\n let c = gid.y;\n if (c >= params.C) { return; }\n let w = gid.x;\n if (w >= params.W) { return; }\n\n // Input: [N, C, 1, W] \u2014 read input[n, c, 0, w]\n let in_idx = (n * params.C + c) * params.W + w;\n // Output: [N, W, C] \u2014 write output[n, w, c]\n let out_idx = (n * params.W + w) * params.C + c;\n output[out_idx] = input[in_idx];\n}\n"; export declare const reshapeToQkvShader = "\n@group(0) @binding(0) var input: array; // [N, T, 3*H*D]\n@group(0) @binding(1) var output: array; // [N, T, 3, H, D]\n\nstruct Params {\n N: u32,\n T: u32,\n Hheads: u32,\n D: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let n = gid.z;\n if (n >= params.N) { return; }\n let t = gid.x;\n if (t >= params.T) { return; }\n let inner = gid.y;\n let inner_size = 3u * params.Hheads * params.D;\n if (inner >= inner_size) { return; }\n\n // Same flat index since reshape only regroups.\n let idx = (n * params.T + t) * inner_size + inner;\n output[idx] = input[idx];\n}\n"; export declare const qkvPackShader = "\n@group(0) @binding(0) var input: array; // [N, T, 3, H, D]\n@group(0) @binding(1) var Q: array; // [N, H, T, D]\n@group(0) @binding(2) var K: array; // [N, H, T, D]\n@group(0) @binding(3) var V: array; // [N, H, T, D]\n\nstruct Params {\n N: u32,\n T: u32,\n Hheads: u32,\n D: u32,\n}\n@group(0) @binding(4) var params: Params;\n\n@compute @workgroup_size(64, 4)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let t = gid.x;\n let h = gid.y;\n let nrest = gid.z;\n if (t >= params.T || h >= params.Hheads) { return; }\n let n = nrest;\n if (n >= params.N) { return; }\n\n // Read Q from input[n, t, 0, h, :]\n let q_src = ((n * params.T + t) * 3u) * params.Hheads * params.D + h * params.D;\n let k_src = q_src + params.Hheads * params.D;\n let v_src = k_src + params.Hheads * params.D;\n\n let dst = ((n * params.Hheads + h) * params.T + t) * params.D;\n\n for (var d = 0u; d < params.D; d++) {\n Q[dst + d] = input[q_src + d];\n K[dst + d] = input[k_src + d];\n V[dst + d] = input[v_src + d];\n }\n}\n"; export declare const concatHeadsShader = "\n@group(0) @binding(0) var input: array; // [N, H, T, D]\n@group(0) @binding(1) var output: array; // [N, T, H*D]\n\nstruct Params {\n N: u32,\n T: u32,\n Hheads: u32,\n D: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let t = gid.x;\n if (t >= params.T) { return; }\n let hd = gid.y;\n let inner_size = params.Hheads * params.D;\n if (hd >= inner_size) { return; }\n let n = gid.z;\n if (n >= params.N) { return; }\n\n let h = hd / params.D;\n let d = hd % params.D;\n\n let src = ((n * params.Hheads + h) * params.T + t) * params.D + d;\n let dst = n * params.T * inner_size + t * inner_size + hd;\n output[dst] = input[src];\n}\n"; export declare const sqrtShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = sqrt(input[idx]);\n}\n"; export declare const erfShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n // Polynomial approximation of the error function (Abramowitz & Stegun 7.1.26).\n // WGSL does not expose an erf() intrinsic on every backend; we approximate from\n // inside the shader instead.\n //\n // erf(x) = sign(x) * (1 - tau)\n // tau = t * exp(-x^2) * (a1 + t * (a2 + t * (a3 + t * (a4 + t * a5))))\n // t = 1 / (1 + p * |x|) \u2190 uses abs(x), NOT x^2\n //\n // Max error ~1.5e-7 over the real line.\n let x = clamp(input[idx], -6.0, 6.0);\n let ax = abs(x);\n let xx = x * x;\n let t = 1.0 / (1.0 + 0.3275911 * ax);\n let poly = 0.254829592\n + t * (-0.284496736\n + t * (1.421413741\n + t * (-1.453152027\n + t * 1.061405429)));\n let tau = t * exp(-xx) * poly;\n output[idx] = sign(x) * (1.0 - tau);\n}\n"; export declare const expShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = exp(input[idx]);\n}\n"; export declare const negShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = -input[idx];\n}\n"; export declare const sigmoidShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = 1.0 / (1.0 + exp(-input[idx]));\n}\n"; export declare const tanhShaderExtra = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params { size: u32 }\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n let v = clamp(input[idx], -15.0, 15.0);\n output[idx] = tanh(v);\n}\n"; export declare const reduceMeanLastShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n batch: u32,\n D: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n// One thread per batch row, 64 threads per workgroup.\n// D=120 in PP-OCRv6 SVTR: serial scan is faster than parallel reduction\n// because barrier overhead >> compute for such a small D.\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let b = gid.x;\n if (b >= params.batch) { return; }\n let base = b * params.D;\n var s = 0.0;\n for (var i = 0u; i < params.D; i++) {\n s += input[base + i];\n }\n output[b] = s / f32(params.D);\n}\n"; export declare const gatherShader = "\n@group(0) @binding(0) var data: array;\n@group(0) @binding(1) var indices: array;\n\n@group(0) @binding(2) var output: array;\n\nstruct Params {\n out_size: u32,\n axis: u32,\n data_rank: u32,\n}\n@group(0) @binding(3) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.out_size) { return; }\n // We treat indices as a 1-D lookup of size data_rank, regardless of axis.\n // The index value determines the source dim coordinate.\n let idx_val = u32(round(indices[idx]));\n output[idx] = data[idx_val];\n}\n"; export declare const reduceMeanGenericShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n out_size: u32,\n O0: u32, O1: u32, O2: u32, O3: u32,\n len0: u32, len1: u32, len2: u32, len3: u32,\n S0: u32, S1: u32, S2: u32, S3: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.out_size) { return; }\n\n let q3 = params.O3;\n let q2 = params.O2;\n let q1 = params.O1;\n let o3 = idx % q3;\n let o2 = (idx / q3) % q2;\n let o1 = (idx / (q3 * q2)) % q1;\n let o0 = idx / (q3 * q2 * q1);\n\n let s3 = select(0u, o3, params.O3 > 1u);\n let s2 = select(0u, o2, params.O2 > 1u);\n let s1 = select(0u, o1, params.O1 > 1u);\n let s0 = select(0u, o0, params.O0 > 1u);\n\n let base = s0 * params.S0 + s1 * params.S1 + s2 * params.S2 + s3 * params.S3;\n\n // Effective reduced lengths (only used when corresponding output dim is 1).\n let n0 = select(1u, params.len0, params.O0 == 1u);\n let n1 = select(1u, params.len1, params.O1 == 1u);\n let n2 = select(1u, params.len2, params.O2 == 1u);\n let n3 = select(1u, params.len3, params.O3 == 1u);\n\n let total = n0 * n1 * n2 * n3;\n let stride_r0 = select(0u, params.S0, params.O0 == 1u);\n let stride_r1 = select(0u, params.S1, params.O1 == 1u);\n let stride_r2 = select(0u, params.S2, params.O2 == 1u);\n let stride_r3 = select(0u, params.S3, params.O3 == 1u);\n\n var sum = 0.0;\n for (var i0 = 0u; i0 < n0; i0++) {\n let off0 = i0 * stride_r0;\n for (var i1 = 0u; i1 < n1; i1++) {\n let off1 = off0 + i1 * stride_r1;\n for (var i2 = 0u; i2 < n2; i2++) {\n let off2 = off1 + i2 * stride_r2;\n for (var i3 = 0u; i3 < n3; i3++) {\n sum += input[base + off2 + i3 * stride_r3];\n }\n }\n }\n }\n\n let count = f32(max(total, 1u));\n output[idx] = sum / count;\n}\n"; export declare const subLastBroadcastShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var scalars: array; // [batch]\n@group(0) @binding(2) var output: array;\n\nstruct Params {\n batch: u32,\n D: u32,\n}\n@group(0) @binding(3) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let bi = gid.x;\n let di = gid.y;\n if (bi >= params.batch || di >= params.D) { return; }\n let idx = bi * params.D + di;\n output[idx] = input[idx] - scalars[bi];\n}\n"; export declare const addLastBroadcastShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var scalars: array; // [batch]\n@group(0) @binding(2) var output: array;\n\nstruct Params {\n batch: u32,\n D: u32,\n}\n@group(0) @binding(3) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let bi = gid.x;\n let di = gid.y;\n if (bi >= params.batch || di >= params.D) { return; }\n let idx = bi * params.D + di;\n output[idx] = input[idx] + scalars[bi];\n}\n"; export declare const mulLastBroadcastShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var scalars: array; // [batch]\n@group(0) @binding(2) var output: array;\n\nstruct Params {\n batch: u32,\n D: u32,\n}\n@group(0) @binding(3) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let bi = gid.x;\n let di = gid.y;\n if (bi >= params.batch || di >= params.D) { return; }\n let idx = bi * params.D + di;\n output[idx] = input[idx] * scalars[bi];\n}\n"; export declare const divLastBroadcastShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var scalars: array; // [batch]\n@group(0) @binding(2) var output: array;\n\nstruct Params {\n batch: u32,\n D: u32,\n}\n@group(0) @binding(3) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let bi = gid.x;\n let di = gid.y;\n if (bi >= params.batch || di >= params.D) { return; }\n let idx = bi * params.D + di;\n output[idx] = input[idx] / scalars[bi];\n}\n"; export declare const affineScalarBroadcastShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var norm_a: array; // mean or std [B]\n@group(0) @binding(2) var perFeat: array; // \u03B3 or \u03B2 [D]\n@group(0) @binding(3) var output: array;\n\nstruct Params {\n batch: u32,\n D: u32,\n}\n@group(0) @binding(4) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let b = gid.x;\n let d = gid.y;\n if (b >= params.batch || d >= params.D) { return; }\n output[b * params.D + d] = a[b * params.D + d] * norm_a[b] * perFeat[d];\n}\n"; export declare const affineScalarBiasBroadcastShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var norm_a: array; // mean or std [B]\n@group(0) @binding(2) var perFeat: array; // \u03B3 [D]\n@group(0) @binding(3) var bias: array; // \u03B2 [D]\n@group(0) @binding(4) var output: array;\n\nstruct Params {\n batch: u32,\n D: u32,\n}\n@group(0) @binding(5) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let b = gid.x;\n let d = gid.y;\n if (b >= params.batch || d >= params.D) { return; }\n output[b * params.D + d] = a[b * params.D + d] * norm_a[b] * perFeat[d] + bias[d];\n}\n"; export declare const genericReshapeShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n total: u32,\n // output dims\n O0: u32, O1: u32, O2: u32, O3: u32,\n // input dims\n I0: u32, I1: u32, I2: u32, I3: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.total) { return; }\n\n let q1 = params.O1 * params.O2 * params.O3;\n let q2 = params.O2 * params.O3;\n let a = idx / q1;\n let r = idx % q1;\n let b = r / q2;\n let r2 = r % q2;\n let c = r2 / params.O3;\n let d = r2 % params.O3;\n\n // Map flat \u2192 flat with input's contiguous strides.\n let in_idx = a * params.I1 * params.I2 * params.I3\n + b * params.I2 * params.I3\n + c * params.I3\n + d;\n output[idx] = input[in_idx];\n}\n"; export declare const concatChannelsFullShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var b: array;\n@group(0) @binding(2) var output: array;\n\nstruct Params {\n N: u32,\n Ca: u32,\n Cb: u32,\n H: u32,\n W: u32,\n}\n@group(0) @binding(3) var params: Params;\n\n@compute @workgroup_size(64, 4)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let n = gid.z;\n let c = gid.y;\n let xhw = gid.x;\n let H = params.H;\n let W = params.W;\n let hwx = H * W;\n if (n >= params.N) { return; }\n if (xhw >= hwx) { return; }\n let c_total = params.Ca + params.Cb;\n if (c >= c_total) { return; }\n\n if (c < params.Ca) {\n let in_idx = (n * params.Ca + c) * hwx + xhw;\n let out_idx = (n * c_total + c) * hwx + xhw;\n output[out_idx] = a[in_idx];\n } else {\n let cb = c - params.Ca;\n let in_idx = (n * params.Cb + cb) * hwx + xhw;\n let out_idx = (n * c_total + c) * hwx + xhw;\n output[out_idx] = b[in_idx];\n }\n}\n"; export declare const divScalarLastShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var output: array;\nstruct Params {\n size: u32,\n _pad: u32, _pad2: u32,\n scalar: f32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = a[idx] / params.scalar;\n}\n"; export declare const mulScalarLastShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var output: array;\nstruct Params {\n size: u32,\n _pad: u32, _pad2: u32,\n scalar: f32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = a[idx] * params.scalar;\n}\n"; export declare const addScalarLastShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var output: array;\nstruct Params {\n size: u32,\n _pad: u32, _pad2: u32,\n scalar: f32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = a[idx] + params.scalar;\n}\n"; export declare const subScalarLastShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var output: array;\nstruct Params {\n size: u32,\n _pad: u32, _pad2: u32,\n scalar: f32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(256)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.size) { return; }\n output[idx] = a[idx] - params.scalar;\n}\n"; export declare const concatAxisZeroShader = "\n@group(0) @binding(0) var a: array;\n@group(0) @binding(1) var b: array;\n@group(0) @binding(2) var output: array;\n\nstruct Params {\n a_size: u32,\n total: u32,\n}\n@group(0) @binding(3) var params: Params;\n\n@compute @workgroup_size(64)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let idx = gid.x;\n if (idx >= params.total) { return; }\n if (idx < params.a_size) {\n output[idx] = a[idx];\n } else {\n output[idx] = b[idx - params.a_size];\n }\n}\n"; export declare const resizeNearestShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n N: u32, C: u32,\n H_in: u32, W_in: u32,\n H_out: u32, W_out: u32,\n}\n@group(0) @binding(2) var params: Params;\n\n@compute @workgroup_size(8, 8, 1)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let n = gid.z;\n let c = gid.y;\n let h_out = gid.x / params.W_out;\n let w_out = gid.x % params.W_out;\n if (h_out >= params.H_out) { return; }\n\n // Nearest neighbor\n let h_in = (h_out * params.H_in) / params.H_out;\n let w_in = (w_out * params.W_in) / params.W_out;\n\n let in_idx = n * params.C * params.H_in * params.W_in + c * params.H_in * params.W_in + h_in * params.W_in + w_in;\n let out_idx = n * params.C * params.H_out * params.W_out + c * params.H_out * params.W_out + h_out * params.W_out + w_out;\n output[out_idx] = input[in_idx];\n}\n"; export declare const resizeBilinearShader = "\n@group(0) @binding(0) var input: array;\n@group(0) @binding(1) var output: array;\n\nstruct Params {\n N: u32, C: u32,\n H_in: u32, W_in: u32,\n H_out: u32, W_out: u32,\n}\n@group(0) @binding(2) var params: Params;\n\nfn get_pixel(n: u32, c: u32, h: i32, w: i32) -> f32 {\n let h_clamp = max(0, min(i32(params.H_in) - 1, h));\n let w_clamp = max(0, min(i32(params.W_in) - 1, w));\n return input[n * params.C * params.H_in * params.W_in + c * params.H_in * params.W_in + u32(h_clamp) * params.W_in + u32(w_clamp)];\n}\n\n@compute @workgroup_size(8, 8, 1)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let n = gid.z;\n let c = gid.y;\n let h_out = gid.x / params.W_out;\n let w_out = gid.x % params.W_out;\n if (h_out >= params.H_out) { return; }\n\n // Bilinear\n let h_scale = f32(params.H_in) / f32(params.H_out);\n let w_scale = f32(params.W_in) / f32(params.W_out);\n let h_f = (f32(h_out) + 0.5) * h_scale - 0.5;\n let w_f = (f32(w_out) + 0.5) * w_scale - 0.5;\n let h0 = i32(floor(h_f));\n let w0 = i32(floor(w_f));\n let h1 = h0 + 1;\n let w1 = w0 + 1;\n let dh = h_f - f32(h0);\n let dw = w_f - f32(w0);\n\n let v00 = get_pixel(n, c, h0, w0);\n let v01 = get_pixel(n, c, h0, w1);\n let v10 = get_pixel(n, c, h1, w0);\n let v11 = get_pixel(n, c, h1, w1);\n\n let v0 = v00 * (1.0 - dw) + v01 * dw;\n let v1 = v10 * (1.0 - dw) + v11 * dw;\n let val = v0 * (1.0 - dh) + v1 * dh;\n\n let out_idx = n * params.C * params.H_out * params.W_out + c * params.H_out * params.W_out + h_out * params.W_out + w_out;\n output[out_idx] = val;\n}\n"; export declare const convTranspose2dShader = "\n@group(0) @binding(0) var input: array; // [N, Cin, H, W]\n@group(0) @binding(1) var weight: array; // [Cin, Cout, kH, kW]\n@group(0) @binding(2) var bias: array; // [Cout]\n@group(0) @binding(3) var output: array; // [N, Cout, Hout, Wout]\n\nstruct Params {\n N: u32, Cin: u32, Cout: u32,\n H: u32, W: u32,\n Hout: u32, Wout: u32,\n kH: u32, kW: u32,\n strideH: u32, strideW: u32,\n padTop: u32, padLeft: u32,\n use_bias: u32,\n}\n@group(0) @binding(4) var params: Params;\n\n@compute @workgroup_size(64, 4, 1)\nfn main(@builtin(global_invocation_id) gid: vec3) {\n let n = gid.z;\n let co = gid.y;\n let h_out = gid.x / params.Wout;\n let w_out = gid.x % params.Wout;\n if (h_out >= params.Hout) { return; }\n\n var sum = 0.0;\n let h_in_start = max(i32(0), (i32(h_out) + i32(params.padTop) - i32(params.kH) + 1) / i32(params.strideH));\n let w_in_start = max(i32(0), (i32(w_out) + i32(params.padLeft) - i32(params.kW) + 1) / i32(params.strideW));\n let h_in_end = min(i32(params.H), (i32(h_out) + i32(params.padTop)) / i32(params.strideH) + 1);\n let w_in_end = min(i32(params.W), (i32(w_out) + i32(params.padLeft)) / i32(params.strideW) + 1);\n\n for (var hi = h_in_start; hi < h_in_end; hi = hi + 1) {\n for (var wi = w_in_start; wi < w_in_end; wi = wi + 1) {\n for (var ci = 0u; ci < params.Cin; ci = ci + 1) {\n let kh = u32(i32(h_out) + i32(params.padTop) - hi * i32(params.strideH));\n let kw = u32(i32(w_out) + i32(params.padLeft) - wi * i32(params.strideW));\n if (kh >= params.kH || kw >= params.kW) { continue; }\n let in_val = input[n * params.Cin * params.H * params.W + ci * params.H * params.W + u32(hi) * params.W + u32(wi)];\n // ONNX ConvTranspose: weights are stored as [Cin, Cout, kH, kW] and\n // used directly (NO kernel flip \u2014 that's a PyTorch internal detail\n // already baked into the exported weights).\n let w_val = weight[ci * params.Cout * params.kH * params.kW + co * params.kH * params.kW + kh * params.kW + kw];\n sum += in_val * w_val;\n }\n }\n }\n\n let out_idx = n * params.Cout * params.Hout * params.Wout + co * params.Hout * params.Wout + h_out * params.Wout + w_out;\n if (params.use_bias != 0u) {\n output[out_idx] = sum + bias[co];\n } else {\n output[out_idx] = sum;\n }\n}\n"; //# sourceMappingURL=shaders-ppocr.d.ts.map