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
|
@@ -98,6 +98,7 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
98
98
|
}
|
|
99
99
|
#endif // INIT_SRC0_SHMEM_Q1_0
|
|
100
100
|
|
|
101
|
+
// legacy-quants
|
|
101
102
|
#if defined(INIT_SRC0_SHMEM_Q4_0) || defined(INIT_SRC0_SHMEM_Q4_1) || defined(INIT_SRC0_SHMEM_Q5_0) || defined(INIT_SRC0_SHMEM_Q5_1) || defined(INIT_SRC0_SHMEM_Q8_0) || defined(INIT_SRC0_SHMEM_Q8_1) || defined(INIT_SRC0_SHMEM_MXFP4)
|
|
102
103
|
const BLOCK_SIZE = 32u;
|
|
103
104
|
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
|
|
@@ -124,7 +125,7 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
124
125
|
if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
|
|
125
126
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
|
|
126
127
|
|
|
127
|
-
#
|
|
128
|
+
#if defined(INIT_SRC0_SHMEM_Q4_0)
|
|
128
129
|
let block_byte_base = src0_idx * 18u; // BLOCK_SIZE_BYTES = 18u;
|
|
129
130
|
let d = load_f16_at_src0(block_byte_base);
|
|
130
131
|
|
|
@@ -134,7 +135,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
134
135
|
let q_packed = load_u32_at_src0(q_byte_offset);
|
|
135
136
|
dequant_q4_0_packed_to_shmem(q_packed, d, shmem_idx + j * BYTES_PER_INNER_LOOP);
|
|
136
137
|
}
|
|
137
|
-
#
|
|
138
|
+
#endif // INIT_SRC0_SHMEM_Q4_0
|
|
139
|
+
|
|
140
|
+
#if defined(INIT_SRC0_SHMEM_Q4_1)
|
|
138
141
|
let block_byte_base = src0_idx * 20u; // BLOCK_SIZE_BYTES = 20u;
|
|
139
142
|
let dm = unpack2x16float(load_u32_at_src0_aligned(block_byte_base));
|
|
140
143
|
let d = f16(dm[0]);
|
|
@@ -153,7 +156,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
153
156
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
|
|
154
157
|
}
|
|
155
158
|
}
|
|
156
|
-
#
|
|
159
|
+
#endif // INIT_SRC0_SHMEM_Q4_1
|
|
160
|
+
|
|
161
|
+
#if defined(INIT_SRC0_SHMEM_Q5_0)
|
|
157
162
|
let block_byte_base = src0_idx * 22u; // BLOCK_SIZE_BYTES = 22u;
|
|
158
163
|
|
|
159
164
|
let d = load_f16_at_src0(block_byte_base);
|
|
@@ -176,7 +181,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
176
181
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
|
|
177
182
|
}
|
|
178
183
|
}
|
|
179
|
-
#
|
|
184
|
+
#endif // INIT_SRC0_SHMEM_Q5_0
|
|
185
|
+
|
|
186
|
+
#if defined(INIT_SRC0_SHMEM_Q5_1)
|
|
180
187
|
let block_byte_base = src0_idx * 24u; // BLOCK_SIZE_BYTES = 24u;
|
|
181
188
|
|
|
182
189
|
let dm = unpack2x16float(load_u32_at_src0_aligned(block_byte_base));
|
|
@@ -201,7 +208,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
201
208
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
|
|
202
209
|
}
|
|
203
210
|
}
|
|
204
|
-
#
|
|
211
|
+
#endif // INIT_SRC0_SHMEM_Q5_1
|
|
212
|
+
|
|
213
|
+
#if defined(INIT_SRC0_SHMEM_Q8_0)
|
|
205
214
|
let block_byte_base = src0_idx * 34u; // BLOCK_SIZE_BYTES = 34u;
|
|
206
215
|
let d = load_f16_at_src0(block_byte_base);
|
|
207
216
|
|
|
@@ -211,7 +220,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
211
220
|
let q_packed = load_u32_at_src0(q_byte_offset);
|
|
212
221
|
dequant_q8_0_packed_to_shmem(q_packed, d, shmem_idx + j * BYTES_PER_INNER_LOOP);
|
|
213
222
|
}
|
|
214
|
-
#
|
|
223
|
+
#endif // INIT_SRC0_SHMEM_Q8_0
|
|
224
|
+
|
|
225
|
+
#if defined(INIT_SRC0_SHMEM_Q8_1)
|
|
215
226
|
let block_byte_base = src0_idx * 36u; // BLOCK_SIZE_BYTES = 36u;
|
|
216
227
|
let dm = unpack2x16float(load_u32_at_src0_aligned(block_byte_base));
|
|
217
228
|
let d = f16(dm[0]);
|
|
@@ -227,8 +238,10 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
227
238
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_val;
|
|
228
239
|
}
|
|
229
240
|
}
|
|
230
|
-
#
|
|
231
|
-
|
|
241
|
+
#endif // INIT_SRC0_SHMEM_Q8_1
|
|
242
|
+
|
|
243
|
+
#if defined(INIT_SRC0_SHMEM_MXFP4)
|
|
244
|
+
let block_byte_base = src0_idx * 17u; // BLOCK_SIZE_BYTES = 17u;
|
|
232
245
|
let eu8 = get_byte(load_u32_at_src0_aligned(block_byte_base), block_byte_base & 3u);
|
|
233
246
|
let e = ldexp(1.0, i32(eu8) - 128);
|
|
234
247
|
|
|
@@ -244,11 +257,52 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
244
257
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = f16(q_hi);
|
|
245
258
|
}
|
|
246
259
|
}
|
|
247
|
-
#endif
|
|
260
|
+
#endif // INIT_SRC0_SHMEM_MXFP4
|
|
248
261
|
}
|
|
249
262
|
}
|
|
250
263
|
}
|
|
251
|
-
#endif
|
|
264
|
+
#endif // legacy-quants
|
|
265
|
+
|
|
266
|
+
#if defined(INIT_SRC0_SHMEM_NVFP4)
|
|
267
|
+
const BLOCK_SIZE = 64u;
|
|
268
|
+
const BLOCK_SIZE_BYTES = 36u;
|
|
269
|
+
const SUB_BLOCK_SIZE = 16u; // elements sharing one UE4M3 scale
|
|
270
|
+
const NQ = 16u;
|
|
271
|
+
const BYTES_PER_THREAD = 8u;
|
|
272
|
+
const BYTES_PER_INNER_LOOP = 4u;
|
|
273
|
+
|
|
274
|
+
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
|
|
275
|
+
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
|
|
276
|
+
let tile_m = i / TILE_K;
|
|
277
|
+
let tile_k_start = i % TILE_K;
|
|
278
|
+
let global_m = offset_m + tile_m;
|
|
279
|
+
let global_k_start = k_outer + tile_k_start;
|
|
280
|
+
|
|
281
|
+
if (global_m >= params.m) {
|
|
282
|
+
break;
|
|
283
|
+
}
|
|
284
|
+
|
|
285
|
+
let block_k = global_k_start / BLOCK_SIZE;
|
|
286
|
+
let sub_block = (global_k_start % BLOCK_SIZE) / SUB_BLOCK_SIZE;
|
|
287
|
+
let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
|
|
288
|
+
|
|
289
|
+
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
|
|
290
|
+
let d_byte_base = block_byte_base;
|
|
291
|
+
let qs_byte_base = block_byte_base + 4u;
|
|
292
|
+
|
|
293
|
+
let d = ue4m3_to_fp32(get_byte(load_u32_at_src0_aligned(d_byte_base), sub_block)) * 0.5;
|
|
294
|
+
|
|
295
|
+
for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j++) {
|
|
296
|
+
let q_packed = load_u32_at_src0_aligned(qs_byte_base + sub_block * 8u + j * 4u);
|
|
297
|
+
for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
|
|
298
|
+
let q_byte = get_byte(q_packed, k);
|
|
299
|
+
shmem[i + j * BYTES_PER_INNER_LOOP + k] = f16(f32(kvalues_mxfp4[q_byte & 0xF]) * d);
|
|
300
|
+
shmem[i + j * BYTES_PER_INNER_LOOP + k + 8u] = f16(f32(kvalues_mxfp4[(q_byte >> 4) & 0xF]) * d);
|
|
301
|
+
}
|
|
302
|
+
}
|
|
303
|
+
}
|
|
304
|
+
}
|
|
305
|
+
#endif // INIT_SRC0_SHMEM_NVFP4
|
|
252
306
|
|
|
253
307
|
// k-quants
|
|
254
308
|
#if defined(INIT_SRC0_SHMEM_Q2_K) || defined(INIT_SRC0_SHMEM_Q3_K) || defined(INIT_SRC0_SHMEM_Q4_K) || defined(INIT_SRC0_SHMEM_Q5_K) || defined(INIT_SRC0_SHMEM_Q6_K)
|
|
@@ -284,7 +338,7 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
284
338
|
|
|
285
339
|
let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
|
|
286
340
|
|
|
287
|
-
#
|
|
341
|
+
#if defined(INIT_SRC0_SHMEM_Q2_K)
|
|
288
342
|
let block_byte_base = src0_idx * 84u; // BLOCK_SIZE_BYTES = 84u;
|
|
289
343
|
let scales_byte_base = block_byte_base;
|
|
290
344
|
let qs_byte_base = block_byte_base + 16u;
|
|
@@ -314,7 +368,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
314
368
|
let ml = dmin * f16(scale >> 4u);
|
|
315
369
|
|
|
316
370
|
store_shmem_kquants(qs_vec4 * dl - ml, elem_idx);
|
|
317
|
-
#
|
|
371
|
+
#endif // INIT_SRC0_SHMEM_Q2_K
|
|
372
|
+
|
|
373
|
+
#if defined(INIT_SRC0_SHMEM_Q3_K)
|
|
318
374
|
let block_byte_base = src0_idx * 110u; // BLOCK_SIZE_BYTES = 110u;
|
|
319
375
|
let hmask_byte_base = block_byte_base + 0u;
|
|
320
376
|
let qs_byte_base = block_byte_base + 32u;
|
|
@@ -355,7 +411,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
355
411
|
let dl = d_all * (f16((scale_hi2 << 4u) | scale_low4) - 32.0);
|
|
356
412
|
|
|
357
413
|
store_shmem_kquants(dl * q_vec4, elem_idx);
|
|
358
|
-
#
|
|
414
|
+
#endif // INIT_SRC0_SHMEM_Q3_K
|
|
415
|
+
|
|
416
|
+
#if defined(INIT_SRC0_SHMEM_Q4_K)
|
|
359
417
|
let block_byte_base = src0_idx * 144u; // BLOCK_SIZE_BYTES = 144u;
|
|
360
418
|
let dm_byte_base = block_byte_base + 0u;
|
|
361
419
|
let scale_byte_base = block_byte_base + 4u;
|
|
@@ -399,7 +457,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
399
457
|
let ml = dmin * f16(mn);
|
|
400
458
|
|
|
401
459
|
store_shmem_kquants(dl * qs_vec4 - vec4(ml, ml, ml, ml), elem_idx);
|
|
402
|
-
#
|
|
460
|
+
#endif // INIT_SRC0_SHMEM_Q4_K
|
|
461
|
+
|
|
462
|
+
#if defined(INIT_SRC0_SHMEM_Q5_K)
|
|
403
463
|
let block_byte_base = src0_idx * 176u; // BLOCK_SIZE_BYTES = 176u;
|
|
404
464
|
let dm_byte_base = block_byte_base + 0u;
|
|
405
465
|
let scale_byte_base = block_byte_base + 4u;
|
|
@@ -456,7 +516,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
456
516
|
let ml = dmin * f16(mn);
|
|
457
517
|
|
|
458
518
|
store_shmem_kquants((qh_vec4 + qs_lo4_vec4) * dl - vec4<f16>(ml, ml, ml, ml), elem_idx);
|
|
459
|
-
#
|
|
519
|
+
#endif // INIT_SRC0_SHMEM_Q5_K
|
|
520
|
+
|
|
521
|
+
#if defined(INIT_SRC0_SHMEM_Q6_K)
|
|
460
522
|
let block_byte_base = src0_idx * 210u; // BLOCK_SIZE_BYTES = 210u;
|
|
461
523
|
let ql_byte_base = block_byte_base;
|
|
462
524
|
let qh_byte_base = block_byte_base + 128u;
|
|
@@ -497,17 +559,18 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
497
559
|
let scale = get_byte_i32(scale_word, scale_byte & 3u);
|
|
498
560
|
|
|
499
561
|
store_shmem_kquants(d * q_vec4 * f16(scale), elem_idx);
|
|
500
|
-
#endif
|
|
562
|
+
#endif // INIT_SRC0_SHMEM_Q6_K
|
|
501
563
|
}
|
|
502
564
|
}
|
|
503
565
|
#endif // k-quants
|
|
504
566
|
|
|
505
|
-
#
|
|
567
|
+
#if defined(INIT_SRC0_SHMEM_IQ4_NL)
|
|
506
568
|
const BLOCK_SIZE = 32u;
|
|
507
569
|
const BLOCK_SIZE_BYTES = 18u;
|
|
570
|
+
const NQ = 4u;
|
|
508
571
|
|
|
509
572
|
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
|
|
510
|
-
for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
|
|
573
|
+
for (var elem_idx = thread_id * NQ; elem_idx < TILE_SRC0_SHMEM; elem_idx += NQ * TOTAL_WORKGROUP_SIZE) {
|
|
511
574
|
let tile_m = elem_idx / TILE_K;
|
|
512
575
|
let tile_k = elem_idx % TILE_K;
|
|
513
576
|
let global_m = offset_m + tile_m;
|
|
@@ -519,408 +582,464 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
519
582
|
}
|
|
520
583
|
|
|
521
584
|
let block_k = global_k / BLOCK_SIZE;
|
|
522
|
-
let k_in_block = global_k % BLOCK_SIZE;
|
|
585
|
+
let k_in_block = global_k % BLOCK_SIZE; // k_in_block % 4 == 0;
|
|
586
|
+
|
|
587
|
+
let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
|
|
523
588
|
|
|
524
|
-
let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
|
|
525
589
|
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
|
|
526
|
-
let
|
|
590
|
+
let d_byte_base = block_byte_base + 0u;
|
|
591
|
+
let qs_byte_base = block_byte_base + 2u;
|
|
592
|
+
|
|
593
|
+
let d = load_f16_at_src0(d_byte_base);
|
|
594
|
+
|
|
595
|
+
let id_qtr = (k_in_block % 16u) / 4u;
|
|
596
|
+
let shift_phase = k_in_block / 16u;
|
|
527
597
|
|
|
528
|
-
let
|
|
529
|
-
let nib_shift = (k_in_block / 16u) * 4u;
|
|
530
|
-
let q_packed = load_u32_at_src0(block_byte_base + 2u + (pos / 4u) * 4u);
|
|
531
|
-
let nib = (get_byte(q_packed, pos % 4u) >> nib_shift) & 0xFu;
|
|
598
|
+
let qs_u32 = load_u32_at_src0(qs_byte_base + 4u * id_qtr);
|
|
532
599
|
|
|
533
|
-
shmem[elem_idx] = d * f16(kvalues_iq4nl[
|
|
600
|
+
shmem[elem_idx + 0u] = d * f16(kvalues_iq4nl[(qs_u32 >> ( 0u + 4u * shift_phase)) & 0xFu]);
|
|
601
|
+
shmem[elem_idx + 1u] = d * f16(kvalues_iq4nl[(qs_u32 >> ( 8u + 4u * shift_phase)) & 0xFu]);
|
|
602
|
+
shmem[elem_idx + 2u] = d * f16(kvalues_iq4nl[(qs_u32 >> (16u + 4u * shift_phase)) & 0xFu]);
|
|
603
|
+
shmem[elem_idx + 3u] = d * f16(kvalues_iq4nl[(qs_u32 >> (24u + 4u * shift_phase)) & 0xFu]);
|
|
534
604
|
}
|
|
535
605
|
}
|
|
536
606
|
#endif // INIT_SRC0_SHMEM_IQ4_NL
|
|
537
607
|
|
|
538
|
-
|
|
608
|
+
// i-quants (super block size: 256)
|
|
609
|
+
#if defined(INIT_SRC0_SHMEM_IQ4_XS) || defined(INIT_SRC0_SHMEM_IQ1_S) || defined(INIT_SRC0_SHMEM_IQ1_M) || defined(INIT_SRC0_SHMEM_IQ2_XXS) \
|
|
610
|
+
|| defined(INIT_SRC0_SHMEM_IQ2_XS) || defined(INIT_SRC0_SHMEM_IQ2_S) || defined(INIT_SRC0_SHMEM_IQ3_XXS) || defined(INIT_SRC0_SHMEM_IQ3_S)
|
|
539
611
|
const BLOCK_SIZE = 256u;
|
|
540
|
-
const
|
|
612
|
+
const NQ = 16u;
|
|
541
613
|
|
|
542
|
-
fn
|
|
543
|
-
|
|
544
|
-
|
|
545
|
-
|
|
546
|
-
|
|
547
|
-
|
|
614
|
+
fn store_shmem_iquants(val: vec4<f16>, idx: u32) {
|
|
615
|
+
shmem[idx] = val.x;
|
|
616
|
+
shmem[idx + 1] = val.y;
|
|
617
|
+
shmem[idx + 2] = val.z;
|
|
618
|
+
shmem[idx + 3] = val.w;
|
|
619
|
+
}
|
|
548
620
|
|
|
549
|
-
|
|
550
|
-
|
|
551
|
-
|
|
552
|
-
}
|
|
621
|
+
fn load_byte_at_src0_aligned(byte_offset: u32) -> u32 {
|
|
622
|
+
return get_byte(load_u32_at_src0_aligned(byte_offset), byte_offset % 4u);
|
|
623
|
+
}
|
|
553
624
|
|
|
554
|
-
|
|
555
|
-
|
|
625
|
+
#if defined(INIT_SRC0_SHMEM_IQ1_M) || defined(INIT_SRC0_SHMEM_IQ1_S)
|
|
626
|
+
fn create_iq_gw4(dl: f32, gw: u32, shift_base: u32, delta: f32) -> vec4<f16> {
|
|
627
|
+
return vec4<f16>(
|
|
628
|
+
f16(dl * (f32((bitcast<i32>(((gw >> (shift_base + 0u)) & 3u) << 30u) >> 30u)) + delta)),
|
|
629
|
+
f16(dl * (f32((bitcast<i32>(((gw >> (shift_base + 2u)) & 3u) << 30u) >> 30u)) + delta)),
|
|
630
|
+
f16(dl * (f32((bitcast<i32>(((gw >> (shift_base + 4u)) & 3u) << 30u) >> 30u)) + delta)),
|
|
631
|
+
f16(dl * (f32((bitcast<i32>(((gw >> (shift_base + 6u)) & 3u) << 30u) >> 30u)) + delta)),
|
|
632
|
+
);
|
|
633
|
+
}
|
|
634
|
+
#endif
|
|
556
635
|
|
|
557
|
-
|
|
558
|
-
|
|
636
|
+
#if defined(INIT_SRC0_SHMEM_IQ4_XS)
|
|
637
|
+
fn create_iq_gw4(dl: f16, qs_u32: u32, shift_phase: u32) -> vec4<f16> {
|
|
638
|
+
return vec4<f16>(
|
|
639
|
+
dl * f16(kvalues_iq4nl[(qs_u32 >> (4 * shift_phase + 0u)) & 0xFu]),
|
|
640
|
+
dl * f16(kvalues_iq4nl[(qs_u32 >> (4 * shift_phase + 8u)) & 0xFu]),
|
|
641
|
+
dl * f16(kvalues_iq4nl[(qs_u32 >> (4 * shift_phase + 16u)) & 0xFu]),
|
|
642
|
+
dl * f16(kvalues_iq4nl[(qs_u32 >> (4 * shift_phase + 24u)) & 0xFu]),
|
|
643
|
+
);
|
|
644
|
+
}
|
|
645
|
+
#endif
|
|
559
646
|
|
|
560
|
-
|
|
561
|
-
|
|
562
|
-
|
|
647
|
+
#if defined(INIT_SRC0_SHMEM_IQ2_XXS)
|
|
648
|
+
fn create_iq_gw4(ig: u32, grid_phase: u32) -> vec4<f32> {
|
|
649
|
+
return vec4<f32>(
|
|
650
|
+
f32(get_byte(iq2xxs_grid[(ig + grid_phase + 0u) / 4u], (ig + grid_phase + 0u) % 4u)),
|
|
651
|
+
f32(get_byte(iq2xxs_grid[(ig + grid_phase + 1u) / 4u], (ig + grid_phase + 1u) % 4u)),
|
|
652
|
+
f32(get_byte(iq2xxs_grid[(ig + grid_phase + 2u) / 4u], (ig + grid_phase + 2u) % 4u)),
|
|
653
|
+
f32(get_byte(iq2xxs_grid[(ig + grid_phase + 3u) / 4u], (ig + grid_phase + 3u) % 4u)),
|
|
654
|
+
);
|
|
655
|
+
}
|
|
656
|
+
#endif
|
|
563
657
|
|
|
564
|
-
|
|
565
|
-
|
|
658
|
+
#if defined(INIT_SRC0_SHMEM_IQ2_XS)
|
|
659
|
+
fn create_iq_gw4(ig: u32, grid_phase: u32) -> vec4<f32> {
|
|
660
|
+
return vec4<f32>(
|
|
661
|
+
f32(get_byte(iq2xs_grid[(ig + grid_phase + 0u) / 4u], (ig + grid_phase + 0u) % 4u)),
|
|
662
|
+
f32(get_byte(iq2xs_grid[(ig + grid_phase + 1u) / 4u], (ig + grid_phase + 1u) % 4u)),
|
|
663
|
+
f32(get_byte(iq2xs_grid[(ig + grid_phase + 2u) / 4u], (ig + grid_phase + 2u) % 4u)),
|
|
664
|
+
f32(get_byte(iq2xs_grid[(ig + grid_phase + 3u) / 4u], (ig + grid_phase + 3u) % 4u)),
|
|
665
|
+
);
|
|
666
|
+
}
|
|
667
|
+
#endif
|
|
566
668
|
|
|
567
|
-
|
|
568
|
-
|
|
569
|
-
|
|
570
|
-
|
|
669
|
+
#if defined(INIT_SRC0_SHMEM_IQ2_S)
|
|
670
|
+
fn create_iq_gw4(ig: u32, grid_phase: u32) -> vec4<f32> {
|
|
671
|
+
return vec4<f32>(
|
|
672
|
+
f32(get_byte(iq2s_grid[(ig + grid_phase + 0u) / 4u], (ig + grid_phase + 0u) % 4u)),
|
|
673
|
+
f32(get_byte(iq2s_grid[(ig + grid_phase + 1u) / 4u], (ig + grid_phase + 1u) % 4u)),
|
|
674
|
+
f32(get_byte(iq2s_grid[(ig + grid_phase + 2u) / 4u], (ig + grid_phase + 2u) % 4u)),
|
|
675
|
+
f32(get_byte(iq2s_grid[(ig + grid_phase + 3u) / 4u], (ig + grid_phase + 3u) % 4u)),
|
|
676
|
+
);
|
|
677
|
+
}
|
|
678
|
+
#endif
|
|
571
679
|
|
|
572
|
-
|
|
573
|
-
|
|
574
|
-
|
|
575
|
-
|
|
680
|
+
#if defined(INIT_SRC0_SHMEM_IQ3_XXS)
|
|
681
|
+
fn create_iq_gw4(ig: u32) -> vec4<f32> {
|
|
682
|
+
return vec4<f32>(
|
|
683
|
+
f32(get_byte(iq3xxs_grid[ig], 0)),
|
|
684
|
+
f32(get_byte(iq3xxs_grid[ig], 1)),
|
|
685
|
+
f32(get_byte(iq3xxs_grid[ig], 2)),
|
|
686
|
+
f32(get_byte(iq3xxs_grid[ig], 3)),
|
|
687
|
+
);
|
|
688
|
+
}
|
|
689
|
+
#endif
|
|
576
690
|
|
|
577
|
-
|
|
578
|
-
|
|
691
|
+
#if defined(INIT_SRC0_SHMEM_IQ3_S)
|
|
692
|
+
fn create_iq_gw4(ig: u32) -> vec4<f32> {
|
|
693
|
+
return vec4<f32>(
|
|
694
|
+
f32(get_byte(iq3s_grid[ig], 0)),
|
|
695
|
+
f32(get_byte(iq3s_grid[ig], 1)),
|
|
696
|
+
f32(get_byte(iq3s_grid[ig], 2)),
|
|
697
|
+
f32(get_byte(iq3s_grid[ig], 3)),
|
|
698
|
+
);
|
|
579
699
|
}
|
|
580
|
-
#endif
|
|
700
|
+
#endif
|
|
581
701
|
|
|
582
|
-
#
|
|
583
|
-
|
|
584
|
-
|
|
702
|
+
#if defined(INIT_SRC0_SHMEM_IQ2_XXS) || defined(INIT_SRC0_SHMEM_IQ2_XS) || defined(INIT_SRC0_SHMEM_IQ2_S) \
|
|
703
|
+
|| defined(INIT_SRC0_SHMEM_IQ3_XXS) || defined(INIT_SRC0_SHMEM_IQ3_S)
|
|
704
|
+
fn create_iq2_m4(signs: u32, mask_phase: u32) -> vec4<f32> {
|
|
705
|
+
return vec4<f32>(
|
|
706
|
+
select(1.0, -1.0, (get_byte(kmask_iq2xs[mask_phase], 0) & signs) != 0u),
|
|
707
|
+
select(1.0, -1.0, (get_byte(kmask_iq2xs[mask_phase], 1) & signs) != 0u),
|
|
708
|
+
select(1.0, -1.0, (get_byte(kmask_iq2xs[mask_phase], 2) & signs) != 0u),
|
|
709
|
+
select(1.0, -1.0, (get_byte(kmask_iq2xs[mask_phase], 3) & signs) != 0u),
|
|
710
|
+
);
|
|
711
|
+
}
|
|
712
|
+
#endif
|
|
585
713
|
|
|
586
714
|
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
|
|
587
|
-
for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
|
|
715
|
+
for (var elem_idx = thread_id * NQ; elem_idx < TILE_SRC0_SHMEM; elem_idx += NQ * TOTAL_WORKGROUP_SIZE) {
|
|
588
716
|
let tile_m = elem_idx / TILE_K;
|
|
589
717
|
let tile_k = elem_idx % TILE_K;
|
|
590
718
|
let global_m = offset_m + tile_m;
|
|
591
719
|
let global_k = k_outer + tile_k;
|
|
592
720
|
|
|
593
721
|
if (global_m >= params.m || global_k >= params.k) {
|
|
594
|
-
|
|
722
|
+
let zero_vec4 = vec4<f16>(f16(0.0), f16(0.0), f16(0.0), f16(0.0));
|
|
723
|
+
store_shmem_iquants(zero_vec4, elem_idx + 0u);
|
|
724
|
+
store_shmem_iquants(zero_vec4, elem_idx + 4u);
|
|
725
|
+
store_shmem_iquants(zero_vec4, elem_idx + 8u);
|
|
726
|
+
store_shmem_iquants(zero_vec4, elem_idx + 12u);
|
|
595
727
|
continue;
|
|
596
728
|
}
|
|
597
729
|
|
|
598
730
|
let block_k = global_k / BLOCK_SIZE;
|
|
599
|
-
let k_in_block = global_k % BLOCK_SIZE;
|
|
731
|
+
let k_in_block = global_k % BLOCK_SIZE; // k_in_block % 16 == 0;
|
|
600
732
|
|
|
601
|
-
let src0_idx
|
|
602
|
-
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
|
|
603
|
-
let d = load_f16_as_f32_at_src0(block_byte_base);
|
|
733
|
+
let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
|
|
604
734
|
|
|
605
|
-
|
|
606
|
-
let
|
|
607
|
-
let
|
|
608
|
-
let
|
|
735
|
+
#if defined(INIT_SRC0_SHMEM_IQ4_XS)
|
|
736
|
+
let block_byte_base = src0_idx * 136u; // BLOCK_SIZE_BYTES = 136u;
|
|
737
|
+
let d_byte_base = block_byte_base + 0u;
|
|
738
|
+
let scales_l_byte_base = block_byte_base + 4u;
|
|
739
|
+
let qs_byte_base = block_byte_base + 8u;
|
|
609
740
|
|
|
610
|
-
let
|
|
611
|
-
let
|
|
612
|
-
let
|
|
741
|
+
let d_scales_h = load_u32_at_src0_aligned(d_byte_base);
|
|
742
|
+
let d = bitcast<vec2<f16>>(d_scales_h).x;
|
|
743
|
+
let scales_h = d_scales_h >> 16u;
|
|
613
744
|
|
|
614
|
-
let
|
|
615
|
-
let
|
|
745
|
+
let sub_block = k_in_block / 32u;
|
|
746
|
+
let phase = (k_in_block / NQ) % 2u;
|
|
616
747
|
|
|
617
|
-
let
|
|
618
|
-
let
|
|
619
|
-
let
|
|
748
|
+
let scales_l_u32 = load_u32_at_src0_aligned(scales_l_byte_base);
|
|
749
|
+
let ls_lo = (get_byte(scales_l_u32, sub_block / 2u) >> (4u * (sub_block % 2u))) & 0xFu;
|
|
750
|
+
let ls_hi = ((scales_h >> (2u * sub_block)) & 3u) << 4u;
|
|
751
|
+
let dl = d * f16(i32(ls_lo | ls_hi) - 32);
|
|
620
752
|
|
|
621
|
-
|
|
622
|
-
|
|
623
|
-
|
|
624
|
-
|
|
753
|
+
let qs_0_3_u32 = load_u32_at_src0_aligned(qs_byte_base + 16u * sub_block + 0u);
|
|
754
|
+
let qs_4_7_u32 = load_u32_at_src0_aligned(qs_byte_base + 16u * sub_block + 4u);
|
|
755
|
+
let qs_8_11_u32 = load_u32_at_src0_aligned(qs_byte_base + 16u * sub_block + 8u);
|
|
756
|
+
let qs_12_15_u32 = load_u32_at_src0_aligned(qs_byte_base + 16u * sub_block + 12u);
|
|
625
757
|
|
|
626
|
-
|
|
627
|
-
|
|
628
|
-
|
|
758
|
+
store_shmem_iquants(create_iq_gw4(dl, qs_0_3_u32, phase), elem_idx + 0u);
|
|
759
|
+
store_shmem_iquants(create_iq_gw4(dl, qs_4_7_u32, phase), elem_idx + 4u);
|
|
760
|
+
store_shmem_iquants(create_iq_gw4(dl, qs_8_11_u32, phase), elem_idx + 8u);
|
|
761
|
+
store_shmem_iquants(create_iq_gw4(dl, qs_12_15_u32, phase), elem_idx + 12u);
|
|
762
|
+
#endif // INIT_SRC0_SHMEM_IQ4_XS
|
|
629
763
|
|
|
630
|
-
|
|
631
|
-
|
|
632
|
-
let
|
|
633
|
-
let
|
|
634
|
-
let
|
|
635
|
-
let global_k = k_outer + tile_k;
|
|
764
|
+
#if defined(INIT_SRC0_SHMEM_IQ1_S)
|
|
765
|
+
let block_byte_base = src0_idx * 50u; // BLOCK_SIZE_BYTES = 50u;
|
|
766
|
+
let d_byte_base = block_byte_base + 0u;
|
|
767
|
+
let qs_byte_base = block_byte_base + 2u;
|
|
768
|
+
let qh_byte_base = block_byte_base + 34u;
|
|
636
769
|
|
|
637
|
-
|
|
638
|
-
shmem[elem_idx] = f16(0.0);
|
|
639
|
-
continue;
|
|
640
|
-
}
|
|
770
|
+
let d = load_f16_as_f32_at_src0(d_byte_base);
|
|
641
771
|
|
|
642
|
-
let
|
|
643
|
-
let k_in_block
|
|
772
|
+
let sub_block = k_in_block / 32u;
|
|
773
|
+
let phase = (k_in_block / NQ) % 2u;
|
|
644
774
|
|
|
645
|
-
let
|
|
646
|
-
let
|
|
775
|
+
let qh_u16 = load_u32_at_src0(qh_byte_base + sub_block * 2u) & 0xFFFFu;
|
|
776
|
+
let qs_u16 = load_u32_at_src0(qs_byte_base + sub_block * 4u + phase * 2u) & 0xFFFFu;
|
|
777
|
+
|
|
778
|
+
let dl = d * (2.0 * f32((qh_u16 >> 12u) & 7u) + 1.0);
|
|
779
|
+
let delta = select(IQ1_DELTA, -IQ1_DELTA, (qh_u16 & 0x8000u) != 0u);
|
|
780
|
+
|
|
781
|
+
let gp0_grid_id = ((qs_u16 & 0xFFu) | (((qh_u16 >> (phase * 6u)) & 7u) << 8u)) * 8u;
|
|
782
|
+
let gp1_grid_id = (((qs_u16 >> 8) & 0xFFu) | (((qh_u16 >> (phase * 6u + 3u)) & 7u) << 8u)) * 8u;
|
|
783
|
+
|
|
784
|
+
let gp0_gw = iq1_grid[(gp0_grid_id) / 16u];
|
|
785
|
+
let gp1_gw = iq1_grid[(gp1_grid_id) / 16u];
|
|
786
|
+
|
|
787
|
+
let gp0_shift_base = (gp0_grid_id % 16u) * 2u;
|
|
788
|
+
let gp1_shift_base = (gp1_grid_id % 16u) * 2u;
|
|
647
789
|
|
|
648
|
-
|
|
649
|
-
|
|
790
|
+
store_shmem_iquants(create_iq_gw4(dl, gp0_gw, gp0_shift_base + 0u, delta), elem_idx + 0u);
|
|
791
|
+
store_shmem_iquants(create_iq_gw4(dl, gp0_gw, gp0_shift_base + 8u, delta), elem_idx + 4u);
|
|
792
|
+
store_shmem_iquants(create_iq_gw4(dl, gp1_gw, gp1_shift_base + 0u, delta), elem_idx + 8u);
|
|
793
|
+
store_shmem_iquants(create_iq_gw4(dl, gp1_gw, gp1_shift_base + 8u, delta), elem_idx + 12u);
|
|
794
|
+
#endif // INIT_SRC0_SHMEM_IQ1_S
|
|
795
|
+
|
|
796
|
+
#if defined(INIT_SRC0_SHMEM_IQ1_M)
|
|
797
|
+
let block_byte_base = src0_idx * 56u; // BLOCK_SIZE_BYTES = 56u;
|
|
798
|
+
let qs_byte_base = block_byte_base + 0u;
|
|
799
|
+
let qh_byte_base = block_byte_base + 32u;
|
|
800
|
+
let scales_byte_base = block_byte_base + 48u;
|
|
801
|
+
|
|
802
|
+
let scales0 = load_u32_at_src0_aligned(scales_byte_base);
|
|
803
|
+
let scales1 = load_u32_at_src0_aligned(scales_byte_base + 4u);
|
|
650
804
|
let scale_packed = ((scales0 >> 12u) & 0xFu) |
|
|
651
805
|
((scales0 >> 24u) & 0x00F0u) |
|
|
652
806
|
((scales1 >> 4u) & 0x0F00u) |
|
|
653
807
|
((scales1 >> 16u) & 0xF000u);
|
|
654
808
|
let d = f32(bitcast<vec2<f16>>(scale_packed).x);
|
|
655
809
|
|
|
656
|
-
let
|
|
657
|
-
let
|
|
658
|
-
let l = pos / 8u;
|
|
659
|
-
let j = pos % 8u;
|
|
810
|
+
let sub_block = k_in_block / 32u;
|
|
811
|
+
let phase = (k_in_block / NQ) % 2u;
|
|
660
812
|
|
|
661
|
-
let
|
|
662
|
-
let
|
|
663
|
-
let
|
|
664
|
-
let dl = d * f32(2u * s_pair + 1u);
|
|
813
|
+
let scale_u32 = select(scales0, scales1, sub_block >= 4u);
|
|
814
|
+
let scale_u3 = (scale_u32 >> (16u * ((sub_block / 2u) % 2u) + 6u * (sub_block % 2u) + 3u * phase)) & 0x7u;
|
|
815
|
+
let dl = d * f32(2u * scale_u3 + 1u);
|
|
665
816
|
|
|
666
|
-
let
|
|
667
|
-
let
|
|
668
|
-
let qh_nib = (qh >> (4u * l)) & 0xFu;
|
|
817
|
+
let qh_u8 = (load_u32_at_src0_aligned(qh_byte_base + 4u * (sub_block / 2u)) >> (16u * (sub_block % 2u) + 8u * phase)) & 0xFFu;
|
|
818
|
+
let qs_u16 = (load_u32_at_src0_aligned(qs_byte_base + 4u * sub_block) >> (16u * phase)) & 0xFFFFu;
|
|
669
819
|
|
|
670
|
-
let
|
|
671
|
-
let
|
|
672
|
-
let delta = select(IQ1_DELTA, -IQ1_DELTA, (qh_nib & 0x8u) != 0u);
|
|
820
|
+
let gp0_grid_id = ((qs_u16 & 0xFFu) | ((qh_u8 & 7u) << 8u)) * 8u;
|
|
821
|
+
let gp0_delta = select(IQ1_DELTA, -IQ1_DELTA, (qh_u8 & 0x8u) != 0u);
|
|
673
822
|
|
|
674
|
-
let
|
|
675
|
-
let
|
|
676
|
-
let g = (gw >> (((ig + j) % 16u) * 2u)) & 3u;
|
|
677
|
-
let gs = bitcast<i32>(g << 30u) >> 30u;
|
|
823
|
+
let gp1_grid_id = (((qs_u16 >> 8u) & 0xFFu) | (((qh_u8 >> 4u) & 7u) << 8u)) * 8u;
|
|
824
|
+
let gp1_delta = select(IQ1_DELTA, -IQ1_DELTA, (qh_u8 & 0x80u) != 0u);
|
|
678
825
|
|
|
679
|
-
|
|
680
|
-
|
|
681
|
-
}
|
|
682
|
-
#endif // INIT_SRC0_SHMEM_IQ1_M
|
|
826
|
+
let gp0_gw = iq1_grid[(gp0_grid_id) / 16u];
|
|
827
|
+
let gp1_gw = iq1_grid[(gp1_grid_id) / 16u];
|
|
683
828
|
|
|
684
|
-
|
|
685
|
-
|
|
686
|
-
const BLOCK_SIZE_BYTES = 66u;
|
|
829
|
+
let gp0_shift_base = (gp0_grid_id % 16u) * 2u;
|
|
830
|
+
let gp1_shift_base = (gp1_grid_id % 16u) * 2u;
|
|
687
831
|
|
|
688
|
-
|
|
689
|
-
|
|
690
|
-
|
|
691
|
-
|
|
692
|
-
|
|
693
|
-
let global_k = k_outer + tile_k;
|
|
832
|
+
store_shmem_iquants(create_iq_gw4(dl, gp0_gw, gp0_shift_base + 0u, gp0_delta), elem_idx + 0u);
|
|
833
|
+
store_shmem_iquants(create_iq_gw4(dl, gp0_gw, gp0_shift_base + 8u, gp0_delta), elem_idx + 4u);
|
|
834
|
+
store_shmem_iquants(create_iq_gw4(dl, gp1_gw, gp1_shift_base + 0u, gp1_delta), elem_idx + 8u);
|
|
835
|
+
store_shmem_iquants(create_iq_gw4(dl, gp1_gw, gp1_shift_base + 8u, gp1_delta), elem_idx + 12u);
|
|
836
|
+
#endif // INIT_SRC0_SHMEM_IQ1_M
|
|
694
837
|
|
|
695
|
-
|
|
696
|
-
|
|
697
|
-
|
|
698
|
-
|
|
838
|
+
#if defined(INIT_SRC0_SHMEM_IQ2_XXS)
|
|
839
|
+
let block_byte_base = src0_idx * 66u; // BLOCK_SIZE_BYTES = 66u;
|
|
840
|
+
let d_byte_base = block_byte_base + 0u;
|
|
841
|
+
let qs_byte_base = block_byte_base + 2u;
|
|
699
842
|
|
|
700
|
-
let
|
|
701
|
-
let k_in_block = global_k % BLOCK_SIZE;
|
|
843
|
+
let d = load_f16_as_f32_at_src0(d_byte_base);
|
|
702
844
|
|
|
703
|
-
let
|
|
704
|
-
let
|
|
705
|
-
let d = load_f16_as_f32_at_src0(block_byte_base);
|
|
845
|
+
let sub_block = k_in_block / 32u;
|
|
846
|
+
let phase = (k_in_block / NQ) % 2u;
|
|
706
847
|
|
|
707
|
-
let
|
|
708
|
-
let
|
|
848
|
+
let aux0 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 0u);
|
|
849
|
+
let aux1 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 4u);
|
|
850
|
+
let db = d * (0.5 + f32(aux1 >> 28u)) * 0.25;
|
|
709
851
|
|
|
710
|
-
let
|
|
711
|
-
let
|
|
852
|
+
let gp0_ig = get_byte(aux0, 2u * phase + 0u) * 8u;
|
|
853
|
+
let gp1_ig = get_byte(aux0, 2u * phase + 1u) * 8u;
|
|
712
854
|
|
|
713
|
-
let
|
|
714
|
-
let
|
|
715
|
-
let db = d * (0.5 + f32(aux1 >> 28u)) * 0.25;
|
|
855
|
+
let gp0_is = (aux1 >> (14u * phase + 0u)) & 127u;
|
|
856
|
+
let gp1_is = (aux1 >> (14u * phase + 7u)) & 127u;
|
|
716
857
|
|
|
717
|
-
let
|
|
718
|
-
let
|
|
719
|
-
let signs = get_byte(ksigns_iq2xs[is / 4u], is % 4u);
|
|
858
|
+
let gp0_signs = get_byte(ksigns_iq2xs[gp0_is / 4u], gp0_is % 4u);
|
|
859
|
+
let gp1_signs = get_byte(ksigns_iq2xs[gp1_is / 4u], gp1_is % 4u);
|
|
720
860
|
|
|
721
|
-
let
|
|
722
|
-
let
|
|
861
|
+
let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
|
|
862
|
+
let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
|
|
863
|
+
let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
|
|
864
|
+
let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
|
|
723
865
|
|
|
724
|
-
|
|
725
|
-
|
|
726
|
-
|
|
866
|
+
let gw_0_3_val4 = create_iq_gw4(gp0_ig, 0);
|
|
867
|
+
let gw_4_7_val4 = create_iq_gw4(gp0_ig, 4);
|
|
868
|
+
let gw_8_11_val4 = create_iq_gw4(gp1_ig, 0);
|
|
869
|
+
let gw_12_15_val4 = create_iq_gw4(gp1_ig, 4);
|
|
870
|
+
|
|
871
|
+
store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
|
|
872
|
+
store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
|
|
873
|
+
store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
|
|
874
|
+
store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
|
|
727
875
|
#endif // INIT_SRC0_SHMEM_IQ2_XXS
|
|
728
876
|
|
|
729
|
-
#
|
|
730
|
-
|
|
731
|
-
|
|
877
|
+
#if defined(INIT_SRC0_SHMEM_IQ2_XS)
|
|
878
|
+
let block_byte_base = src0_idx * 74u; // BLOCK_SIZE_BYTES = 74u;
|
|
879
|
+
let d_byte_base = block_byte_base + 0u;
|
|
880
|
+
let qs_byte_base = block_byte_base + 2u;
|
|
881
|
+
let scales_byte_base = block_byte_base + 66u;
|
|
732
882
|
|
|
733
|
-
|
|
734
|
-
for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
|
|
735
|
-
let tile_m = elem_idx / TILE_K;
|
|
736
|
-
let tile_k = elem_idx % TILE_K;
|
|
737
|
-
let global_m = offset_m + tile_m;
|
|
738
|
-
let global_k = k_outer + tile_k;
|
|
883
|
+
let d = load_f16_as_f32_at_src0(d_byte_base);
|
|
739
884
|
|
|
740
|
-
|
|
741
|
-
|
|
742
|
-
continue;
|
|
743
|
-
}
|
|
885
|
+
let sub_block = k_in_block / 32u;
|
|
886
|
+
let phase = (k_in_block / NQ) % 2u;
|
|
744
887
|
|
|
745
|
-
let
|
|
746
|
-
let
|
|
888
|
+
let scale = (load_byte_at_src0_aligned(scales_byte_base + 1u * sub_block) >> (4u * phase)) & 0xFu;
|
|
889
|
+
let db = d * (0.5 + f32(scale)) * 0.25;
|
|
747
890
|
|
|
748
|
-
let
|
|
749
|
-
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
|
|
750
|
-
let d = load_f16_as_f32_at_src0(block_byte_base);
|
|
891
|
+
let qs_u32 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 4u * phase);
|
|
751
892
|
|
|
752
|
-
let
|
|
753
|
-
let
|
|
893
|
+
let gp0_ig = (qs_u32 & 0x1FFu) * 8u;
|
|
894
|
+
let gp1_ig = ((qs_u32 >> 16u) & 0x1FFu) * 8u;
|
|
754
895
|
|
|
755
|
-
let
|
|
756
|
-
let
|
|
896
|
+
let gp0_is = (qs_u32 >> 9u) & 0x7Fu;
|
|
897
|
+
let gp1_is = (qs_u32 >> 25u) & 0x7Fu;
|
|
757
898
|
|
|
758
|
-
let
|
|
759
|
-
let
|
|
760
|
-
let s_nib = select(s & 0xFu, (s >> 4u) & 0xFu, (l / 2u) != 0u);
|
|
761
|
-
let dl = d * (0.5 + f32(s_nib)) * 0.25;
|
|
899
|
+
let gp0_signs = get_byte(ksigns_iq2xs[gp0_is / 4u], gp0_is % 4u);
|
|
900
|
+
let gp1_signs = get_byte(ksigns_iq2xs[gp1_is / 4u], gp1_is % 4u);
|
|
762
901
|
|
|
763
|
-
let
|
|
764
|
-
let
|
|
765
|
-
let
|
|
766
|
-
let
|
|
767
|
-
let signs = get_byte(ksigns_iq2xs[is / 4u], is % 4u);
|
|
902
|
+
let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
|
|
903
|
+
let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
|
|
904
|
+
let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
|
|
905
|
+
let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
|
|
768
906
|
|
|
769
|
-
let
|
|
770
|
-
let
|
|
907
|
+
let gw_0_3_val4 = create_iq_gw4(gp0_ig, 0);
|
|
908
|
+
let gw_4_7_val4 = create_iq_gw4(gp0_ig, 4);
|
|
909
|
+
let gw_8_11_val4 = create_iq_gw4(gp1_ig, 0);
|
|
910
|
+
let gw_12_15_val4 = create_iq_gw4(gp1_ig, 4);
|
|
771
911
|
|
|
772
|
-
|
|
773
|
-
|
|
774
|
-
|
|
912
|
+
store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
|
|
913
|
+
store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
|
|
914
|
+
store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
|
|
915
|
+
store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
|
|
775
916
|
#endif // INIT_SRC0_SHMEM_IQ2_XS
|
|
776
917
|
|
|
777
|
-
#
|
|
778
|
-
|
|
779
|
-
|
|
918
|
+
#if defined(INIT_SRC0_SHMEM_IQ2_S)
|
|
919
|
+
let block_byte_base = src0_idx * 82u; // BLOCK_SIZE_BYTES = 82u;
|
|
920
|
+
let d_byte_base = block_byte_base + 0u;
|
|
921
|
+
let qs_byte_base = block_byte_base + 2u;
|
|
922
|
+
let qh_byte_base = block_byte_base + 66u;
|
|
923
|
+
let scales_byte_base = block_byte_base + 74u;
|
|
780
924
|
|
|
781
|
-
|
|
782
|
-
for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
|
|
783
|
-
let tile_m = elem_idx / TILE_K;
|
|
784
|
-
let tile_k = elem_idx % TILE_K;
|
|
785
|
-
let global_m = offset_m + tile_m;
|
|
786
|
-
let global_k = k_outer + tile_k;
|
|
925
|
+
let d = load_f16_as_f32_at_src0(d_byte_base);
|
|
787
926
|
|
|
788
|
-
|
|
789
|
-
|
|
790
|
-
continue;
|
|
791
|
-
}
|
|
792
|
-
|
|
793
|
-
let block_k = global_k / BLOCK_SIZE;
|
|
794
|
-
let k_in_block = global_k % BLOCK_SIZE;
|
|
927
|
+
let sub_block = k_in_block / 32u;
|
|
928
|
+
let phase = (k_in_block / NQ) % 2u;
|
|
795
929
|
|
|
796
|
-
let
|
|
797
|
-
let
|
|
798
|
-
let d = load_f16_as_f32_at_src0(block_byte_base);
|
|
930
|
+
let scale = (load_byte_at_src0_aligned(scales_byte_base + 1u * sub_block) >> (4u * phase)) & 0xFu;
|
|
931
|
+
let db = d * (0.5 + f32(scale)) * 0.25;
|
|
799
932
|
|
|
800
|
-
let
|
|
801
|
-
let
|
|
802
|
-
let
|
|
933
|
+
let qs_u16 = load_u32_at_src0(qs_byte_base + 4u * sub_block + 2u * phase) & 0xFFFFu;
|
|
934
|
+
let signs_u16 = load_u32_at_src0(qs_byte_base + 32u + 4u * sub_block + 2u * phase) & 0xFFFFu;
|
|
935
|
+
let qh_u4 = (load_byte_at_src0_aligned(qh_byte_base + 1u * sub_block) >> (4u * phase)) & 0xFu;
|
|
803
936
|
|
|
804
|
-
let
|
|
805
|
-
let
|
|
806
|
-
let s_nib = select(s & 0xFu, (s >> 4u) & 0xFu, (l / 2u) != 0u);
|
|
807
|
-
let dl = d * (0.5 + f32(s_nib)) * 0.25;
|
|
937
|
+
let gp0_ig = ((qs_u16 & 0xFFu) | ((qh_u4 & 0x3u) << 8u)) * 8u;
|
|
938
|
+
let gp1_ig = (((qs_u16 >> 8u) & 0xFFu) | ((qh_u4 & 0xCu) << 6u)) * 8u;
|
|
808
939
|
|
|
809
|
-
let
|
|
810
|
-
let
|
|
811
|
-
let qh_b = (get_byte(qh_word, ib % 4u) << (8u - 2u * l)) & 0x300u;
|
|
812
|
-
let ig = (get_byte(qs_word, l) | qh_b) * 8u;
|
|
940
|
+
let gp0_signs = get_byte(signs_u16, 0);
|
|
941
|
+
let gp1_signs = get_byte(signs_u16, 1);
|
|
813
942
|
|
|
814
|
-
let
|
|
815
|
-
let
|
|
943
|
+
let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
|
|
944
|
+
let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
|
|
945
|
+
let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
|
|
946
|
+
let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
|
|
816
947
|
|
|
817
|
-
let
|
|
818
|
-
let
|
|
948
|
+
let gw_0_3_val4 = create_iq_gw4(gp0_ig, 0);
|
|
949
|
+
let gw_4_7_val4 = create_iq_gw4(gp0_ig, 4);
|
|
950
|
+
let gw_8_11_val4 = create_iq_gw4(gp1_ig, 0);
|
|
951
|
+
let gw_12_15_val4 = create_iq_gw4(gp1_ig, 4);
|
|
819
952
|
|
|
820
|
-
|
|
821
|
-
|
|
822
|
-
|
|
953
|
+
store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
|
|
954
|
+
store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
|
|
955
|
+
store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
|
|
956
|
+
store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
|
|
823
957
|
#endif // INIT_SRC0_SHMEM_IQ2_S
|
|
824
958
|
|
|
825
|
-
#
|
|
826
|
-
|
|
827
|
-
|
|
959
|
+
#if defined(INIT_SRC0_SHMEM_IQ3_XXS)
|
|
960
|
+
let block_byte_base = src0_idx * 98u; // BLOCK_SIZE_BYTES = 98u;
|
|
961
|
+
let d_byte_base = block_byte_base + 0u;
|
|
962
|
+
let qs_byte_base = block_byte_base + 2u;
|
|
828
963
|
|
|
829
|
-
|
|
830
|
-
for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
|
|
831
|
-
let tile_m = elem_idx / TILE_K;
|
|
832
|
-
let tile_k = elem_idx % TILE_K;
|
|
833
|
-
let global_m = offset_m + tile_m;
|
|
834
|
-
let global_k = k_outer + tile_k;
|
|
964
|
+
let d = load_f16_as_f32_at_src0(d_byte_base);
|
|
835
965
|
|
|
836
|
-
|
|
837
|
-
|
|
838
|
-
continue;
|
|
839
|
-
}
|
|
966
|
+
let sub_block = k_in_block / 32u;
|
|
967
|
+
let phase = (k_in_block / NQ) % 2u;
|
|
840
968
|
|
|
841
|
-
let
|
|
842
|
-
let
|
|
969
|
+
let qs_u32 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 4u * phase);
|
|
970
|
+
let sign_u32 = load_u32_at_src0(qs_byte_base + 64u + 4u * sub_block);
|
|
971
|
+
let db = d * (0.5 + f32(sign_u32 >> 28u)) * 0.5;
|
|
843
972
|
|
|
844
|
-
let
|
|
845
|
-
let
|
|
846
|
-
let
|
|
847
|
-
|
|
848
|
-
let ib_pair = k_in_block / 32u;
|
|
849
|
-
let in_pair = k_in_block % 32u;
|
|
850
|
-
let l = in_pair / 8u;
|
|
851
|
-
let in_l = in_pair % 8u;
|
|
852
|
-
let k2 = in_l / 4u;
|
|
853
|
-
let j = in_l % 4u;
|
|
854
|
-
|
|
855
|
-
let ib = ib_pair * 2u;
|
|
856
|
-
let sc_sign_off = block_byte_base + 2u + (ib + 32u) * 2u;
|
|
857
|
-
let sc_sign = load_u32_at_src0(sc_sign_off);
|
|
858
|
-
let db = d * (0.5 + f32(sc_sign >> 28u)) * 0.5;
|
|
859
|
-
let is = (sc_sign >> (7u * l)) & 127u;
|
|
860
|
-
let signs = get_byte(ksigns_iq2xs[is / 4u], is % 4u);
|
|
861
|
-
|
|
862
|
-
let ig_word = load_u32_at_src0(block_byte_base + 2u + (ib * 2u + l) * 2u) & 0xFFFFu;
|
|
863
|
-
let ig_byte = get_byte(ig_word, k2);
|
|
864
|
-
let g = get_byte(iq3xxs_grid[ig_byte], j);
|
|
865
|
-
let m = select(1.0, -1.0, (get_byte(kmask_iq2xs[k2], j) & signs) != 0u);
|
|
866
|
-
|
|
867
|
-
shmem[elem_idx] = f16(db * f32(g) * m);
|
|
868
|
-
}
|
|
869
|
-
}
|
|
870
|
-
#endif // INIT_SRC0_SHMEM_IQ3_XXS
|
|
973
|
+
let ig_0_3 = get_byte(qs_u32, 0);
|
|
974
|
+
let ig_4_7 = get_byte(qs_u32, 1);
|
|
975
|
+
let ig_8_11 = get_byte(qs_u32, 2);
|
|
976
|
+
let ig_12_15 = get_byte(qs_u32, 3);
|
|
871
977
|
|
|
872
|
-
|
|
873
|
-
|
|
874
|
-
const BLOCK_SIZE_BYTES = 110u;
|
|
978
|
+
let gp0_is = (sign_u32 >> (14u * phase + 0u)) & 0x7Fu;
|
|
979
|
+
let gp1_is = (sign_u32 >> (14u * phase + 7u)) & 0x7Fu;
|
|
875
980
|
|
|
876
|
-
|
|
877
|
-
|
|
878
|
-
let tile_m = elem_idx / TILE_K;
|
|
879
|
-
let tile_k = elem_idx % TILE_K;
|
|
880
|
-
let global_m = offset_m + tile_m;
|
|
881
|
-
let global_k = k_outer + tile_k;
|
|
981
|
+
let gp0_signs = get_byte(ksigns_iq2xs[gp0_is / 4u], gp0_is % 4u);
|
|
982
|
+
let gp1_signs = get_byte(ksigns_iq2xs[gp1_is / 4u], gp1_is % 4u);
|
|
882
983
|
|
|
883
|
-
|
|
884
|
-
|
|
885
|
-
|
|
886
|
-
|
|
984
|
+
let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
|
|
985
|
+
let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
|
|
986
|
+
let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
|
|
987
|
+
let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
|
|
887
988
|
|
|
888
|
-
let
|
|
889
|
-
let
|
|
890
|
-
|
|
891
|
-
let
|
|
892
|
-
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
|
|
893
|
-
let d = load_f16_as_f32_at_src0(block_byte_base);
|
|
894
|
-
|
|
895
|
-
let ib = k_in_block / 64u;
|
|
896
|
-
let rest = k_in_block % 64u;
|
|
897
|
-
let k = rest / 32u;
|
|
898
|
-
let in_k = rest % 32u;
|
|
899
|
-
let l = in_k / 8u;
|
|
900
|
-
let in_l = in_k % 8u;
|
|
901
|
-
let k2 = in_l / 4u;
|
|
902
|
-
let j = in_l % 4u;
|
|
903
|
-
|
|
904
|
-
let scales_word = load_u32_at_src0(block_byte_base + 106u);
|
|
905
|
-
let s = get_byte(scales_word, ib);
|
|
906
|
-
let s_nib = select(s & 0xFu, (s >> 4u) & 0xFu, k != 0u);
|
|
907
|
-
let dl = d * (1.0 + 2.0 * f32(s_nib));
|
|
908
|
-
|
|
909
|
-
let qh_word = load_u32_at_src0(block_byte_base + 66u + (ib / 2u) * 4u);
|
|
910
|
-
let qh_byte = get_byte(qh_word, (ib % 2u) * 2u + k);
|
|
989
|
+
let gw_0_3_val4 = create_iq_gw4(ig_0_3);
|
|
990
|
+
let gw_4_7_val4 = create_iq_gw4(ig_4_7);
|
|
991
|
+
let gw_8_11_val4 = create_iq_gw4(ig_8_11);
|
|
992
|
+
let gw_12_15_val4 = create_iq_gw4(ig_12_15);
|
|
911
993
|
|
|
912
|
-
|
|
913
|
-
|
|
914
|
-
|
|
915
|
-
|
|
916
|
-
|
|
917
|
-
let signs_word = load_u32_at_src0(block_byte_base + 74u + (ib * 2u + k) * 4u);
|
|
918
|
-
let signs = get_byte(signs_word, l);
|
|
919
|
-
|
|
920
|
-
let g = get_byte(iq3s_grid[ig], j);
|
|
921
|
-
let m = select(1.0, -1.0, (get_byte(kmask_iq2xs[k2], j) & signs) != 0u);
|
|
994
|
+
store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
|
|
995
|
+
store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
|
|
996
|
+
store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
|
|
997
|
+
store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
|
|
998
|
+
#endif // INIT_SRC0_SHMEM_IQ3_XXS
|
|
922
999
|
|
|
923
|
-
|
|
1000
|
+
#if defined(INIT_SRC0_SHMEM_IQ3_S)
|
|
1001
|
+
let block_byte_base = src0_idx * 110u; // BLOCK_SIZE_BYTES = 110u;
|
|
1002
|
+
let d_byte_base = block_byte_base + 0u;
|
|
1003
|
+
let qs_byte_base = block_byte_base + 2u;
|
|
1004
|
+
let qh_byte_base = block_byte_base + 66u;
|
|
1005
|
+
let signs_byte_base = block_byte_base + 74u;
|
|
1006
|
+
let scales_byte_base = block_byte_base + 106u;
|
|
1007
|
+
|
|
1008
|
+
let d = load_f16_as_f32_at_src0(d_byte_base);
|
|
1009
|
+
|
|
1010
|
+
let sub_block = k_in_block / 32u;
|
|
1011
|
+
let phase = (k_in_block / NQ) % 2u;
|
|
1012
|
+
|
|
1013
|
+
let scale = (load_byte_at_src0_aligned(scales_byte_base + 1u * (sub_block / 2u)) >> (4u * (sub_block % 2u))) & 0xFu;
|
|
1014
|
+
let db = d * (1.0 + 2.0 * f32(scale));
|
|
1015
|
+
|
|
1016
|
+
let qs_u32 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 4u * phase);
|
|
1017
|
+
let qh_u4 = (load_byte_at_src0_aligned(qh_byte_base + 1u * sub_block) >> (4u * phase)) & 0xFu;
|
|
1018
|
+
let signs_u16 = (load_u32_at_src0(signs_byte_base + 4u * sub_block + 2u * phase)) & 0xFFFFu;
|
|
1019
|
+
|
|
1020
|
+
let ig_0_3 = ((qs_u32 >> 0u) & 0xFFu) | ((qh_u4 & 0x1u) << 8u);
|
|
1021
|
+
let ig_4_7 = ((qs_u32 >> 8u) & 0xFFu) | ((qh_u4 & 0x2u) << 7u);
|
|
1022
|
+
let ig_8_11 = ((qs_u32 >> 16u) & 0xFFu) | ((qh_u4 & 0x4u) << 6u);
|
|
1023
|
+
let ig_12_15 = ((qs_u32 >> 24u) & 0xFFu) | ((qh_u4 & 0x8u) << 5u);
|
|
1024
|
+
|
|
1025
|
+
let gp0_signs = get_byte(signs_u16, 0);
|
|
1026
|
+
let gp1_signs = get_byte(signs_u16, 1);
|
|
1027
|
+
|
|
1028
|
+
let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
|
|
1029
|
+
let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
|
|
1030
|
+
let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
|
|
1031
|
+
let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
|
|
1032
|
+
|
|
1033
|
+
let gw_0_3_val4 = create_iq_gw4(ig_0_3);
|
|
1034
|
+
let gw_4_7_val4 = create_iq_gw4(ig_4_7);
|
|
1035
|
+
let gw_8_11_val4 = create_iq_gw4(ig_8_11);
|
|
1036
|
+
let gw_12_15_val4 = create_iq_gw4(ig_12_15);
|
|
1037
|
+
|
|
1038
|
+
store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
|
|
1039
|
+
store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
|
|
1040
|
+
store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
|
|
1041
|
+
store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
|
|
1042
|
+
#endif // INIT_SRC0_SHMEM_IQ3_S
|
|
924
1043
|
}
|
|
925
1044
|
}
|
|
926
|
-
#endif //
|
|
1045
|
+
#endif // i-quants (super block size: 256)
|