whispercpp 1.3.7 → 1.3.8
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- checksums.yaml +4 -4
- data/README.md +5 -4
- data/ext/options.rb +1 -1
- data/ext/ruby_whisper.c +0 -1
- data/ext/ruby_whisper.h +7 -1
- data/ext/ruby_whisper_context.c +50 -1
- data/ext/ruby_whisper_log_settable.h +1 -2
- data/ext/ruby_whisper_params.c +9 -8
- data/ext/ruby_whisper_transcribe.cpp +0 -19
- data/ext/ruby_whisper_vad_context.c +30 -10
- data/ext/ruby_whisper_vad_context_detect.cpp +8 -9
- data/ext/ruby_whisper_vad_params.c +4 -4
- data/ext/ruby_whisper_vad_segment.c +2 -2
- data/ext/sources/CMakeLists.txt +2 -1
- data/ext/sources/cmake/parakeet.pc.in +2 -2
- data/ext/sources/cmake/whisper.pc.in +2 -2
- data/ext/sources/examples/cli/cli.cpp +9 -1
- data/ext/sources/examples/common-ggml.cpp +2 -0
- data/ext/sources/examples/vad-speech-segments/speech.cpp +3 -2
- data/ext/sources/ggml/CMakeLists.txt +3 -4
- data/ext/sources/ggml/include/ggml-cuda.h +0 -3
- data/ext/sources/ggml/include/ggml-sycl.h +8 -0
- data/ext/sources/ggml/include/ggml.h +3 -1
- data/ext/sources/ggml/src/CMakeLists.txt +8 -1
- data/ext/sources/ggml/src/ggml-backend-meta.cpp +7 -4
- data/ext/sources/ggml/src/ggml-common.h +13 -2
- data/ext/sources/ggml/src/ggml-cpu/CMakeLists.txt +1 -1
- data/ext/sources/ggml/src/ggml-cpu/amx/mmq.cpp +5 -6
- data/ext/sources/ggml/src/ggml-cpu/arch/arm/quants.c +78 -4
- data/ext/sources/ggml/src/ggml-cpu/arch/x86/quants.c +142 -4
- data/ext/sources/ggml/src/ggml-cpu/arch-fallback.h +7 -2
- data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.c +14 -0
- data/ext/sources/ggml/src/ggml-cpu/llamafile/sgemm.cpp +26 -19
- data/ext/sources/ggml/src/ggml-cpu/ops.cpp +129 -46
- data/ext/sources/ggml/src/ggml-cpu/quants.c +51 -0
- data/ext/sources/ggml/src/ggml-cpu/quants.h +3 -0
- data/ext/sources/ggml/src/ggml-cpu/simd-gemm.h +1 -1
- data/ext/sources/ggml/src/ggml-cpu/simd-mappings.h +11 -0
- data/ext/sources/ggml/src/ggml-cpu/vec.cpp +2 -2
- data/ext/sources/ggml/src/ggml-cuda/binbcast.cu +90 -46
- data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cu +81 -0
- data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cuh +3 -0
- data/ext/sources/ggml/src/ggml-cuda/common.cuh +4 -0
- data/ext/sources/ggml/src/ggml-cuda/concat.cu +33 -21
- data/ext/sources/ggml/src/ggml-cuda/conv-transpose-1d.cu +14 -12
- data/ext/sources/ggml/src/ggml-cuda/convert.cu +86 -34
- data/ext/sources/ggml/src/ggml-cuda/cpy.cu +80 -29
- data/ext/sources/ggml/src/ggml-cuda/fattn-common.cuh +9 -5
- data/ext/sources/ggml/src/ggml-cuda/fattn-mma-f16.cuh +4 -0
- data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cuh +9 -5
- data/ext/sources/ggml/src/ggml-cuda/fattn.cu +27 -21
- data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cu +40 -25
- data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cuh +10 -0
- data/ext/sources/ggml/src/ggml-cuda/getrows.cu +15 -12
- data/ext/sources/ggml/src/ggml-cuda/ggml-cuda.cu +718 -1248
- data/ext/sources/ggml/src/ggml-cuda/mmq.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/mmvq.cu +77 -40
- data/ext/sources/ggml/src/ggml-cuda/out-prod.cu +55 -12
- data/ext/sources/ggml/src/ggml-cuda/set-rows.cu +64 -4
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_2.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_32-ncols2_2.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_2.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_2.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/topk-moe.cu +7 -1
- data/ext/sources/ggml/src/ggml-cuda/vendors/hip.h +1 -0
- data/ext/sources/ggml/src/ggml-cuda/vendors/musa.h +1 -0
- data/ext/sources/ggml/src/ggml-hexagon/CMakeLists.txt +0 -5
- data/ext/sources/ggml/src/ggml-hexagon/ggml-hexagon.cpp +1634 -1293
- data/ext/sources/ggml/src/ggml-hexagon/htp/CMakeLists.txt +11 -40
- data/ext/sources/ggml/src/ggml-hexagon/htp/cmake-toolchain.cmake +13 -15
- data/ext/sources/ggml/src/ggml-hexagon/htp/concat-ops.c +1 -1
- data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +1749 -399
- data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +303 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-common.h +80 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.h +26 -23
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-profile.h +64 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-utils.h +1 -83
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +555 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h +1303 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.c +9 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.h +27 -4
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-utils.h +59 -37
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ctx.h +11 -3
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ops.h +52 -12
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp-vtcm.h +19 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp_iface.idl +2 -1
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-base.h +14 -30
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-exp.h +39 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h +232 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h +1511 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h +1200 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h +39 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/main.c +127 -32
- data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.c +3023 -4425
- data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.h +650 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/rope-ops.c +48 -13
- data/ext/sources/ggml/src/ggml-hexagon/htp/ssm-conv.c +10 -9
- data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.c +15 -3
- data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.h +8 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp-opnode.h +168 -50
- data/ext/sources/ggml/src/ggml-hexagon/libggml-htp.inf +0 -4
- data/ext/sources/ggml/src/ggml-hip/CMakeLists.txt +5 -0
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.cpp +69 -5
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.h +4 -1
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.m +27 -6
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-impl.h +38 -0
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.cpp +132 -2
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.h +2 -0
- data/ext/sources/ggml/src/ggml-metal/ggml-metal.metal +345 -87
- data/ext/sources/ggml/src/ggml-opencl/CMakeLists.txt +13 -0
- data/ext/sources/ggml/src/ggml-opencl/fa_tune.h +92 -0
- data/ext/sources/ggml/src/ggml-opencl/ggml-opencl.cpp +4060 -357
- data/ext/sources/ggml/src/ggml-opencl/kernels/cvt.cl +198 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f16.cl +81 -41
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32.cl +88 -39
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl +1995 -96
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q4_0.cl +1615 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q8_0.cl +1486 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_pre_f16.cl +156 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_mxfp4_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_0_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_1_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_k_f32_ns.cl +71 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_0_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_1_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_k_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q6_k_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q1_0_f32.cl +94 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q1_0_f32.cl +121 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q8_0_f32.cl +1 -1
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q1_0_f32_l4_lm.cl +156 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_l4.cl +1149 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32.cl +141 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32_flat.cl +190 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/norm.cl +5 -2
- data/ext/sources/ggml/src/ggml-opencl/kernels/set_rows.cl +500 -0
- data/ext/sources/ggml/src/ggml-opencl/libdl.h +79 -0
- data/ext/sources/ggml/src/ggml-openvino/.clang-format +0 -5
- data/ext/sources/ggml/src/ggml-openvino/CMakeLists.txt +2 -4
- data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.cpp +733 -130
- data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.h +76 -23
- data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.cpp +57 -3
- data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.h +29 -8
- data/ext/sources/ggml/src/ggml-openvino/ggml-openvino.cpp +307 -59
- data/ext/sources/ggml/src/ggml-openvino/ggml-quants.cpp +66 -0
- data/ext/sources/ggml/src/ggml-openvino/ggml-quants.h +10 -4
- data/ext/sources/ggml/src/ggml-openvino/openvino/decoder.h +56 -16
- data/ext/sources/ggml/src/ggml-openvino/openvino/frontend.h +1 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/input_model.h +4 -4
- data/ext/sources/ggml/src/ggml-openvino/openvino/node_context.h +94 -37
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/add_id.cpp +76 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/argsort.cpp +47 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/clamp.cpp +33 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/concat.cpp +48 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/cont.cpp +8 -16
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/cpy.cpp +14 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/div.cpp +146 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp +108 -21
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp +282 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.hpp +65 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/get_rows.cpp +2 -9
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_geglu.cpp +21 -7
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp +41 -8
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/im2col.cpp +120 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/l2_norm.cpp +44 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/mul_mat_id.cpp +226 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/mulmat.cpp +19 -9
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/norm.cpp +58 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/pad.cpp +95 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/permute.cpp +58 -13
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/repeat.cpp +74 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/reshape.cpp +13 -6
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/rms_norm.cpp +1 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/rope.cpp +134 -38
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/set_rows.cpp +3 -3
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/softmax.cpp +126 -49
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/ssm_conv.cpp +59 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/sum_rows.cpp +27 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/transpose.cpp +32 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_silu.cpp +1 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_softplus.cpp +38 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/view.cpp +90 -25
- data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.cpp +41 -23
- data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.h +18 -5
- data/ext/sources/ggml/src/ggml-openvino/openvino/pass/mark_decompression_convert_constant_folding.h +1 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.cpp +43 -40
- data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.h +5 -4
- data/ext/sources/ggml/src/ggml-openvino/openvino/utils.cpp +548 -3
- data/ext/sources/ggml/src/ggml-openvino/openvino/utils.h +28 -26
- data/ext/sources/ggml/src/ggml-openvino/utils.cpp +383 -94
- data/ext/sources/ggml/src/ggml-openvino/utils.h +11 -8
- data/ext/sources/ggml/src/ggml-quants.c +76 -0
- data/ext/sources/ggml/src/ggml-quants.h +3 -0
- data/ext/sources/ggml/src/ggml-sycl/CMakeLists.txt +5 -5
- data/ext/sources/ggml/src/ggml-sycl/backend.hpp +2 -0
- data/ext/sources/ggml/src/ggml-sycl/binbcast.cpp +12 -0
- data/ext/sources/ggml/src/ggml-sycl/col2im-1d.cpp +102 -0
- data/ext/sources/ggml/src/ggml-sycl/col2im-1d.hpp +8 -0
- data/ext/sources/ggml/src/ggml-sycl/common.cpp +6 -8
- data/ext/sources/ggml/src/ggml-sycl/common.hpp +19 -2
- data/ext/sources/ggml/src/ggml-sycl/concat.cpp +21 -1
- data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.cpp +158 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.hpp +10 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.cpp +125 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.hpp +10 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d.cpp +150 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d.hpp +10 -0
- data/ext/sources/ggml/src/ggml-sycl/conv3d.cpp +224 -0
- data/ext/sources/ggml/src/ggml-sycl/conv3d.hpp +8 -0
- data/ext/sources/ggml/src/ggml-sycl/convert.cpp +6 -0
- data/ext/sources/ggml/src/ggml-sycl/cpy.cpp +706 -0
- data/ext/sources/ggml/src/ggml-sycl/cpy.hpp +281 -0
- data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.cpp +255 -0
- data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.hpp +7 -0
- data/ext/sources/ggml/src/ggml-sycl/dequantize.hpp +15 -0
- data/ext/sources/ggml/src/ggml-sycl/dmmv.cpp +492 -319
- data/ext/sources/ggml/src/ggml-sycl/dpct/helper.hpp +15 -7
- data/ext/sources/ggml/src/ggml-sycl/element_wise.cpp +215 -115
- data/ext/sources/ggml/src/ggml-sycl/element_wise.hpp +2 -0
- data/ext/sources/ggml/src/ggml-sycl/ggml-sycl.cpp +1006 -336
- data/ext/sources/ggml/src/ggml-sycl/mmvq.cpp +252 -67
- data/ext/sources/ggml/src/ggml-sycl/mmvq.hpp +17 -0
- data/ext/sources/ggml/src/ggml-sycl/norm.cpp +103 -49
- data/ext/sources/ggml/src/ggml-sycl/outprod.cpp +45 -9
- data/ext/sources/ggml/src/ggml-sycl/pool.cpp +185 -0
- data/ext/sources/ggml/src/ggml-sycl/pool.hpp +22 -0
- data/ext/sources/ggml/src/ggml-sycl/presets.hpp +3 -1
- data/ext/sources/ggml/src/ggml-sycl/set_rows.cpp +10 -2
- data/ext/sources/ggml/src/ggml-sycl/softmax.cpp +9 -10
- data/ext/sources/ggml/src/ggml-sycl/vecdotq.hpp +35 -0
- data/ext/sources/ggml/src/ggml-vulkan/CMakeLists.txt +5 -0
- data/ext/sources/ggml/src/ggml-vulkan/ggml-vulkan.cpp +833 -215
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/col2im_1d.comp +61 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp +1 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp +431 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/diag.comp +3 -3
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +1 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +1 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl +21 -19
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/get_rows_back.comp +25 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl +23 -4
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_main.glsl +14 -18
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/l2_norm.comp +4 -7
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +21 -24
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +31 -23
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +6 -5
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +84 -67
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/norm.comp +10 -10
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/repeat_back.comp +3 -3
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/roll.comp +3 -3
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tri.comp +3 -3
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/unary.comp +168 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +121 -74
- data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +26 -19
- data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu.cpp +31 -36
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +16 -2
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl +7 -7
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl +21 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +439 -320
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_vec.wgsl +2 -2
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl +45 -39
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl +586 -465
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl +63 -69
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl +14 -9
- data/ext/sources/ggml/src/ggml.c +36 -14
- data/ext/sources/include/whisper.h +21 -0
- data/ext/sources/src/whisper.cpp +164 -14
- data/lib/whisper/log_settable.rb +5 -8
- data/lib/whisper/model/uri.rb +0 -7
- data/sig/whisper.rbs +6 -0
- data/test/test_vad.rb +9 -0
- data/test/test_vad_context.rb +2 -2
- data/whispercpp.gemspec +1 -1
- metadata +62 -37
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-flash-attn-ops.c +0 -1878
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-matmul-ops.c +0 -2066
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-ops.c +0 -6
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-ops.h +0 -88
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-profile.h +0 -34
- data/ext/sources/ggml/src/ggml-hexagon/htp/vtcm-utils.h +0 -16
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_gelu.cpp +0 -25
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/abs.comp +0 -21
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ceil.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/clamp.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/cos.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/elu.comp +0 -27
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/exp.comp +0 -20
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/floor.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu.comp +0 -25
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_erf.comp +0 -39
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_quick.comp +0 -23
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardsigmoid.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardswish.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/leaky_relu.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/neg.comp +0 -20
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/relu.comp +0 -21
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/round.comp +0 -29
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sgn.comp +0 -21
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sigmoid.comp +0 -20
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/silu.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sin.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/softplus.comp +0 -23
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sqrt.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/square.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/step.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tanh.comp +0 -20
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/trunc.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/xielu.comp +0 -35
|
@@ -32,8 +32,8 @@ fn inner_dot(src0_val: SRC0_TYPE, src1_val: SRC1_TYPE) -> f32 {
|
|
|
32
32
|
#endif
|
|
33
33
|
|
|
34
34
|
#ifdef MUL_ACC_FLOAT
|
|
35
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
36
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
35
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
36
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
37
37
|
|
|
38
38
|
let k_vec = params.k / VEC_SIZE;
|
|
39
39
|
let src1_idx_base_vec = src1_idx_base / VEC_SIZE;
|
|
@@ -41,12 +41,18 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
41
41
|
// Each thread walks K, loads from the vector, and updates
|
|
42
42
|
// a small block of output rows held in registers.
|
|
43
43
|
for (var k = thread_id; k < k_vec; k += WG_SIZE) {
|
|
44
|
-
|
|
44
|
+
var x_vals: array<SRC1_TYPE, NUM_COLS>;
|
|
45
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
46
|
+
x_vals[col] = src1[src1_idx_base_vec + col * (params.stride_11 / VEC_SIZE) + k];
|
|
47
|
+
}
|
|
45
48
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
46
49
|
let output_row = row_base + row;
|
|
47
50
|
if (output_row < params.m) {
|
|
48
51
|
let src0_idx = (src0_batch_offset + output_row * params.stride_01) / VEC_SIZE + k;
|
|
49
|
-
|
|
52
|
+
let w = src0[src0_idx];
|
|
53
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
54
|
+
acc[col][row] += inner_dot(w, x_vals[col]);
|
|
55
|
+
}
|
|
50
56
|
}
|
|
51
57
|
}
|
|
52
58
|
}
|
|
@@ -60,30 +66,33 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
60
66
|
#define BLOCK_SIZE_BYTES 18
|
|
61
67
|
#define THREADS_PER_BLOCK 16
|
|
62
68
|
#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
|
|
63
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
64
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
69
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
70
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
65
71
|
|
|
66
72
|
let num_blocks = params.k / BLOCK_SIZE;
|
|
67
73
|
let thread_within_block = thread_id % THREADS_PER_BLOCK;
|
|
68
74
|
for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
|
|
69
75
|
let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * ELEMS_PER_THREAD;
|
|
70
|
-
var x_block: array<f32, ELEMS_PER_THREAD>;
|
|
71
|
-
for (var
|
|
72
|
-
|
|
76
|
+
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
|
77
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
78
|
+
for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
|
|
79
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
80
|
+
}
|
|
73
81
|
}
|
|
74
|
-
|
|
75
82
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
76
83
|
let output_row = row_base + row;
|
|
77
84
|
if (output_row < params.m) {
|
|
78
85
|
let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
|
|
79
86
|
let d = f32(load_f16_at_src0(block_byte_base));
|
|
80
87
|
let q_byte = load_u32_at_src0(block_byte_base + 2u + thread_within_block) & 0xFFu;
|
|
81
|
-
var
|
|
82
|
-
|
|
83
|
-
|
|
84
|
-
|
|
88
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
89
|
+
var row_sum = 0.0;
|
|
90
|
+
for (var bit = 0u; bit < 8u; bit++) {
|
|
91
|
+
let w = select(-d, d, ((q_byte >> bit) & 1u) != 0u);
|
|
92
|
+
row_sum += w * x_block[col][bit];
|
|
93
|
+
}
|
|
94
|
+
acc[col][row] += row_sum;
|
|
85
95
|
}
|
|
86
|
-
acc[row] += row_sum;
|
|
87
96
|
}
|
|
88
97
|
}
|
|
89
98
|
}
|
|
@@ -97,35 +106,37 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
97
106
|
#define BLOCK_SIZE_BYTES 18
|
|
98
107
|
#define THREADS_PER_BLOCK 4
|
|
99
108
|
#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
|
|
100
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
101
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
109
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
110
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
102
111
|
|
|
103
112
|
let num_blocks = params.k / BLOCK_SIZE;
|
|
104
113
|
let thread_within_block = thread_id % 4;
|
|
105
114
|
for (var block = thread_id/THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE/THREADS_PER_BLOCK) {
|
|
106
115
|
let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
|
|
107
|
-
var x_block: array<f32, ELEMS_PER_THREAD>;
|
|
108
|
-
for (var
|
|
109
|
-
|
|
110
|
-
|
|
116
|
+
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
|
117
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
118
|
+
for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
|
|
119
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
120
|
+
x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
|
|
121
|
+
}
|
|
111
122
|
}
|
|
112
|
-
|
|
113
123
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
114
124
|
let output_row = row_base + row;
|
|
115
125
|
if (output_row < params.m) {
|
|
116
126
|
let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
|
|
117
127
|
let d = f32(load_f16_at_src0(block_byte_base));
|
|
118
|
-
var row_sum = 0.0;
|
|
119
|
-
|
|
120
128
|
let q_packed = load_u32_at_src0(block_byte_base + 2u + 4u * thread_within_block);
|
|
121
|
-
for (var
|
|
122
|
-
|
|
123
|
-
|
|
124
|
-
|
|
125
|
-
|
|
126
|
-
|
|
129
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
130
|
+
var row_sum = 0.0;
|
|
131
|
+
for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
|
|
132
|
+
let q_byte = get_byte(q_packed, byte_idx);
|
|
133
|
+
let q_lo = (f32(q_byte & 0xFu) - 8.0) * d;
|
|
134
|
+
let q_hi = (f32((q_byte >> 4u) & 0xFu) - 8.0) * d;
|
|
135
|
+
row_sum += q_lo * x_block[col][byte_idx];
|
|
136
|
+
row_sum += q_hi * x_block[col][byte_idx + 4u];
|
|
137
|
+
}
|
|
138
|
+
acc[col][row] += row_sum;
|
|
127
139
|
}
|
|
128
|
-
acc[row] += row_sum;
|
|
129
140
|
}
|
|
130
141
|
}
|
|
131
142
|
}
|
|
@@ -139,36 +150,38 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
139
150
|
#define BLOCK_SIZE_BYTES 20
|
|
140
151
|
#define THREADS_PER_BLOCK 4
|
|
141
152
|
#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
|
|
142
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
143
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
153
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
154
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
144
155
|
|
|
145
156
|
let num_blocks = params.k / BLOCK_SIZE;
|
|
146
157
|
let thread_within_block = thread_id % THREADS_PER_BLOCK;
|
|
147
158
|
for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
|
|
148
159
|
let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
|
|
149
|
-
var x_block: array<f32, ELEMS_PER_THREAD>;
|
|
150
|
-
for (var
|
|
151
|
-
|
|
152
|
-
|
|
160
|
+
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
|
161
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
162
|
+
for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
|
|
163
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
164
|
+
x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
|
|
165
|
+
}
|
|
153
166
|
}
|
|
154
|
-
|
|
155
167
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
156
168
|
let output_row = row_base + row;
|
|
157
169
|
if (output_row < params.m) {
|
|
158
170
|
let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
|
|
159
171
|
let d = f32(load_f16_at_src0(block_byte_base));
|
|
160
172
|
let m = f32(load_f16_at_src0(block_byte_base + 2u));
|
|
161
|
-
var row_sum = 0.0;
|
|
162
|
-
|
|
163
173
|
let q_packed = load_u32_at_src0(block_byte_base + 4u + 4u * thread_within_block);
|
|
164
|
-
for (var
|
|
165
|
-
|
|
166
|
-
|
|
167
|
-
|
|
168
|
-
|
|
169
|
-
|
|
174
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
175
|
+
var row_sum = 0.0;
|
|
176
|
+
for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
|
|
177
|
+
let q_byte = get_byte(q_packed, byte_idx);
|
|
178
|
+
let q_lo = f32(q_byte & 0xFu) * d + m;
|
|
179
|
+
let q_hi = f32((q_byte >> 4u) & 0xFu) * d + m;
|
|
180
|
+
row_sum += q_lo * x_block[col][byte_idx];
|
|
181
|
+
row_sum += q_hi * x_block[col][byte_idx + 4u];
|
|
182
|
+
}
|
|
183
|
+
acc[col][row] += row_sum;
|
|
170
184
|
}
|
|
171
|
-
acc[row] += row_sum;
|
|
172
185
|
}
|
|
173
186
|
}
|
|
174
187
|
}
|
|
@@ -182,19 +195,20 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
182
195
|
#define BLOCK_SIZE_BYTES 22
|
|
183
196
|
#define THREADS_PER_BLOCK 4
|
|
184
197
|
#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
|
|
185
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
186
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
198
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
199
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
187
200
|
|
|
188
201
|
let num_blocks = params.k / BLOCK_SIZE;
|
|
189
202
|
let thread_within_block = thread_id % THREADS_PER_BLOCK;
|
|
190
203
|
for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
|
|
191
204
|
let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
|
|
192
|
-
var x_block: array<f32, ELEMS_PER_THREAD>;
|
|
193
|
-
for (var
|
|
194
|
-
|
|
195
|
-
|
|
205
|
+
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
|
206
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
207
|
+
for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
|
|
208
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
209
|
+
x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
|
|
210
|
+
}
|
|
196
211
|
}
|
|
197
|
-
|
|
198
212
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
199
213
|
let output_row = row_base + row;
|
|
200
214
|
if (output_row < params.m) {
|
|
@@ -203,18 +217,19 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
203
217
|
let qh_packed = load_u32_at_src0(block_byte_base + 2u);
|
|
204
218
|
let q_packed = load_u32_at_src0(block_byte_base + 6u + 4u * thread_within_block);
|
|
205
219
|
let qh_shift = thread_within_block * 4u;
|
|
206
|
-
var
|
|
207
|
-
|
|
208
|
-
|
|
209
|
-
|
|
210
|
-
|
|
211
|
-
|
|
212
|
-
|
|
213
|
-
|
|
214
|
-
|
|
215
|
-
|
|
220
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
221
|
+
var row_sum = 0.0;
|
|
222
|
+
for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
|
|
223
|
+
let q_byte = get_byte(q_packed, byte_idx);
|
|
224
|
+
let qh_lo = ((qh_packed >> (qh_shift + byte_idx)) << 4u) & 0x10u;
|
|
225
|
+
let qh_hi = (qh_packed >> (qh_shift + byte_idx + 12u)) & 0x10u;
|
|
226
|
+
let q_lo = (f32((q_byte & 0xFu) | qh_lo) - 16.0) * d;
|
|
227
|
+
let q_hi = (f32(((q_byte >> 4u) & 0xFu) | qh_hi) - 16.0) * d;
|
|
228
|
+
row_sum += q_lo * x_block[col][byte_idx];
|
|
229
|
+
row_sum += q_hi * x_block[col][byte_idx + 4u];
|
|
230
|
+
}
|
|
231
|
+
acc[col][row] += row_sum;
|
|
216
232
|
}
|
|
217
|
-
acc[row] += row_sum;
|
|
218
233
|
}
|
|
219
234
|
}
|
|
220
235
|
}
|
|
@@ -228,19 +243,20 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
228
243
|
#define BLOCK_SIZE_BYTES 24
|
|
229
244
|
#define THREADS_PER_BLOCK 4
|
|
230
245
|
#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
|
|
231
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
232
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
246
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
247
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
233
248
|
|
|
234
249
|
let num_blocks = params.k / BLOCK_SIZE;
|
|
235
250
|
let thread_within_block = thread_id % THREADS_PER_BLOCK;
|
|
236
251
|
for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
|
|
237
252
|
let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
|
|
238
|
-
var x_block: array<f32, ELEMS_PER_THREAD>;
|
|
239
|
-
for (var
|
|
240
|
-
|
|
241
|
-
|
|
253
|
+
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
|
254
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
255
|
+
for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
|
|
256
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
257
|
+
x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
|
|
258
|
+
}
|
|
242
259
|
}
|
|
243
|
-
|
|
244
260
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
245
261
|
let output_row = row_base + row;
|
|
246
262
|
if (output_row < params.m) {
|
|
@@ -250,18 +266,19 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
250
266
|
let qh_packed = load_u32_at_src0(block_byte_base + 4u);
|
|
251
267
|
let q_packed = load_u32_at_src0(block_byte_base + 8u + 4u * thread_within_block);
|
|
252
268
|
let qh_shift = thread_within_block * 4u;
|
|
253
|
-
var
|
|
254
|
-
|
|
255
|
-
|
|
256
|
-
|
|
257
|
-
|
|
258
|
-
|
|
259
|
-
|
|
260
|
-
|
|
261
|
-
|
|
262
|
-
|
|
269
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
270
|
+
var row_sum = 0.0;
|
|
271
|
+
for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
|
|
272
|
+
let q_byte = get_byte(q_packed, byte_idx);
|
|
273
|
+
let qh_lo = ((qh_packed >> (qh_shift + byte_idx)) << 4u) & 0x10u;
|
|
274
|
+
let qh_hi = (qh_packed >> (qh_shift + byte_idx + 12u)) & 0x10u;
|
|
275
|
+
let q_lo = f32((q_byte & 0xFu) | qh_lo) * d + m;
|
|
276
|
+
let q_hi = f32(((q_byte >> 4u) & 0xFu) | qh_hi) * d + m;
|
|
277
|
+
row_sum += q_lo * x_block[col][byte_idx];
|
|
278
|
+
row_sum += q_hi * x_block[col][byte_idx + 4u];
|
|
279
|
+
}
|
|
280
|
+
acc[col][row] += row_sum;
|
|
263
281
|
}
|
|
264
|
-
acc[row] += row_sum;
|
|
265
282
|
}
|
|
266
283
|
}
|
|
267
284
|
}
|
|
@@ -275,33 +292,38 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
275
292
|
#define BLOCK_SIZE_BYTES 34
|
|
276
293
|
#define THREADS_PER_BLOCK 4
|
|
277
294
|
#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
|
|
278
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
279
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
295
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
296
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
280
297
|
|
|
281
298
|
let num_blocks = params.k / BLOCK_SIZE;
|
|
282
299
|
let thread_within_block = thread_id % THREADS_PER_BLOCK;
|
|
283
300
|
for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
|
|
284
301
|
let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * ELEMS_PER_THREAD;
|
|
285
|
-
var x_block: array<f32, ELEMS_PER_THREAD>;
|
|
286
|
-
for (var
|
|
287
|
-
|
|
302
|
+
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
|
303
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
304
|
+
for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
|
|
305
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
306
|
+
}
|
|
288
307
|
}
|
|
289
|
-
|
|
290
308
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
291
309
|
let output_row = row_base + row;
|
|
292
310
|
if (output_row < params.m) {
|
|
293
311
|
let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
|
|
294
312
|
let d = f32(load_f16_at_src0(block_byte_base));
|
|
295
|
-
var
|
|
296
|
-
|
|
313
|
+
var q_packed: array<u32, ELEMS_PER_THREAD / 4u>;
|
|
297
314
|
for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
|
|
298
|
-
|
|
299
|
-
|
|
300
|
-
|
|
301
|
-
|
|
315
|
+
q_packed[packed_idx] = load_u32_at_src0(block_byte_base + 2u + 4u * (thread_within_block * 2u + packed_idx));
|
|
316
|
+
}
|
|
317
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
318
|
+
var row_sum = 0.0;
|
|
319
|
+
for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
|
|
320
|
+
for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
|
|
321
|
+
let q_val = f32(get_byte_i32(q_packed[packed_idx], byte_idx)) * d;
|
|
322
|
+
row_sum += q_val * x_block[col][packed_idx * 4u + byte_idx];
|
|
323
|
+
}
|
|
302
324
|
}
|
|
325
|
+
acc[col][row] += row_sum;
|
|
303
326
|
}
|
|
304
|
-
acc[row] += row_sum;
|
|
305
327
|
}
|
|
306
328
|
}
|
|
307
329
|
}
|
|
@@ -315,34 +337,39 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
315
337
|
#define BLOCK_SIZE_BYTES 36
|
|
316
338
|
#define THREADS_PER_BLOCK 4
|
|
317
339
|
#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
|
|
318
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
319
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
340
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
341
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
320
342
|
|
|
321
343
|
let num_blocks = params.k / BLOCK_SIZE;
|
|
322
344
|
let thread_within_block = thread_id % THREADS_PER_BLOCK;
|
|
323
345
|
for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
|
|
324
346
|
let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * ELEMS_PER_THREAD;
|
|
325
|
-
var x_block: array<f32, ELEMS_PER_THREAD>;
|
|
326
|
-
for (var
|
|
327
|
-
|
|
347
|
+
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
|
348
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
349
|
+
for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
|
|
350
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
351
|
+
}
|
|
328
352
|
}
|
|
329
|
-
|
|
330
353
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
331
354
|
let output_row = row_base + row;
|
|
332
355
|
if (output_row < params.m) {
|
|
333
356
|
let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
|
|
334
357
|
let d = f32(load_f16_at_src0(block_byte_base));
|
|
335
358
|
let m = f32(load_f16_at_src0(block_byte_base + 2u));
|
|
336
|
-
var
|
|
337
|
-
|
|
359
|
+
var q_packed: array<u32, ELEMS_PER_THREAD / 4u>;
|
|
338
360
|
for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
|
|
339
|
-
|
|
340
|
-
|
|
341
|
-
|
|
342
|
-
|
|
361
|
+
q_packed[packed_idx] = load_u32_at_src0(block_byte_base + 4u + 4u * (thread_within_block * 2u + packed_idx));
|
|
362
|
+
}
|
|
363
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
364
|
+
var row_sum = 0.0;
|
|
365
|
+
for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
|
|
366
|
+
for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
|
|
367
|
+
let q_val = f32(get_byte_i32(q_packed[packed_idx], byte_idx)) * d + m;
|
|
368
|
+
row_sum += q_val * x_block[col][packed_idx * 4u + byte_idx];
|
|
369
|
+
}
|
|
343
370
|
}
|
|
371
|
+
acc[col][row] += row_sum;
|
|
344
372
|
}
|
|
345
|
-
acc[row] += row_sum;
|
|
346
373
|
}
|
|
347
374
|
}
|
|
348
375
|
}
|
|
@@ -355,8 +382,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
355
382
|
#define BLOCK_SIZE 256
|
|
356
383
|
#define BLOCK_SIZE_BYTES 84
|
|
357
384
|
#define THREADS_PER_BLOCK 16
|
|
358
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
359
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
385
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
386
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
360
387
|
|
|
361
388
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
362
389
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -379,14 +406,15 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
379
406
|
|
|
380
407
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
381
408
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
382
|
-
var x_block: array<f32, 16>;
|
|
383
|
-
for (var
|
|
384
|
-
|
|
385
|
-
|
|
386
|
-
|
|
387
|
-
|
|
409
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
410
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
411
|
+
for (var i = 0u; i < 4u; i++) {
|
|
412
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
413
|
+
x_block[col][i + 4u] = f32(src1[x_base + col * params.stride_11 + 32u + i]);
|
|
414
|
+
x_block[col][i + 8u] = f32(src1[x_base + col * params.stride_11 + 64u + i]);
|
|
415
|
+
x_block[col][i + 12u] = f32(src1[x_base + col * params.stride_11 + 96u + i]);
|
|
416
|
+
}
|
|
388
417
|
}
|
|
389
|
-
|
|
390
418
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
391
419
|
let output_row = row_base + row;
|
|
392
420
|
if (output_row < params.m) {
|
|
@@ -404,30 +432,32 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
404
432
|
let qs0 = q_u32 & 0xFFFFu;
|
|
405
433
|
let qs1 = q_u32 >> 16u;
|
|
406
434
|
|
|
407
|
-
var
|
|
408
|
-
|
|
409
|
-
|
|
410
|
-
|
|
411
|
-
|
|
412
|
-
|
|
413
|
-
|
|
414
|
-
|
|
415
|
-
|
|
416
|
-
|
|
417
|
-
|
|
418
|
-
|
|
419
|
-
|
|
420
|
-
|
|
421
|
-
|
|
422
|
-
|
|
423
|
-
|
|
424
|
-
|
|
425
|
-
|
|
426
|
-
|
|
427
|
-
|
|
428
|
-
|
|
429
|
-
|
|
430
|
-
|
|
435
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
436
|
+
var sumy = vec4<f32>(0.0, 0.0, 0.0, 0.0);
|
|
437
|
+
var acc1 = vec4<f32>(0.0, 0.0, 0.0, 0.0);
|
|
438
|
+
var acc2 = vec4<f32>(0.0, 0.0, 0.0, 0.0);
|
|
439
|
+
|
|
440
|
+
sumy[0] = x_block[col][0] + x_block[col][1] + x_block[col][2] + x_block[col][3];
|
|
441
|
+
sumy[1] = x_block[col][4] + x_block[col][5] + x_block[col][6] + x_block[col][7];
|
|
442
|
+
sumy[2] = x_block[col][8] + x_block[col][9] + x_block[col][10] + x_block[col][11];
|
|
443
|
+
sumy[3] = x_block[col][12] + x_block[col][13] + x_block[col][14] + x_block[col][15];
|
|
444
|
+
|
|
445
|
+
acc1[0] = x_block[col][0] * f32(qs0 & 0x0003u) + x_block[col][2] * f32(qs1 & 0x0003u);
|
|
446
|
+
acc2[0] = x_block[col][1] * f32(qs0 & 0x0300u) + x_block[col][3] * f32(qs1 & 0x0300u);
|
|
447
|
+
acc1[1] = x_block[col][4] * f32(qs0 & 0x000Cu) + x_block[col][6] * f32(qs1 & 0x000Cu);
|
|
448
|
+
acc2[1] = x_block[col][5] * f32(qs0 & 0x0C00u) + x_block[col][7] * f32(qs1 & 0x0C00u);
|
|
449
|
+
acc1[2] = x_block[col][8] * f32(qs0 & 0x0030u) + x_block[col][10] * f32(qs1 & 0x0030u);
|
|
450
|
+
acc2[2] = x_block[col][9] * f32(qs0 & 0x3000u) + x_block[col][11] * f32(qs1 & 0x3000u);
|
|
451
|
+
acc1[3] = x_block[col][12] * f32(qs0 & 0x00C0u) + x_block[col][14] * f32(qs1 & 0x00C0u);
|
|
452
|
+
acc2[3] = x_block[col][13] * f32(qs0 & 0xC000u) + x_block[col][15] * f32(qs1 & 0xC000u);
|
|
453
|
+
|
|
454
|
+
acc[col][row] += dall * ((acc1[0] + (1.0/256.0) * acc2[0]) * f32(sc0 & 0xFu) +
|
|
455
|
+
(acc1[1] + (1.0/256.0) * acc2[1]) * f32(sc2 & 0xFu) / 4.0 +
|
|
456
|
+
(acc1[2] + (1.0/256.0) * acc2[2]) * f32(sc4 & 0xFu) / 16.0 +
|
|
457
|
+
(acc1[3] + (1.0/256.0) * acc2[3]) * f32(sc6 & 0xFu) / 64.0)
|
|
458
|
+
- dmin * (sumy[0] * f32(sc0 & 0xF0u) + sumy[1] * f32(sc2 & 0xF0u) +
|
|
459
|
+
sumy[2] * f32(sc4 & 0xF0u) + sumy[3] * f32(sc6 & 0xF0u));
|
|
460
|
+
}
|
|
431
461
|
}
|
|
432
462
|
}
|
|
433
463
|
}
|
|
@@ -440,8 +470,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
440
470
|
#define BLOCK_SIZE 256
|
|
441
471
|
#define BLOCK_SIZE_BYTES 110
|
|
442
472
|
#define THREADS_PER_BLOCK 16
|
|
443
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
444
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
473
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
474
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
445
475
|
|
|
446
476
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
447
477
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -485,12 +515,13 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
485
515
|
|
|
486
516
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
487
517
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
488
|
-
var x_block: array<f32, 16>;
|
|
489
|
-
for (var
|
|
490
|
-
|
|
491
|
-
|
|
518
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
519
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
520
|
+
for (var i = 0u; i < 8u; i++) {
|
|
521
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
522
|
+
x_block[col][i + 8u] = f32(src1[x_base + col * params.stride_11 + 32u + i]);
|
|
523
|
+
}
|
|
492
524
|
}
|
|
493
|
-
|
|
494
525
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
495
526
|
let output_row = row_base + row;
|
|
496
527
|
if (output_row < params.m) {
|
|
@@ -516,28 +547,30 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
516
547
|
let h_u32_0 = load_u32_at_src0(block_byte_base + h_byte + 0u);
|
|
517
548
|
let h_u32_1 = load_u32_at_src0(block_byte_base + h_byte + 4u);
|
|
518
549
|
|
|
519
|
-
var
|
|
520
|
-
|
|
521
|
-
|
|
522
|
-
|
|
523
|
-
|
|
524
|
-
|
|
525
|
-
|
|
526
|
-
|
|
527
|
-
|
|
528
|
-
|
|
529
|
-
|
|
530
|
-
|
|
531
|
-
|
|
532
|
-
|
|
533
|
-
|
|
534
|
-
|
|
535
|
-
|
|
536
|
-
|
|
550
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
551
|
+
var s1 = 0.0; var s2 = 0.0; var s3 = 0.0;
|
|
552
|
+
var s4 = 0.0; var s5 = 0.0; var s6 = 0.0;
|
|
553
|
+
|
|
554
|
+
for (var l = 0u; l < 8u; l += 2u) {
|
|
555
|
+
let q_u32 = select(q_u32_0, q_u32_1, l >= 4u);
|
|
556
|
+
let qs = select(q_u32 & 0xFFFFu, q_u32 >> 16u, (l & 2u) != 0u);
|
|
557
|
+
let h_u32 = select(h_u32_0, h_u32_1, l >= 4u);
|
|
558
|
+
let hv = select(h_u32 & 0xFFFFu, h_u32 >> 16u, (l & 2u) != 0u);
|
|
559
|
+
|
|
560
|
+
s1 += x_block[col][l + 0u] * f32(qs & qm0);
|
|
561
|
+
s2 += x_block[col][l + 1u] * f32(qs & qm1);
|
|
562
|
+
s3 += select(0.0, x_block[col][l + 0u], (hv & hm0) == 0u) +
|
|
563
|
+
select(0.0, x_block[col][l + 1u], (hv & hm1) == 0u);
|
|
564
|
+
s4 += x_block[col][l + 8u] * f32(qs & qm2);
|
|
565
|
+
s5 += x_block[col][l + 9u] * f32(qs & qm3);
|
|
566
|
+
s6 += select(0.0, x_block[col][l + 8u], (hv & hm2) == 0u) +
|
|
567
|
+
select(0.0, x_block[col][l + 9u], (hv & hm3) == 0u);
|
|
568
|
+
}
|
|
537
569
|
|
|
538
|
-
|
|
539
|
-
|
|
540
|
-
|
|
570
|
+
let d1 = d * (s1 + (1.0/256.0) * s2 - s3 * v1);
|
|
571
|
+
let d2 = d * (s4 + (1.0/256.0) * s5 - s6 * v2);
|
|
572
|
+
acc[col][row] += (d1 * scale0 + 0.25 * d2 * scale1) / f32(1u << shift);
|
|
573
|
+
}
|
|
541
574
|
}
|
|
542
575
|
}
|
|
543
576
|
}
|
|
@@ -550,8 +583,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
550
583
|
#define BLOCK_SIZE 256
|
|
551
584
|
#define BLOCK_SIZE_BYTES 144
|
|
552
585
|
#define THREADS_PER_BLOCK 16
|
|
553
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
554
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
586
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
587
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
555
588
|
|
|
556
589
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
557
590
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -573,12 +606,15 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
573
606
|
|
|
574
607
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
575
608
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
576
|
-
var x_block: array<f32, 16>;
|
|
577
|
-
for (var
|
|
578
|
-
|
|
579
|
-
|
|
580
|
-
|
|
581
|
-
|
|
609
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
610
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
611
|
+
let col_base = x_base + col * params.stride_11;
|
|
612
|
+
for (var i = 0u; i < 4u; i++) {
|
|
613
|
+
x_block[col][i] = f32(src1[col_base + i]);
|
|
614
|
+
x_block[col][i + 4u] = f32(src1[col_base + 32u + i]);
|
|
615
|
+
x_block[col][i + 8u] = f32(src1[col_base + 128u + i]);
|
|
616
|
+
x_block[col][i + 12u] = f32(src1[col_base + 160u + i]);
|
|
617
|
+
}
|
|
582
618
|
}
|
|
583
619
|
|
|
584
620
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
@@ -613,23 +649,25 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
613
649
|
let q1_u32 = load_u32_at_src0_aligned(block_byte_base + 16u + q_offset);
|
|
614
650
|
let q2_u32 = load_u32_at_src0_aligned(block_byte_base + 80u + q_offset);
|
|
615
651
|
|
|
616
|
-
var
|
|
617
|
-
|
|
618
|
-
|
|
619
|
-
|
|
620
|
-
|
|
621
|
-
|
|
622
|
-
|
|
623
|
-
|
|
624
|
-
|
|
625
|
-
|
|
626
|
-
|
|
627
|
-
|
|
628
|
-
|
|
629
|
-
|
|
652
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
653
|
+
var dot = vec4<f32>(0.0, 0.0, 0.0, 0.0);
|
|
654
|
+
var sumx = vec4<f32>(0.0, 0.0, 0.0, 0.0);
|
|
655
|
+
for (var i = 0u; i < 4u; i++) {
|
|
656
|
+
let q1b = byte_of(q1_u32, i);
|
|
657
|
+
let q2b = byte_of(q2_u32, i);
|
|
658
|
+
dot[0] += x_block[col][i] * f32(q1b & 0x0Fu);
|
|
659
|
+
dot[1] += x_block[col][i + 4u] * f32(q1b >> 4u);
|
|
660
|
+
dot[2] += x_block[col][i + 8u] * f32(q2b & 0x0Fu);
|
|
661
|
+
dot[3] += x_block[col][i + 12u] * f32(q2b >> 4u);
|
|
662
|
+
sumx[0] += x_block[col][i];
|
|
663
|
+
sumx[1] += x_block[col][i + 4u];
|
|
664
|
+
sumx[2] += x_block[col][i + 8u];
|
|
665
|
+
sumx[3] += x_block[col][i + 12u];
|
|
666
|
+
}
|
|
630
667
|
|
|
631
|
-
|
|
632
|
-
|
|
668
|
+
acc[col][row] += d * (dot[0] * scale0 + dot[1] * scale1 + dot[2] * scale2 + dot[3] * scale3)
|
|
669
|
+
- dmin * (sumx[0] * min0 + sumx[1] * min1 + sumx[2] * min2 + sumx[3] * min3);
|
|
670
|
+
}
|
|
633
671
|
}
|
|
634
672
|
}
|
|
635
673
|
}
|
|
@@ -642,8 +680,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
642
680
|
#define BLOCK_SIZE 256
|
|
643
681
|
#define BLOCK_SIZE_BYTES 176
|
|
644
682
|
#define THREADS_PER_BLOCK 16
|
|
645
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
646
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
683
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
684
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
647
685
|
|
|
648
686
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
649
687
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -671,14 +709,16 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
671
709
|
|
|
672
710
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
673
711
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
674
|
-
var x_block: array<f32, 16>;
|
|
675
|
-
for (var
|
|
676
|
-
|
|
677
|
-
|
|
678
|
-
|
|
679
|
-
|
|
712
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
713
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
714
|
+
let col_base = x_base + col * params.stride_11;
|
|
715
|
+
for (var i = 0u; i < 4u; i++) {
|
|
716
|
+
x_block[col][i] = f32(src1[col_base + i]);
|
|
717
|
+
x_block[col][i + 4u] = f32(src1[col_base + 32u + i]);
|
|
718
|
+
x_block[col][i + 8u] = f32(src1[col_base + 128u + i]);
|
|
719
|
+
x_block[col][i + 12u] = f32(src1[col_base + 160u + i]);
|
|
720
|
+
}
|
|
680
721
|
}
|
|
681
|
-
|
|
682
722
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
683
723
|
let output_row = row_base + row;
|
|
684
724
|
if (output_row < params.m) {
|
|
@@ -712,37 +752,39 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
712
752
|
let q2_u32 = load_u32_at_src0_aligned(block_byte_base + q_offset + 64u);
|
|
713
753
|
let qh_u32 = load_u32_at_src0_aligned(block_byte_base + qh_offset);
|
|
714
754
|
|
|
715
|
-
var
|
|
716
|
-
|
|
717
|
-
|
|
718
|
-
|
|
719
|
-
|
|
720
|
-
|
|
721
|
-
|
|
722
|
-
|
|
723
|
-
|
|
724
|
-
|
|
725
|
-
|
|
726
|
-
|
|
727
|
-
|
|
728
|
-
|
|
729
|
-
|
|
730
|
-
|
|
731
|
-
|
|
732
|
-
|
|
733
|
-
|
|
734
|
-
|
|
735
|
-
|
|
736
|
-
|
|
737
|
-
|
|
738
|
-
|
|
739
|
-
|
|
740
|
-
|
|
741
|
-
|
|
755
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
756
|
+
var vals = vec4<f32>(0.0, 0.0, 0.0, 0.0);
|
|
757
|
+
var sumy = vec4<f32>(0.0, 0.0, 0.0, 0.0);
|
|
758
|
+
for (var i = 0u; i < 4u; i++) {
|
|
759
|
+
let q1b = byte_of(q1_u32, i);
|
|
760
|
+
let q2b = byte_of(q2_u32, i);
|
|
761
|
+
let qhb = byte_of(qh_u32, i);
|
|
762
|
+
|
|
763
|
+
let yl0 = x_block[col][i];
|
|
764
|
+
let yl8 = x_block[col][i + 4u];
|
|
765
|
+
let yh0 = x_block[col][i + 8u];
|
|
766
|
+
let yh8 = x_block[col][i + 12u];
|
|
767
|
+
|
|
768
|
+
sumy[0] += yl0;
|
|
769
|
+
sumy[1] += yl8;
|
|
770
|
+
sumy[2] += yh0;
|
|
771
|
+
sumy[3] += yh8;
|
|
772
|
+
|
|
773
|
+
let q0 = f32((q1b & 0x0Fu) | select(0u, 0x10u, (qhb & hm1) != 0u));
|
|
774
|
+
let q1 = f32((q1b >> 4u) | select(0u, 0x10u, (qhb & hm2) != 0u));
|
|
775
|
+
let q2 = f32((q2b & 0x0Fu) | select(0u, 0x10u, (qhb & hm3) != 0u));
|
|
776
|
+
let q3 = f32((q2b >> 4u) | select(0u, 0x10u, (qhb & hm4) != 0u));
|
|
777
|
+
|
|
778
|
+
vals[0] += yl0 * q0;
|
|
779
|
+
vals[1] += yl8 * q1;
|
|
780
|
+
vals[2] += yh0 * q2;
|
|
781
|
+
vals[3] += yh8 * q3;
|
|
782
|
+
}
|
|
742
783
|
|
|
743
|
-
|
|
744
|
-
|
|
745
|
-
|
|
784
|
+
acc[col][row] += d * (f0 * vals[0] + f1 * vals[1] + f4 * vals[2] + f5 * vals[3])
|
|
785
|
+
- dmin * (sumy[0] * m0 + sumy[1] * m1 +
|
|
786
|
+
sumy[2] * m4 + sumy[3] * m5);
|
|
787
|
+
}
|
|
746
788
|
}
|
|
747
789
|
}
|
|
748
790
|
}
|
|
@@ -755,8 +797,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
755
797
|
#define BLOCK_SIZE 256
|
|
756
798
|
#define BLOCK_SIZE_BYTES 210
|
|
757
799
|
#define THREADS_PER_BLOCK 16
|
|
758
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
759
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
800
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
801
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
760
802
|
|
|
761
803
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
762
804
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -777,14 +819,16 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
777
819
|
|
|
778
820
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
779
821
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
780
|
-
var x_block: array<f32, 16>;
|
|
781
|
-
for (var
|
|
782
|
-
|
|
783
|
-
|
|
784
|
-
|
|
785
|
-
|
|
822
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
823
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
824
|
+
let col_base = x_base + col * params.stride_11;
|
|
825
|
+
for (var l = 0u; l < 4u; l++) {
|
|
826
|
+
x_block[col][l] = f32(src1[col_base + l]);
|
|
827
|
+
x_block[col][l + 4u] = f32(src1[col_base + 32u + l]);
|
|
828
|
+
x_block[col][l + 8u] = f32(src1[col_base + 64u + l]);
|
|
829
|
+
x_block[col][l + 12u] = f32(src1[col_base + 96u + l]);
|
|
830
|
+
}
|
|
786
831
|
}
|
|
787
|
-
|
|
788
832
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
789
833
|
let output_row = row_base + row;
|
|
790
834
|
if (output_row < params.m) {
|
|
@@ -802,26 +846,28 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
802
846
|
let sc4 = sbyte_of(sc_u32_1, sc_byte_pos);
|
|
803
847
|
let sc6 = sbyte_of(sc_u32_1, sc_byte_pos + 2u);
|
|
804
848
|
|
|
805
|
-
var
|
|
849
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
850
|
+
var sums = vec4<f32>(0.0, 0.0, 0.0, 0.0);
|
|
806
851
|
|
|
807
|
-
|
|
808
|
-
|
|
809
|
-
|
|
810
|
-
|
|
852
|
+
for (var l = 0u; l < 4u; l++) {
|
|
853
|
+
let q1b = byte_of(ql1_u32, l);
|
|
854
|
+
let q2b = byte_of(ql2_u32, l);
|
|
855
|
+
let qhb = byte_of(qh_u32, l);
|
|
811
856
|
|
|
812
|
-
|
|
813
|
-
|
|
814
|
-
|
|
815
|
-
|
|
857
|
+
let dq0 = f32(i32((q1b & 0x0Fu) | ((qhb & 0x03u) << 4u)) - 32);
|
|
858
|
+
let dq1 = f32(i32((q2b & 0x0Fu) | ((qhb & 0x0Cu) << 2u)) - 32);
|
|
859
|
+
let dq2 = f32(i32((q1b >> 4u) | (qhb & 0x30u)) - 32);
|
|
860
|
+
let dq3 = f32(i32((q2b >> 4u) | ((qhb & 0xC0u) >> 2u)) - 32);
|
|
816
861
|
|
|
817
|
-
|
|
818
|
-
|
|
819
|
-
|
|
820
|
-
|
|
821
|
-
|
|
862
|
+
sums[0] += x_block[col][l] * dq0;
|
|
863
|
+
sums[1] += x_block[col][l + 4u] * dq1;
|
|
864
|
+
sums[2] += x_block[col][l + 8u] * dq2;
|
|
865
|
+
sums[3] += x_block[col][l + 12u] * dq3;
|
|
866
|
+
}
|
|
822
867
|
|
|
823
|
-
|
|
824
|
-
|
|
868
|
+
acc[col][row] += d * (sums[0] * f32(sc0) + sums[1] * f32(sc2) +
|
|
869
|
+
sums[2] * f32(sc4) + sums[3] * f32(sc6));
|
|
870
|
+
}
|
|
825
871
|
}
|
|
826
872
|
}
|
|
827
873
|
}
|
|
@@ -834,8 +880,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
834
880
|
#define BLOCK_SIZE 256
|
|
835
881
|
#define BLOCK_SIZE_BYTES 50
|
|
836
882
|
#define THREADS_PER_BLOCK 16
|
|
837
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
838
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
883
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
884
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
839
885
|
|
|
840
886
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
841
887
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -850,11 +896,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
850
896
|
|
|
851
897
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
852
898
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
853
|
-
var x_block: array<f32, 16>;
|
|
854
|
-
for (var
|
|
855
|
-
|
|
899
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
900
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
901
|
+
for (var i = 0u; i < 16u; i++) {
|
|
902
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
903
|
+
}
|
|
856
904
|
}
|
|
857
|
-
|
|
858
905
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
859
906
|
let output_row = row_base + row;
|
|
860
907
|
if (output_row < params.m) {
|
|
@@ -866,20 +913,22 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
866
913
|
let delta = select(IQ1_DELTA, -IQ1_DELTA, (qh & 0x8000u) != 0u);
|
|
867
914
|
let qs_w = load_u32_at_src0(block_byte_base + 2u + sub_blk * 4u);
|
|
868
915
|
|
|
869
|
-
var
|
|
870
|
-
|
|
871
|
-
|
|
872
|
-
|
|
873
|
-
|
|
874
|
-
|
|
875
|
-
|
|
876
|
-
|
|
877
|
-
|
|
878
|
-
|
|
879
|
-
|
|
916
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
917
|
+
var row_sum = 0.0;
|
|
918
|
+
for (var ll = 0u; ll < 2u; ll++) {
|
|
919
|
+
let l = slot0 + ll;
|
|
920
|
+
let qs_byte = get_byte(qs_w, l);
|
|
921
|
+
let ig = (qs_byte | (((qh >> (3u * l)) & 7u) << 8u)) * 8u;
|
|
922
|
+
let gw = iq1_grid[ig / 16u];
|
|
923
|
+
let bit_base = (ig % 16u) * 2u;
|
|
924
|
+
for (var j = 0u; j < 8u; j++) {
|
|
925
|
+
let g = (gw >> (bit_base + j * 2u)) & 3u;
|
|
926
|
+
let gs = select(f32(g), f32(g) - 4.0, (g & 2u) != 0u);
|
|
927
|
+
row_sum += dl * (gs + delta) * x_block[col][ll * 8u + j];
|
|
928
|
+
}
|
|
880
929
|
}
|
|
930
|
+
acc[col][row] += row_sum;
|
|
881
931
|
}
|
|
882
|
-
acc[row] += row_sum;
|
|
883
932
|
}
|
|
884
933
|
}
|
|
885
934
|
}
|
|
@@ -892,8 +941,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
892
941
|
#define BLOCK_SIZE 256
|
|
893
942
|
#define BLOCK_SIZE_BYTES 56
|
|
894
943
|
#define THREADS_PER_BLOCK 16
|
|
895
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
896
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
944
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
945
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
897
946
|
|
|
898
947
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
899
948
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -908,11 +957,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
908
957
|
|
|
909
958
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
910
959
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
911
|
-
var x_block: array<f32, 16>;
|
|
912
|
-
for (var
|
|
913
|
-
|
|
960
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
961
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
962
|
+
for (var i = 0u; i < 16u; i++) {
|
|
963
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
964
|
+
}
|
|
914
965
|
}
|
|
915
|
-
|
|
916
966
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
917
967
|
let output_row = row_base + row;
|
|
918
968
|
if (output_row < params.m) {
|
|
@@ -936,26 +986,28 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
936
986
|
let qh_lo = qh & 0xFFu;
|
|
937
987
|
let qh_hi = (qh >> 8u) & 0xFFu;
|
|
938
988
|
|
|
939
|
-
var
|
|
940
|
-
|
|
941
|
-
|
|
942
|
-
|
|
943
|
-
|
|
944
|
-
|
|
945
|
-
|
|
946
|
-
|
|
947
|
-
|
|
948
|
-
|
|
949
|
-
|
|
950
|
-
|
|
951
|
-
|
|
952
|
-
|
|
953
|
-
|
|
954
|
-
|
|
955
|
-
|
|
989
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
990
|
+
var row_sum = 0.0;
|
|
991
|
+
for (var ll = 0u; ll < 2u; ll++) {
|
|
992
|
+
let l = slot0 + ll;
|
|
993
|
+
let bit_off = 6u * (sub_blk % 2u) + 3u * (l / 2u);
|
|
994
|
+
let sub_scale = (sc_u16 >> bit_off) & 0x7u;
|
|
995
|
+
let dl = d * f32(2u * sub_scale + 1u);
|
|
996
|
+
let qh_byte = select(qh_lo, qh_hi, l >= 2u);
|
|
997
|
+
let ll2 = l % 2u;
|
|
998
|
+
let grid_idx = get_byte(qs_w, l) | (((qh_byte >> (4u * ll2)) & 7u) << 8u);
|
|
999
|
+
let delta = select(IQ1_DELTA, -IQ1_DELTA, ((qh_byte >> (3u + 4u * ll2)) & 1u) != 0u);
|
|
1000
|
+
let ig = grid_idx * 8u;
|
|
1001
|
+
let gw = iq1_grid[ig / 16u];
|
|
1002
|
+
let bit_base = (ig % 16u) * 2u;
|
|
1003
|
+
for (var j = 0u; j < 8u; j++) {
|
|
1004
|
+
let g = (gw >> (bit_base + j * 2u)) & 3u;
|
|
1005
|
+
let gs = select(f32(g), f32(g) - 4.0, (g & 2u) != 0u);
|
|
1006
|
+
row_sum += dl * (gs + delta) * x_block[col][ll * 8u + j];
|
|
1007
|
+
}
|
|
956
1008
|
}
|
|
1009
|
+
acc[col][row] += row_sum;
|
|
957
1010
|
}
|
|
958
|
-
acc[row] += row_sum;
|
|
959
1011
|
}
|
|
960
1012
|
}
|
|
961
1013
|
}
|
|
@@ -968,8 +1020,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
968
1020
|
#define BLOCK_SIZE 256
|
|
969
1021
|
#define BLOCK_SIZE_BYTES 66
|
|
970
1022
|
#define THREADS_PER_BLOCK 16
|
|
971
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
972
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
1023
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
1024
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
973
1025
|
|
|
974
1026
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
975
1027
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -984,11 +1036,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
984
1036
|
|
|
985
1037
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
986
1038
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
987
|
-
var x_block: array<f32, 16>;
|
|
988
|
-
for (var
|
|
989
|
-
|
|
1039
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
1040
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
1041
|
+
for (var i = 0u; i < 16u; i++) {
|
|
1042
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
1043
|
+
}
|
|
990
1044
|
}
|
|
991
|
-
|
|
992
1045
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
993
1046
|
let output_row = row_base + row;
|
|
994
1047
|
if (output_row < params.m) {
|
|
@@ -999,22 +1052,24 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
999
1052
|
let ls = aux_hi >> 28u;
|
|
1000
1053
|
let db = d * (0.5 + f32(ls)) * 0.25;
|
|
1001
1054
|
|
|
1002
|
-
var
|
|
1003
|
-
|
|
1004
|
-
|
|
1005
|
-
|
|
1006
|
-
|
|
1007
|
-
|
|
1008
|
-
|
|
1009
|
-
|
|
1010
|
-
|
|
1011
|
-
|
|
1012
|
-
|
|
1013
|
-
|
|
1014
|
-
|
|
1055
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
1056
|
+
var row_sum = 0.0;
|
|
1057
|
+
for (var ll = 0u; ll < 2u; ll++) {
|
|
1058
|
+
let l = slot0 + ll;
|
|
1059
|
+
let grid_idx = (aux_lo >> (8u * l)) & 0xFFu;
|
|
1060
|
+
let signs_idx = (aux_hi >> (7u * l)) & 0x7Fu;
|
|
1061
|
+
let signs = (ksigns_iq2xs[signs_idx / 4u] >> ((signs_idx % 4u) * 8u)) & 0xFFu;
|
|
1062
|
+
let gw_lo = iq2xxs_grid[grid_idx * 2u];
|
|
1063
|
+
let gw_hi = iq2xxs_grid[grid_idx * 2u + 1u];
|
|
1064
|
+
for (var j = 0u; j < 8u; j++) {
|
|
1065
|
+
let gw = select(gw_hi, gw_lo, j < 4u);
|
|
1066
|
+
let b = f32((gw >> ((j & 3u) * 8u)) & 0xFFu);
|
|
1067
|
+
let s = select(1.0, -1.0, ((signs >> j) & 1u) != 0u);
|
|
1068
|
+
row_sum += db * b * s * x_block[col][ll * 8u + j];
|
|
1069
|
+
}
|
|
1015
1070
|
}
|
|
1071
|
+
acc[col][row] += row_sum;
|
|
1016
1072
|
}
|
|
1017
|
-
acc[row] += row_sum;
|
|
1018
1073
|
}
|
|
1019
1074
|
}
|
|
1020
1075
|
}
|
|
@@ -1027,8 +1082,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1027
1082
|
#define BLOCK_SIZE 256
|
|
1028
1083
|
#define BLOCK_SIZE_BYTES 74
|
|
1029
1084
|
#define THREADS_PER_BLOCK 16
|
|
1030
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
1031
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
1085
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
1086
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
1032
1087
|
|
|
1033
1088
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
1034
1089
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -1043,11 +1098,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1043
1098
|
|
|
1044
1099
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
1045
1100
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
1046
|
-
var x_block: array<f32, 16>;
|
|
1047
|
-
for (var
|
|
1048
|
-
|
|
1101
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
1102
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
1103
|
+
for (var i = 0u; i < 16u; i++) {
|
|
1104
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
1105
|
+
}
|
|
1049
1106
|
}
|
|
1050
|
-
|
|
1051
1107
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
1052
1108
|
let output_row = row_base + row;
|
|
1053
1109
|
if (output_row < params.m) {
|
|
@@ -1058,27 +1114,29 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1058
1114
|
let scales_word = load_u32_at_src0(block_byte_base + 66u + (sub_blk / 4u) * 4u);
|
|
1059
1115
|
let scales_byte = get_byte(scales_word, sub_blk % 4u);
|
|
1060
1116
|
|
|
1061
|
-
var
|
|
1062
|
-
|
|
1063
|
-
|
|
1064
|
-
|
|
1065
|
-
|
|
1066
|
-
|
|
1067
|
-
|
|
1068
|
-
|
|
1069
|
-
|
|
1070
|
-
|
|
1071
|
-
|
|
1072
|
-
|
|
1073
|
-
|
|
1074
|
-
|
|
1075
|
-
|
|
1076
|
-
|
|
1077
|
-
|
|
1078
|
-
|
|
1117
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
1118
|
+
var row_sum = 0.0;
|
|
1119
|
+
for (var ll = 0u; ll < 2u; ll++) {
|
|
1120
|
+
let l = slot0 + ll;
|
|
1121
|
+
let qs_word = select(qs_hi, qs_lo, l < 2u);
|
|
1122
|
+
let half2 = (l % 2u) * 16u;
|
|
1123
|
+
let qs_val = (qs_word >> half2) & 0xFFFFu;
|
|
1124
|
+
let grid_idx = qs_val & 0x1FFu;
|
|
1125
|
+
let signs_idx = (qs_val >> 9u) & 0x7Fu;
|
|
1126
|
+
let sub_scale = (scales_byte >> (4u * (l / 2u))) & 0xFu;
|
|
1127
|
+
let db = d * (0.5 + f32(sub_scale)) * 0.25;
|
|
1128
|
+
let signs = (ksigns_iq2xs[signs_idx / 4u] >> ((signs_idx % 4u) * 8u)) & 0xFFu;
|
|
1129
|
+
let gw_lo = iq2xs_grid[grid_idx * 2u];
|
|
1130
|
+
let gw_hi = iq2xs_grid[grid_idx * 2u + 1u];
|
|
1131
|
+
for (var j = 0u; j < 8u; j++) {
|
|
1132
|
+
let gw = select(gw_hi, gw_lo, j < 4u);
|
|
1133
|
+
let b = f32((gw >> ((j & 3u) * 8u)) & 0xFFu);
|
|
1134
|
+
let s = select(1.0, -1.0, ((signs >> j) & 1u) != 0u);
|
|
1135
|
+
row_sum += db * b * s * x_block[col][ll * 8u + j];
|
|
1136
|
+
}
|
|
1079
1137
|
}
|
|
1138
|
+
acc[col][row] += row_sum;
|
|
1080
1139
|
}
|
|
1081
|
-
acc[row] += row_sum;
|
|
1082
1140
|
}
|
|
1083
1141
|
}
|
|
1084
1142
|
}
|
|
@@ -1091,8 +1149,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1091
1149
|
#define BLOCK_SIZE 256
|
|
1092
1150
|
#define BLOCK_SIZE_BYTES 82
|
|
1093
1151
|
#define THREADS_PER_BLOCK 16
|
|
1094
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
1095
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
1152
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
1153
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
1096
1154
|
|
|
1097
1155
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
1098
1156
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -1107,11 +1165,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1107
1165
|
|
|
1108
1166
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
1109
1167
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
1110
|
-
var x_block: array<f32, 16>;
|
|
1111
|
-
for (var
|
|
1112
|
-
|
|
1168
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
1169
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
1170
|
+
for (var i = 0u; i < 16u; i++) {
|
|
1171
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
1172
|
+
}
|
|
1113
1173
|
}
|
|
1114
|
-
|
|
1115
1174
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
1116
1175
|
let output_row = row_base + row;
|
|
1117
1176
|
if (output_row < params.m) {
|
|
@@ -1124,24 +1183,26 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1124
1183
|
let sc_word = load_u32_at_src0(block_byte_base + 74u + (sub_blk / 4u) * 4u);
|
|
1125
1184
|
let scales_byte = get_byte(sc_word, sub_blk % 4u);
|
|
1126
1185
|
|
|
1127
|
-
var
|
|
1128
|
-
|
|
1129
|
-
|
|
1130
|
-
|
|
1131
|
-
|
|
1132
|
-
|
|
1133
|
-
|
|
1134
|
-
|
|
1135
|
-
|
|
1136
|
-
|
|
1137
|
-
|
|
1138
|
-
|
|
1139
|
-
|
|
1140
|
-
|
|
1141
|
-
|
|
1186
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
1187
|
+
var row_sum = 0.0;
|
|
1188
|
+
for (var ll = 0u; ll < 2u; ll++) {
|
|
1189
|
+
let l = slot0 + ll;
|
|
1190
|
+
let qs_byte = get_byte(qs_w, l);
|
|
1191
|
+
let sign_byte = get_byte(sg_w, l);
|
|
1192
|
+
let grid_idx = qs_byte | (((qh_byte >> (2u * l)) & 3u) << 8u);
|
|
1193
|
+
let sub_scale = (scales_byte >> (4u * (l / 2u))) & 0xFu;
|
|
1194
|
+
let db = d * (0.5 + f32(sub_scale)) * 0.25;
|
|
1195
|
+
let gw_lo = iq2s_grid[grid_idx * 2u];
|
|
1196
|
+
let gw_hi = iq2s_grid[grid_idx * 2u + 1u];
|
|
1197
|
+
for (var j = 0u; j < 8u; j++) {
|
|
1198
|
+
let gw = select(gw_hi, gw_lo, j < 4u);
|
|
1199
|
+
let b = f32((gw >> ((j & 3u) * 8u)) & 0xFFu);
|
|
1200
|
+
let s = select(1.0, -1.0, ((sign_byte >> j) & 1u) != 0u);
|
|
1201
|
+
row_sum += db * b * s * x_block[col][ll * 8u + j];
|
|
1202
|
+
}
|
|
1142
1203
|
}
|
|
1204
|
+
acc[col][row] += row_sum;
|
|
1143
1205
|
}
|
|
1144
|
-
acc[row] += row_sum;
|
|
1145
1206
|
}
|
|
1146
1207
|
}
|
|
1147
1208
|
}
|
|
@@ -1154,8 +1215,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1154
1215
|
#define BLOCK_SIZE 256
|
|
1155
1216
|
#define BLOCK_SIZE_BYTES 98
|
|
1156
1217
|
#define THREADS_PER_BLOCK 16
|
|
1157
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
1158
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
1218
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
1219
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
1159
1220
|
|
|
1160
1221
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
1161
1222
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -1170,11 +1231,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1170
1231
|
|
|
1171
1232
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
1172
1233
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
1173
|
-
var x_block: array<f32, 16>;
|
|
1174
|
-
for (var
|
|
1175
|
-
|
|
1234
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
1235
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
1236
|
+
for (var i = 0u; i < 16u; i++) {
|
|
1237
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
1238
|
+
}
|
|
1176
1239
|
}
|
|
1177
|
-
|
|
1178
1240
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
1179
1241
|
let output_row = row_base + row;
|
|
1180
1242
|
if (output_row < params.m) {
|
|
@@ -1186,27 +1248,29 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1186
1248
|
let ls = aux >> 28u;
|
|
1187
1249
|
let db = d * (0.5 + f32(ls)) * 0.5;
|
|
1188
1250
|
|
|
1189
|
-
var
|
|
1190
|
-
|
|
1191
|
-
|
|
1192
|
-
|
|
1193
|
-
|
|
1194
|
-
|
|
1195
|
-
|
|
1196
|
-
|
|
1197
|
-
|
|
1198
|
-
|
|
1199
|
-
|
|
1200
|
-
|
|
1201
|
-
|
|
1202
|
-
|
|
1203
|
-
|
|
1204
|
-
|
|
1205
|
-
|
|
1206
|
-
|
|
1251
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
1252
|
+
var row_sum = 0.0;
|
|
1253
|
+
for (var ll = 0u; ll < 2u; ll++) {
|
|
1254
|
+
let l = slot0 + ll;
|
|
1255
|
+
let qs_word = select(qs_hi, qs_lo, l < 2u);
|
|
1256
|
+
let byte_pos = (l % 2u) * 2u;
|
|
1257
|
+
let grid_idx_0 = (qs_word >> (byte_pos * 8u)) & 0xFFu;
|
|
1258
|
+
let grid_idx_1 = (qs_word >> ((byte_pos + 1u) * 8u)) & 0xFFu;
|
|
1259
|
+
let signs_idx = (aux >> (7u * l)) & 0x7Fu;
|
|
1260
|
+
let signs = (ksigns_iq2xs[signs_idx / 4u] >> ((signs_idx % 4u) * 8u)) & 0xFFu;
|
|
1261
|
+
let grid1 = iq3xxs_grid[grid_idx_0];
|
|
1262
|
+
let grid2 = iq3xxs_grid[grid_idx_1];
|
|
1263
|
+
for (var j = 0u; j < 4u; j++) {
|
|
1264
|
+
let b1 = f32((grid1 >> (j * 8u)) & 0xFFu);
|
|
1265
|
+
let b2 = f32((grid2 >> (j * 8u)) & 0xFFu);
|
|
1266
|
+
let s1 = select(1.0, -1.0, ((signs >> j) & 1u) != 0u);
|
|
1267
|
+
let s2 = select(1.0, -1.0, ((signs >> (j + 4u)) & 1u) != 0u);
|
|
1268
|
+
row_sum += db * b1 * s1 * x_block[col][ll * 8u + j];
|
|
1269
|
+
row_sum += db * b2 * s2 * x_block[col][ll * 8u + j + 4u];
|
|
1270
|
+
}
|
|
1207
1271
|
}
|
|
1272
|
+
acc[col][row] += row_sum;
|
|
1208
1273
|
}
|
|
1209
|
-
acc[row] += row_sum;
|
|
1210
1274
|
}
|
|
1211
1275
|
}
|
|
1212
1276
|
}
|
|
@@ -1219,8 +1283,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1219
1283
|
#define BLOCK_SIZE 256
|
|
1220
1284
|
#define BLOCK_SIZE_BYTES 110
|
|
1221
1285
|
#define THREADS_PER_BLOCK 16
|
|
1222
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
1223
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
1286
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
1287
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
1224
1288
|
|
|
1225
1289
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
1226
1290
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -1235,11 +1299,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1235
1299
|
|
|
1236
1300
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
1237
1301
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
1238
|
-
var x_block: array<f32, 16>;
|
|
1239
|
-
for (var
|
|
1240
|
-
|
|
1302
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
1303
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
1304
|
+
for (var i = 0u; i < 16u; i++) {
|
|
1305
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
1306
|
+
}
|
|
1241
1307
|
}
|
|
1242
|
-
|
|
1243
1308
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
1244
1309
|
let output_row = row_base + row;
|
|
1245
1310
|
if (output_row < params.m) {
|
|
@@ -1255,28 +1320,30 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1255
1320
|
let sub_scale = (scales_byte >> (4u * (sub_blk % 2u))) & 0xFu;
|
|
1256
1321
|
let db = d * (1.0 + 2.0 * f32(sub_scale));
|
|
1257
1322
|
|
|
1258
|
-
var
|
|
1259
|
-
|
|
1260
|
-
|
|
1261
|
-
|
|
1262
|
-
|
|
1263
|
-
|
|
1264
|
-
|
|
1265
|
-
|
|
1266
|
-
|
|
1267
|
-
|
|
1268
|
-
|
|
1269
|
-
|
|
1270
|
-
|
|
1271
|
-
|
|
1272
|
-
|
|
1273
|
-
|
|
1274
|
-
|
|
1275
|
-
|
|
1276
|
-
|
|
1323
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
1324
|
+
var row_sum = 0.0;
|
|
1325
|
+
for (var ll = 0u; ll < 2u; ll++) {
|
|
1326
|
+
let l = slot0 + ll;
|
|
1327
|
+
let qs_word = select(qs_hi, qs_lo, l < 2u);
|
|
1328
|
+
let byte_pos = (l % 2u) * 2u;
|
|
1329
|
+
let qs0 = (qs_word >> (byte_pos * 8u)) & 0xFFu;
|
|
1330
|
+
let qs1 = (qs_word >> ((byte_pos + 1u) * 8u)) & 0xFFu;
|
|
1331
|
+
let grid_idx_1 = qs0 | (((qh_byte >> (2u * l)) & 1u) << 8u);
|
|
1332
|
+
let grid_idx_2 = qs1 | (((qh_byte >> (2u * l + 1u)) & 1u) << 8u);
|
|
1333
|
+
let sign_byte = get_byte(sg_w, l);
|
|
1334
|
+
let grid1 = iq3s_grid[grid_idx_1];
|
|
1335
|
+
let grid2 = iq3s_grid[grid_idx_2];
|
|
1336
|
+
for (var j = 0u; j < 4u; j++) {
|
|
1337
|
+
let b1 = f32((grid1 >> (j * 8u)) & 0xFFu);
|
|
1338
|
+
let b2 = f32((grid2 >> (j * 8u)) & 0xFFu);
|
|
1339
|
+
let s1 = select(1.0, -1.0, ((sign_byte >> j) & 1u) != 0u);
|
|
1340
|
+
let s2 = select(1.0, -1.0, ((sign_byte >> (j + 4u)) & 1u) != 0u);
|
|
1341
|
+
row_sum += db * b1 * s1 * x_block[col][ll * 8u + j];
|
|
1342
|
+
row_sum += db * b2 * s2 * x_block[col][ll * 8u + j + 4u];
|
|
1343
|
+
}
|
|
1277
1344
|
}
|
|
1345
|
+
acc[col][row] += row_sum;
|
|
1278
1346
|
}
|
|
1279
|
-
acc[row] += row_sum;
|
|
1280
1347
|
}
|
|
1281
1348
|
}
|
|
1282
1349
|
}
|
|
@@ -1290,35 +1357,37 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1290
1357
|
#define BLOCK_SIZE_BYTES 18
|
|
1291
1358
|
#define THREADS_PER_BLOCK 4
|
|
1292
1359
|
#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
|
|
1293
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
1294
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
1360
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
1361
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
1295
1362
|
|
|
1296
1363
|
let num_blocks = params.k / BLOCK_SIZE;
|
|
1297
1364
|
let thread_within_block = thread_id % THREADS_PER_BLOCK;
|
|
1298
1365
|
for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
|
|
1299
1366
|
let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4u;
|
|
1300
|
-
var x_block: array<f32, ELEMS_PER_THREAD>;
|
|
1301
|
-
for (var
|
|
1302
|
-
|
|
1303
|
-
|
|
1367
|
+
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
|
1368
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
1369
|
+
for (var i = 0u; i < ELEMS_PER_THREAD / 2u; i++) {
|
|
1370
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
1371
|
+
x_block[col][i + 4u] = f32(src1[x_base + col * params.stride_11 + i + 16u]);
|
|
1372
|
+
}
|
|
1304
1373
|
}
|
|
1305
|
-
|
|
1306
1374
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
1307
1375
|
let output_row = row_base + row;
|
|
1308
1376
|
if (output_row < params.m) {
|
|
1309
1377
|
let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
|
|
1310
1378
|
let d = f32(load_f16_at_src0(block_byte_base));
|
|
1311
|
-
var row_sum = 0.0;
|
|
1312
|
-
|
|
1313
1379
|
let q_packed = load_u32_at_src0(block_byte_base + 2u + 4u * thread_within_block);
|
|
1314
|
-
for (var
|
|
1315
|
-
|
|
1316
|
-
|
|
1317
|
-
|
|
1318
|
-
|
|
1319
|
-
|
|
1380
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
1381
|
+
var row_sum = 0.0;
|
|
1382
|
+
for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
|
|
1383
|
+
let q_byte = get_byte(q_packed, byte_idx);
|
|
1384
|
+
let q_lo = f32(kvalues_iq4nl[q_byte & 0xFu]) * d;
|
|
1385
|
+
let q_hi = f32(kvalues_iq4nl[(q_byte >> 4u) & 0xFu]) * d;
|
|
1386
|
+
row_sum += q_lo * x_block[col][byte_idx];
|
|
1387
|
+
row_sum += q_hi * x_block[col][byte_idx + 4u];
|
|
1388
|
+
}
|
|
1389
|
+
acc[col][row] += row_sum;
|
|
1320
1390
|
}
|
|
1321
|
-
acc[row] += row_sum;
|
|
1322
1391
|
}
|
|
1323
1392
|
}
|
|
1324
1393
|
}
|
|
@@ -1331,8 +1400,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1331
1400
|
#define BLOCK_SIZE 256
|
|
1332
1401
|
#define BLOCK_SIZE_BYTES 136
|
|
1333
1402
|
#define THREADS_PER_BLOCK 16
|
|
1334
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
1335
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
1403
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
1404
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
1336
1405
|
|
|
1337
1406
|
let tid = thread_id % THREADS_PER_BLOCK;
|
|
1338
1407
|
let block_group = thread_id / THREADS_PER_BLOCK;
|
|
@@ -1346,11 +1415,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1346
1415
|
|
|
1347
1416
|
for (var block = block_group; block < num_blocks; block += num_block_groups) {
|
|
1348
1417
|
let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
|
|
1349
|
-
var x_block: array<f32, 16>;
|
|
1350
|
-
for (var
|
|
1351
|
-
|
|
1418
|
+
var x_block: array<array<f32, 16>, NUM_COLS>;
|
|
1419
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
1420
|
+
for (var i = 0u; i < 16u; i++) {
|
|
1421
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
1422
|
+
}
|
|
1352
1423
|
}
|
|
1353
|
-
|
|
1354
1424
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
1355
1425
|
let output_row = row_base + row;
|
|
1356
1426
|
if (output_row < params.m) {
|
|
@@ -1370,17 +1440,19 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1370
1440
|
let q_w2 = load_u32_at_src0(block_byte_base + qs_byte_off + 8u);
|
|
1371
1441
|
let q_w3 = load_u32_at_src0(block_byte_base + qs_byte_off + 12u);
|
|
1372
1442
|
|
|
1373
|
-
var
|
|
1374
|
-
|
|
1375
|
-
|
|
1376
|
-
|
|
1377
|
-
|
|
1378
|
-
|
|
1379
|
-
|
|
1380
|
-
|
|
1381
|
-
|
|
1443
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
1444
|
+
var row_sum = 0.0;
|
|
1445
|
+
for (var i = 0u; i < 16u; i++) {
|
|
1446
|
+
let q_word = select(
|
|
1447
|
+
select(q_w0, q_w1, i >= 4u),
|
|
1448
|
+
select(q_w2, q_w3, i >= 12u),
|
|
1449
|
+
i >= 8u);
|
|
1450
|
+
let q_byte = get_byte(q_word, i % 4u);
|
|
1451
|
+
let nib = select(q_byte & 0xFu, (q_byte >> 4u) & 0xFu, half == 1u);
|
|
1452
|
+
row_sum += f32(kvalues_iq4nl[nib]) * dl * x_block[col][i];
|
|
1453
|
+
}
|
|
1454
|
+
acc[col][row] += row_sum;
|
|
1382
1455
|
}
|
|
1383
|
-
acc[row] += row_sum;
|
|
1384
1456
|
}
|
|
1385
1457
|
}
|
|
1386
1458
|
}
|
|
@@ -1394,35 +1466,84 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
|
|
1394
1466
|
#define BLOCK_SIZE_BYTES 17
|
|
1395
1467
|
#define THREADS_PER_BLOCK 4
|
|
1396
1468
|
#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
|
|
1397
|
-
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
|
|
1398
|
-
var acc: array<f32, OUTPUTS_PER_WG>;
|
|
1469
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
1470
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
1399
1471
|
|
|
1400
1472
|
let num_blocks = params.k / BLOCK_SIZE;
|
|
1401
1473
|
let thread_within_block = thread_id % 4;
|
|
1402
1474
|
for (var block = thread_id/THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE/THREADS_PER_BLOCK) {
|
|
1403
1475
|
let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
|
|
1404
|
-
var x_block: array<f32, ELEMS_PER_THREAD>;
|
|
1405
|
-
for (var
|
|
1406
|
-
|
|
1407
|
-
|
|
1476
|
+
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
|
1477
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
1478
|
+
for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
|
|
1479
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
1480
|
+
x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
|
|
1481
|
+
}
|
|
1408
1482
|
}
|
|
1409
|
-
|
|
1410
1483
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
1411
1484
|
let output_row = row_base + row;
|
|
1412
1485
|
if (output_row < params.m) {
|
|
1413
1486
|
let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
|
|
1414
1487
|
let eu8 = get_byte(load_u32_at_src0(block_byte_base), 0);
|
|
1415
1488
|
let e = ldexp(1.0, i32(eu8) - 128);
|
|
1416
|
-
var row_sum = 0.0;
|
|
1417
1489
|
let q_packed = load_u32_at_src0(block_byte_base + 1u + 4u * thread_within_block);
|
|
1418
|
-
for (var
|
|
1419
|
-
|
|
1420
|
-
|
|
1421
|
-
|
|
1422
|
-
|
|
1423
|
-
|
|
1490
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
1491
|
+
var row_sum = 0.0;
|
|
1492
|
+
for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
|
|
1493
|
+
let q_byte = get_byte(q_packed, byte_idx);
|
|
1494
|
+
let q_lo = f32(kvalues_mxfp4[q_byte & 0xFu]) * e;
|
|
1495
|
+
let q_hi = f32(kvalues_mxfp4[(q_byte >> 4u) & 0xFu]) * e;
|
|
1496
|
+
row_sum += q_lo * x_block[col][byte_idx];
|
|
1497
|
+
row_sum += q_hi * x_block[col][byte_idx + 4u];
|
|
1498
|
+
}
|
|
1499
|
+
acc[col][row] += row_sum;
|
|
1500
|
+
}
|
|
1501
|
+
}
|
|
1502
|
+
}
|
|
1503
|
+
}
|
|
1504
|
+
|
|
1505
|
+
return acc;
|
|
1506
|
+
}
|
|
1507
|
+
#endif
|
|
1508
|
+
|
|
1509
|
+
#ifdef MUL_ACC_NVFP4
|
|
1510
|
+
#define BLOCK_SIZE 64
|
|
1511
|
+
#define BLOCK_SIZE_BYTES 36
|
|
1512
|
+
#define THREADS_PER_BLOCK 4
|
|
1513
|
+
#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
|
|
1514
|
+
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
|
1515
|
+
var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
|
|
1516
|
+
|
|
1517
|
+
let num_blocks = params.k / BLOCK_SIZE;
|
|
1518
|
+
let sub = thread_id % THREADS_PER_BLOCK;
|
|
1519
|
+
for (var block = thread_id/THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE/THREADS_PER_BLOCK) {
|
|
1520
|
+
let x_base = src1_idx_base + block * BLOCK_SIZE + sub * ELEMS_PER_THREAD;
|
|
1521
|
+
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
|
1522
|
+
for (var col = 0u; col < NUM_COLS;col += 1) {
|
|
1523
|
+
for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
|
|
1524
|
+
x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
|
|
1525
|
+
x_block[col][i + 8] = f32(src1[x_base + col * params.stride_11 + i + 8]);
|
|
1526
|
+
}
|
|
1527
|
+
}
|
|
1528
|
+
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
|
1529
|
+
let output_row = row_base + row;
|
|
1530
|
+
if (output_row < params.m) {
|
|
1531
|
+
let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
|
|
1532
|
+
let d = ue4m3_to_fp32(get_byte(load_u32_at_src0_aligned(block_byte_base), sub)) * 0.5;
|
|
1533
|
+
let q_w0 = load_u32_at_src0_aligned(block_byte_base + 4u + 8u * sub);
|
|
1534
|
+
let q_w1 = load_u32_at_src0_aligned(block_byte_base + 8u + 8u * sub);
|
|
1535
|
+
for (var col = 0u;col < NUM_COLS;col += 1) {
|
|
1536
|
+
var row_sum = 0.0;
|
|
1537
|
+
for (var l = 0u; l < 8u; l++) {
|
|
1538
|
+
let q_word = select(q_w0, q_w1, l >= 4u);
|
|
1539
|
+
let q_byte = get_byte(q_word, l % 4u);
|
|
1540
|
+
let q_lo = f32(kvalues_mxfp4[q_byte & 0xFu]) * d;
|
|
1541
|
+
let q_hi = f32(kvalues_mxfp4[(q_byte >> 4u) & 0xFu]) * d;
|
|
1542
|
+
row_sum += q_lo * x_block[col][l];
|
|
1543
|
+
row_sum += q_hi * x_block[col][l + 8u];
|
|
1544
|
+
}
|
|
1545
|
+
acc[col][row] += row_sum;
|
|
1424
1546
|
}
|
|
1425
|
-
acc[row] += row_sum;
|
|
1426
1547
|
}
|
|
1427
1548
|
}
|
|
1428
1549
|
}
|