Brunobkr/llama.cpp_AlgMor24_github
ΩFFFΣLLIa • llama.cpp • AlgMor24 ██████╗ ███████╗███████╗███████╗██╗ ██╗ ██╗ █████╗ ██╔═══██╗██╔════╝██╔════╝██╔════╝██║ ██║ ██║██╔══██╗ ██║ ██║█████╗ █████╗ █████╗ ██║ ██║ ██║███████║ ██║ ██║██╔══╝ ██╔══╝ ██╔══╝ ██║ ██║ ██║██╔══██║ ╚██████╔╝██║ ██║ ███████╗███████╗███████╗██║██║ ██║ ╚═════╝ ╚═╝ ╚═╝ ╚══════╝╚══════╝╚══════╝╚═╝╚═╝ ╚═╝ High-Performance LLM / VLM Inference & Autonomous Agentic Ecosystem… See the full description on the dataset page: https://huggingface.co/datasets/Brunobkr/llama.cpp_AlgMor24_github.
03.1k
1#ifndef GGML_WEBGPU_SHADER_LIB_HPP2#define GGML_WEBGPU_SHADER_LIB_HPP3 4#include "ggml-impl.h"5#include "ggml-wgsl-shaders.hpp"6#include "ggml.h"7#include "pre_wgsl.hpp"8 9#include <webgpu/webgpu_cpp.h>10 11#include <algorithm>12#include <memory>13#include <string>14#include <unordered_map>15#include <vector>16 17#define GGML_WEBGPU_F16_SIZE_BYTES 218#define GGML_WEBGPU_F32_SIZE_BYTES 419#define GGML_WEBGPU_I32_SIZE_BYTES 420#define GGML_WEBGPU_FLASH_ATTN_PREFERRED_KV_SG_TILES 8u21#define GGML_WEBGPU_FLASH_ATTN_VEC_MAX_SEQ_LEN 20u22#define GGML_WEBGPU_FLASH_ATTN_VEC_MAX_KV_TILE 32u23#define GGML_WEBGPU_FLASH_ATTN_TILE_MAX_KV_TILE 64u24#define GGML_WEBGPU_FLASH_ATTN_PREFERRED_WG_SIZE 128u25// Matches GGML_PAD(..., 256) in src/llama-context.cpp for KV cache sizing.26#define GGML_WEBGPU_KV_SEQ_PAD 256u27 28#define GGML_WEBGPU_ARGSORT_MERGE_MAX_WG_SIZE 512u29 30// Matrix multiplication parameters31 32// Register tiling parameters33#define WEBGPU_MUL_MAT_TILE_M 434#define WEBGPU_MUL_MAT_TILE_N 435#define WEBGPU_MUL_MAT_WG_SIZE_M 836#define WEBGPU_MUL_MAT_WG_SIZE_N 837#define WEBGPU_MUL_MAT_REG_TILE_K_FLOAT 838#define WEBGPU_MUL_MAT_REG_TILE_K_QUANT 3239 40// Subgroup matrix parameters41// The number of subgroups in the M dimension42#define WEBGPU_MUL_MAT_SUBGROUP_M 243// The number of subgroups in the N dimension44#define WEBGPU_MUL_MAT_SUBGROUP_N 445// The number of subgroup matrices each subgroup accumulates over46#define WEBGPU_MUL_MAT_SUBGROUP_MATRIX_M 447#define WEBGPU_MUL_MAT_SUBGROUP_MATRIX_N 248#define WEBGPU_MUL_MAT_SUBGROUP_TILE_K_FLOAT 3249#define WEBGPU_MUL_MAT_SUBGROUP_TILE_K_QUANT 3250 51// Matrix-vector multiplication parameters52#define WEBGPU_MUL_MAT_VEC_WG_SIZE 25653 54#define WEBGPU_MUL_MAT_VEC_FLOAT_OUTPUTS_PER_WG 455#define WEBGPU_MUL_MAT_VEC_LEGACY_Q_OUTPUTS_PER_WG 456#define WEBGPU_MUL_MAT_VEC_K_Q_OUTPUTS_PER_WG 457 58// default size for reg-tile matrix multiplication59#define WEBGPU_MUL_MAT_WG_SIZE 25660 61// Same hash combine function as in boost62template <typename T> inline void ggml_webgpu_hash_combine(size_t & seed, const T & value) {63 seed ^= std::hash<T>{}(value) + 0x9e3779b9 + (seed << 6) + (seed >> 2);64}65 66// Calculates base address of a tensor ignoring the fake base pointer67inline uintptr_t ggml_webgpu_tensor_addr(const ggml_tensor * tensor) {68 const ggml_tensor * base_tensor = tensor->view_src ? tensor->view_src : tensor;69 return (uintptr_t) base_tensor->data + tensor->view_offs;70}71 72inline bool ggml_webgpu_tensor_equal(const ggml_tensor * a, const ggml_tensor * b) {73 return a->buffer == b->buffer && ggml_webgpu_tensor_addr(a) == ggml_webgpu_tensor_addr(b);74}75 76struct ggml_webgpu_shader_lib_context {77 ggml_tensor * src0;78 ggml_tensor * src1;79 ggml_tensor * src2;80 ggml_tensor * src3;81 ggml_tensor * src4;82 ggml_tensor * src5;83 ggml_tensor * dst;84 85 uint32_t max_wg_size;86 size_t wg_mem_limit_bytes = 0;87 bool supports_subgroups = false;88 bool supports_subgroup_matrix = false;89 uint32_t sg_mat_m = 0;90 uint32_t sg_mat_n = 0;91 uint32_t sg_mat_k = 0;92 uint32_t min_subgroup_size = 0;93 uint32_t max_subgroup_size = 0;94 bool supports_dot_product = false;95 std::string vendor;96};97 98struct webgpu_pipeline {99 wgpu::ComputePipeline pipeline;100 std::string name;101 std::shared_ptr<void> context = nullptr;102};103 104struct ggml_webgpu_generic_shader_decisions {105 uint32_t wg_size = 0;106 bool inplace = false;107};108 109struct ggml_webgpu_binary_shader_decisions {110 uint32_t wg_size = 0;111 bool inplace = false;112 bool overlap = false;113 bool src_overlap = false;114};115 116struct ggml_webgpu_glu_shader_decisions {117 uint32_t wg_size = 0;118 bool src_overlap = false;119};120 121struct ggml_webgpu_processed_shader {122 std::string wgsl;123 std::string variant;124 std::shared_ptr<void> decisions;125};126 127struct ggml_webgpu_ssm_conv_shader_decisions {128 uint32_t block_size;129 uint32_t tokens_per_wg;130};131 132struct ggml_webgpu_ssm_scan_pipeline_key {133 int type;134 int d_state;135 bool xbc_overlap;136 bool a_overlap;137 bool ids_overlap;138 139 bool operator==(const ggml_webgpu_ssm_scan_pipeline_key & other) const {140 return type == other.type && d_state == other.d_state && xbc_overlap == other.xbc_overlap &&141 a_overlap == other.a_overlap && ids_overlap == other.ids_overlap;142 }143};144 145struct ggml_webgpu_ssm_scan_pipeline_key_hash {146 size_t operator()(const ggml_webgpu_ssm_scan_pipeline_key & key) const {147 size_t seed = 0;148 ggml_webgpu_hash_combine(seed, key.type);149 ggml_webgpu_hash_combine(seed, key.d_state);150 ggml_webgpu_hash_combine(seed, key.xbc_overlap);151 ggml_webgpu_hash_combine(seed, key.a_overlap);152 ggml_webgpu_hash_combine(seed, key.ids_overlap);153 return seed;154 }155};156 157struct ggml_webgpu_ssm_scan_shader_decisions {158 uint32_t wg_size;159 uint32_t tokens_per_tile;160 bool xbc_overlap = false;161 bool a_overlap = false;162 bool ids_overlap = false;163};164 165/** Argsort **/166 167struct ggml_webgpu_argsort_shader_lib_context {168 uint32_t max_wg_size;169 size_t wg_mem_limit_bytes;170 int32_t order;171};172 173/** Set Rows **/174 175struct ggml_webgpu_set_rows_pipeline_key {176 int dst_type;177 int vec4;178 int i64_idx;179 int pair_blocks;180 181 bool operator==(const ggml_webgpu_set_rows_pipeline_key & other) const {182 return dst_type == other.dst_type && vec4 == other.vec4 && i64_idx == other.i64_idx &&183 pair_blocks == other.pair_blocks;184 }185};186 187struct ggml_webgpu_set_rows_pipeline_key_hash {188 size_t operator()(const ggml_webgpu_set_rows_pipeline_key & key) const {189 size_t seed = 0;190 ggml_webgpu_hash_combine(seed, key.dst_type);191 ggml_webgpu_hash_combine(seed, key.vec4);192 ggml_webgpu_hash_combine(seed, key.i64_idx);193 ggml_webgpu_hash_combine(seed, key.pair_blocks);194 return seed;195 }196};197 198struct ggml_webgpu_set_rows_shader_decisions {199 bool vec4;200 bool i64_idx;201 bool pair_blocks;202 uint32_t wg_size;203};204 205/** Set **/206 207struct ggml_webgpu_set_pipeline_key {208 ggml_type type;209 bool inplace;210 211 bool operator==(const ggml_webgpu_set_pipeline_key & other) const {212 return type == other.type && inplace == other.inplace;213 }214};215 216struct ggml_webgpu_set_pipeline_key_hash {217 size_t operator()(const ggml_webgpu_set_pipeline_key & key) const {218 size_t seed = 0;219 ggml_webgpu_hash_combine(seed, key.type);220 ggml_webgpu_hash_combine(seed, key.inplace);221 return seed;222 }223};224 225/** Get Rows **/226 227struct ggml_webgpu_get_rows_pipeline_key {228 ggml_type src_type;229 int vectorized;230 231 bool operator==(const ggml_webgpu_get_rows_pipeline_key & other) const {232 return src_type == other.src_type && vectorized == other.vectorized;233 }234};235 236struct ggml_webgpu_get_rows_pipeline_key_hash {237 size_t operator()(const ggml_webgpu_get_rows_pipeline_key & key) const {238 size_t seed = 0;239 ggml_webgpu_hash_combine(seed, key.src_type);240 ggml_webgpu_hash_combine(seed, key.vectorized);241 return seed;242 }243};244 245/** Row Norm **/246 247struct ggml_webgpu_row_norm_pipeline_key {248 ggml_op op;249 ggml_type src_type;250 ggml_type dst_type;251 bool inplace;252 253 bool operator==(const ggml_webgpu_row_norm_pipeline_key & other) const {254 return op == other.op && src_type == other.src_type && dst_type == other.dst_type && inplace == other.inplace;255 }256};257 258struct ggml_webgpu_row_norm_pipeline_key_hash {259 size_t operator()(const ggml_webgpu_row_norm_pipeline_key & key) const {260 size_t seed = 0;261 ggml_webgpu_hash_combine(seed, key.op);262 ggml_webgpu_hash_combine(seed, key.src_type);263 ggml_webgpu_hash_combine(seed, key.dst_type);264 ggml_webgpu_hash_combine(seed, key.inplace);265 return seed;266 }267};268 269/** RMS_NORM + MUL **/270 271struct ggml_webgpu_rms_norm_mul_pipeline_key {272 bool inplace; // rn_src == dst273 bool overlap; // mul_src == dst274 bool src_overlap; // rn_src binding overlaps mul_src binding275 276 bool operator==(const ggml_webgpu_rms_norm_mul_pipeline_key & other) const {277 return inplace == other.inplace && overlap == other.overlap && src_overlap == other.src_overlap;278 }279};280 281struct ggml_webgpu_rms_norm_mul_pipeline_key_hash {282 size_t operator()(const ggml_webgpu_rms_norm_mul_pipeline_key & key) const {283 size_t seed = 0;284 ggml_webgpu_hash_combine(seed, key.inplace);285 ggml_webgpu_hash_combine(seed, key.overlap);286 ggml_webgpu_hash_combine(seed, key.src_overlap);287 return seed;288 }289};290 291struct ggml_webgpu_rms_norm_mul_shader_decisions {292 uint32_t wg_size = 0;293 bool inplace = false;294 bool overlap = false;295 bool src_overlap = false;296};297 298/** Pad **/299struct ggml_webgpu_pad_pipeline_key {300 bool circular;301 302 bool operator==(const ggml_webgpu_pad_pipeline_key & other) const { return circular == other.circular; }303};304 305struct ggml_webgpu_pad_pipeline_key_hash {306 size_t operator()(const ggml_webgpu_pad_pipeline_key & key) const {307 size_t seed = 0;308 ggml_webgpu_hash_combine(seed, key.circular);309 return seed;310 }311};312 313/** Solve Tri **/314struct ggml_webgpu_solve_tri_pipeline_key {315 int type;316 int n;317 int k;318 319 bool operator==(const ggml_webgpu_solve_tri_pipeline_key & other) const {320 return type == other.type && n == other.n && k == other.k;321 }322};323 324struct ggml_webgpu_solve_tri_pipeline_key_hash {325 size_t operator()(const ggml_webgpu_solve_tri_pipeline_key & key) const {326 size_t seed = 0;327 ggml_webgpu_hash_combine(seed, key.type);328 ggml_webgpu_hash_combine(seed, key.n);329 ggml_webgpu_hash_combine(seed, key.k);330 return seed;331 }332};333 334/** SSM Conv **/335struct ggml_webgpu_ssm_conv_pipeline_key {336 int type;337 int vectorized;338 339 bool operator==(const ggml_webgpu_ssm_conv_pipeline_key & other) const {340 return type == other.type && vectorized == other.vectorized;341 }342};343 344/** CONV 2D */345struct ggml_webgpu_conv2d_pipeline_key {346 ggml_type weight_type;347 ggml_type input_type;348 ggml_type output_type;349 350 bool operator==(const ggml_webgpu_conv2d_pipeline_key & other) const {351 return weight_type == other.weight_type && input_type == other.input_type && output_type == other.output_type;352 }353};354 355struct ggml_webgpu_conv2d_pipeline_key_hash {356 size_t operator()(const ggml_webgpu_conv2d_pipeline_key & key) const {357 size_t seed = 0;358 ggml_webgpu_hash_combine(seed, key.weight_type);359 ggml_webgpu_hash_combine(seed, key.input_type);360 ggml_webgpu_hash_combine(seed, key.output_type);361 return seed;362 }363};364 365// Same type fields as conv2d plus the input layout (WHCN vs CWHN).366struct ggml_webgpu_conv2d_dw_pipeline_key {367 ggml_type weight_type;368 ggml_type input_type;369 ggml_type output_type;370 bool whcn;371 372 bool operator==(const ggml_webgpu_conv2d_dw_pipeline_key & other) const {373 return weight_type == other.weight_type && input_type == other.input_type && output_type == other.output_type &&374 whcn == other.whcn;375 }376};377 378struct ggml_webgpu_conv2d_dw_pipeline_key_hash {379 size_t operator()(const ggml_webgpu_conv2d_dw_pipeline_key & key) const {380 size_t seed = 0;381 ggml_webgpu_hash_combine(seed, key.weight_type);382 ggml_webgpu_hash_combine(seed, key.input_type);383 ggml_webgpu_hash_combine(seed, key.output_type);384 ggml_webgpu_hash_combine(seed, key.whcn);385 return seed;386 }387};388 389/** Im2Col **/390struct ggml_webgpu_im2col_pipeline_key {391 ggml_type input_type;392 ggml_type output_type;393 394 bool operator==(const ggml_webgpu_im2col_pipeline_key & other) const {395 return input_type == other.input_type && output_type == other.output_type;396 }397};398 399struct ggml_webgpu_im2col_pipeline_key_hash {400 size_t operator()(const ggml_webgpu_im2col_pipeline_key & key) const {401 size_t seed = 0;402 ggml_webgpu_hash_combine(seed, key.input_type);403 ggml_webgpu_hash_combine(seed, key.output_type);404 return seed;405 }406};407 408/** Gated Delta Net **/409struct ggml_webgpu_gated_delta_net_pipeline_key {410 int type;411 int s_v;412 int kda;413 414 bool operator==(const ggml_webgpu_gated_delta_net_pipeline_key & other) const {415 return type == other.type && s_v == other.s_v && kda == other.kda;416 }417};418 419struct ggml_webgpu_gated_delta_net_pipeline_key_hash {420 size_t operator()(const ggml_webgpu_gated_delta_net_pipeline_key & key) const {421 size_t seed = 0;422 ggml_webgpu_hash_combine(seed, key.type);423 ggml_webgpu_hash_combine(seed, key.s_v);424 ggml_webgpu_hash_combine(seed, key.kda);425 return seed;426 }427};428 429struct ggml_webgpu_ssm_conv_pipeline_key_hash {430 size_t operator()(const ggml_webgpu_ssm_conv_pipeline_key & key) const {431 size_t seed = 0;432 ggml_webgpu_hash_combine(seed, key.type);433 ggml_webgpu_hash_combine(seed, key.vectorized);434 return seed;435 }436};437 438/** Scale **/439 440struct ggml_webgpu_scale_pipeline_key {441 int inplace;442 443 bool operator==(const ggml_webgpu_scale_pipeline_key & other) const { return inplace == other.inplace; }444};445 446struct ggml_webgpu_scale_pipeline_key_hash {447 size_t operator()(const ggml_webgpu_scale_pipeline_key & key) const {448 size_t seed = 0;449 ggml_webgpu_hash_combine(seed, key.inplace);450 return seed;451 }452};453 454/** Upscale **/455 456struct ggml_webgpu_upscale_pipeline_key {457 ggml_type input_type;458 ggml_type output_type;459 uint32_t base_mode;460 bool antialias;461 462 bool operator==(const ggml_webgpu_upscale_pipeline_key & other) const {463 return input_type == other.input_type && output_type == other.output_type && base_mode == other.base_mode &&464 antialias == other.antialias;465 }466};467 468struct ggml_webgpu_upscale_pipeline_key_hash {469 size_t operator()(const ggml_webgpu_upscale_pipeline_key & key) const {470 size_t seed = 0;471 ggml_webgpu_hash_combine(seed, key.input_type);472 ggml_webgpu_hash_combine(seed, key.output_type);473 ggml_webgpu_hash_combine(seed, key.base_mode);474 ggml_webgpu_hash_combine(seed, key.antialias);475 return seed;476 }477};478 479/** Concat **/480 481struct ggml_webgpu_concat_pipeline_key {482 int type;483 bool src_overlap;484 485 bool operator==(const ggml_webgpu_concat_pipeline_key & other) const {486 return type == other.type && src_overlap == other.src_overlap;487 }488};489 490struct ggml_webgpu_concat_pipeline_key_hash {491 size_t operator()(const ggml_webgpu_concat_pipeline_key & key) const {492 size_t seed = 0;493 ggml_webgpu_hash_combine(seed, key.type);494 ggml_webgpu_hash_combine(seed, key.src_overlap);495 return seed;496 }497};498 499/** Repeat **/500 501struct ggml_webgpu_repeat_pipeline_key {502 int type;503 504 bool operator==(const ggml_webgpu_repeat_pipeline_key & other) const { return type == other.type; }505};506 507struct ggml_webgpu_repeat_pipeline_key_hash {508 size_t operator()(const ggml_webgpu_repeat_pipeline_key & key) const {509 size_t seed = 0;510 ggml_webgpu_hash_combine(seed, key.type);511 return seed;512 }513};514 515/** Binary **/516 517struct ggml_webgpu_binary_pipeline_key {518 int type;519 int op;520 bool inplace;521 bool overlap;522 bool src_overlap;523 524 bool operator==(const ggml_webgpu_binary_pipeline_key & other) const {525 return type == other.type && op == other.op && inplace == other.inplace && overlap == other.overlap &&526 src_overlap == other.src_overlap;527 }528};529 530struct ggml_webgpu_binary_pipeline_key_hash {531 size_t operator()(const ggml_webgpu_binary_pipeline_key & key) const {532 size_t seed = 0;533 ggml_webgpu_hash_combine(seed, key.type);534 ggml_webgpu_hash_combine(seed, key.op);535 ggml_webgpu_hash_combine(seed, key.inplace);536 ggml_webgpu_hash_combine(seed, key.overlap);537 ggml_webgpu_hash_combine(seed, key.src_overlap);538 return seed;539 }540};541 542/* Add_Id */543 544struct ggml_webgpu_add_id_pipeline_key {545 bool inplace;546 547 bool operator==(const ggml_webgpu_add_id_pipeline_key & other) const { return inplace == other.inplace; }548};549 550struct ggml_webgpu_add_id_pipeline_key_hash {551 size_t operator()(const ggml_webgpu_add_id_pipeline_key & key) const {552 size_t seed = 0;553 ggml_webgpu_hash_combine(seed, key.inplace);554 return seed;555 }556};557 558/** Unary **/559 560struct ggml_webgpu_unary_pipeline_key {561 int type;562 int op;563 bool is_unary; // many unary operators fall under the GGML_OP_UNARY umbrella564 bool inplace;565 ggml_tri_type ttype; // only used for GGML_OP_TRI566 567 bool operator==(const ggml_webgpu_unary_pipeline_key & other) const {568 return type == other.type && op == other.op && is_unary == other.is_unary && inplace == other.inplace &&569 ttype == other.ttype;570 }571};572 573struct ggml_webgpu_unary_pipeline_key_hash {574 size_t operator()(const ggml_webgpu_unary_pipeline_key & key) const {575 size_t seed = 0;576 ggml_webgpu_hash_combine(seed, key.type);577 ggml_webgpu_hash_combine(seed, key.op);578 ggml_webgpu_hash_combine(seed, key.is_unary);579 ggml_webgpu_hash_combine(seed, key.inplace);580 ggml_webgpu_hash_combine(seed, key.ttype);581 return seed;582 }583};584 585/** FlashAttention */586 587struct ggml_webgpu_flash_attn_common_pipeline_key {588 ggml_type q_type;589 ggml_type k_type;590 ggml_type v_type;591 ggml_type dst_type;592 uint32_t head_dim_qk;593 uint32_t head_dim_v;594 bool k_direct;595 bool v_direct;596 bool kv_overlap;597 bool has_mask;598 bool has_sinks;599 bool uses_logit_softcap;600 601 bool operator==(const ggml_webgpu_flash_attn_common_pipeline_key & other) const {602 return q_type == other.q_type && k_type == other.k_type && v_type == other.v_type &&603 dst_type == other.dst_type && head_dim_qk == other.head_dim_qk && head_dim_v == other.head_dim_v &&604 k_direct == other.k_direct && v_direct == other.v_direct && kv_overlap == other.kv_overlap &&605 has_mask == other.has_mask && has_sinks == other.has_sinks &&606 uses_logit_softcap == other.uses_logit_softcap;607 }608};609 610inline void ggml_webgpu_flash_attn_hash_common_pipeline_key(size_t & seed,611 const ggml_webgpu_flash_attn_common_pipeline_key & key) {612 ggml_webgpu_hash_combine(seed, key.q_type);613 ggml_webgpu_hash_combine(seed, key.k_type);614 ggml_webgpu_hash_combine(seed, key.v_type);615 ggml_webgpu_hash_combine(seed, key.dst_type);616 ggml_webgpu_hash_combine(seed, key.head_dim_qk);617 ggml_webgpu_hash_combine(seed, key.head_dim_v);618 ggml_webgpu_hash_combine(seed, key.k_direct);619 ggml_webgpu_hash_combine(seed, key.v_direct);620 ggml_webgpu_hash_combine(seed, key.kv_overlap);621 ggml_webgpu_hash_combine(seed, key.has_mask);622 ggml_webgpu_hash_combine(seed, key.has_sinks);623 ggml_webgpu_hash_combine(seed, key.uses_logit_softcap);624}625 626struct ggml_webgpu_flash_attn_vec_pipeline_key {627 ggml_webgpu_flash_attn_common_pipeline_key common;628 629 bool operator==(const ggml_webgpu_flash_attn_vec_pipeline_key & other) const { return common == other.common; }630};631 632struct ggml_webgpu_flash_attn_vec_pipeline_key_hash {633 size_t operator()(const ggml_webgpu_flash_attn_vec_pipeline_key & key) const {634 size_t seed = 0;635 ggml_webgpu_flash_attn_hash_common_pipeline_key(seed, key.common);636 return seed;637 }638};639 640struct ggml_webgpu_flash_attn_pipeline_key {641 ggml_webgpu_flash_attn_common_pipeline_key common;642 bool use_sg_matrix;643 644 bool operator==(const ggml_webgpu_flash_attn_pipeline_key & other) const {645 return common == other.common && use_sg_matrix == other.use_sg_matrix;646 }647};648 649struct ggml_webgpu_flash_attn_pipeline_key_hash {650 size_t operator()(const ggml_webgpu_flash_attn_pipeline_key & key) const {651 size_t seed = 0;652 ggml_webgpu_flash_attn_hash_common_pipeline_key(seed, key.common);653 ggml_webgpu_hash_combine(seed, key.use_sg_matrix);654 return seed;655 }656};657 658struct ggml_webgpu_flash_attn_vec_decisions {659 uint32_t kv_tile = 0;660 uint32_t wg_size = 0;661};662 663struct ggml_webgpu_flash_attn_decisions {664 bool use_sg_matrix = false;665 uint32_t q_tile = 0;666 uint32_t kv_tile = 0;667 uint32_t wg_size = 0;668};669 670inline constexpr uint32_t GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH = 4u;671inline constexpr uint32_t GGML_WEBGPU_FLASH_ATTN_TILE_Q_TILE = 4u;672 673inline size_t ggml_webgpu_flash_attn_tensor_offset(const ggml_tensor * tensor) {674 constexpr uintptr_t ptr_base_addr = 0x1000u;675 const ggml_tensor * base = tensor->view_src != nullptr ? tensor->view_src : tensor;676 return reinterpret_cast<uintptr_t>(base->data) - ptr_base_addr + tensor->view_offs;677}678 679inline bool ggml_webgpu_flash_attn_float_vec4_aligned(const ggml_tensor * K, size_t storage_offset_alignment) {680 const uint32_t offset_elems =681 (uint32_t) ((ggml_webgpu_flash_attn_tensor_offset(K) & (storage_offset_alignment - 1)) /682 ggml_type_size(K->type));683 return offset_elems % GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH == 0u;684}685 686inline bool ggml_webgpu_flash_attn_float_vec4_aligned(const ggml_tensor * K,687 const ggml_tensor * V,688 size_t storage_offset_alignment) {689 return ggml_webgpu_flash_attn_float_vec4_aligned(K, storage_offset_alignment) &&690 ggml_webgpu_flash_attn_float_vec4_aligned(V, storage_offset_alignment);691}692 693inline bool ggml_webgpu_flash_attn_k_direct(const ggml_tensor * Q, const ggml_tensor * K, uint32_t kv_direct_align) {694 return (K->type == GGML_TYPE_F16 || K->type == GGML_TYPE_Q8_0 || K->type == GGML_TYPE_Q4_0) &&695 (Q->ne[0] % kv_direct_align == 0) && (K->ne[1] % GGML_WEBGPU_KV_SEQ_PAD == 0);696}697 698inline bool ggml_webgpu_flash_attn_v_direct(const ggml_tensor * Q, const ggml_tensor * V, uint32_t kv_direct_align) {699 return ggml_webgpu_flash_attn_k_direct(Q, V, kv_direct_align);700}701 702inline ggml_webgpu_flash_attn_common_pipeline_key ggml_webgpu_flash_attn_make_common_pipeline_key(703 const ggml_webgpu_shader_lib_context & context,704 uint32_t kv_direct_align,705 bool kv_overlap) {706 ggml_webgpu_flash_attn_common_pipeline_key key = {};707 key.q_type = context.src0->type;708 key.k_type = context.src1->type;709 key.v_type = context.src2->type;710 key.dst_type = context.dst->type;711 key.head_dim_qk = (uint32_t) context.src0->ne[0];712 key.head_dim_v = (uint32_t) context.src2->ne[0];713 key.k_direct = ggml_webgpu_flash_attn_k_direct(context.src0, context.src1, kv_direct_align);714 key.v_direct = ggml_webgpu_flash_attn_v_direct(context.src0, context.src2, kv_direct_align);715 key.kv_overlap = kv_overlap;716 key.has_mask = context.src3 != nullptr;717 key.has_sinks = context.src4 != nullptr;718 key.uses_logit_softcap = ggml_get_op_params_f32(context.dst, 2) != 0.0f;719 return key;720}721 722inline std::vector<std::string> ggml_webgpu_flash_attn_common_defines(723 const ggml_webgpu_flash_attn_common_pipeline_key & key,724 std::string & variant,725 uint32_t q_tile,726 uint32_t kv_tile,727 uint32_t wg_size) {728 std::vector<std::string> defines;729 730 switch (key.k_type) {731 case GGML_TYPE_F32:732 defines.push_back("K_F32");733 break;734 case GGML_TYPE_F16:735 defines.push_back("K_F16");736 break;737 case GGML_TYPE_Q4_0:738 defines.push_back("K_Q4_0");739 break;740 case GGML_TYPE_Q8_0:741 defines.push_back("K_Q8_0");742 break;743 default:744 GGML_ABORT("Unsupported K type for flash attention shader");745 }746 variant += std::string("_k") + ggml_type_name(key.k_type);747 748 switch (key.v_type) {749 case GGML_TYPE_F32:750 defines.push_back("V_F32");751 break;752 case GGML_TYPE_F16:753 defines.push_back("V_F16");754 break;755 case GGML_TYPE_Q4_0:756 defines.push_back("V_Q4_0");757 break;758 case GGML_TYPE_Q8_0:759 defines.push_back("V_Q8_0");760 break;761 default:762 GGML_ABORT("Unsupported V type for flash attention shader");763 }764 variant += std::string("_v") + ggml_type_name(key.v_type);765 766 switch (key.q_type) {767 case GGML_TYPE_F32:768 defines.push_back("Q_F32");769 break;770 case GGML_TYPE_F16:771 defines.push_back("Q_F16");772 break;773 default:774 GGML_ABORT("Unsupported Q type for flash attention shader");775 }776 variant += std::string("_q") + ggml_type_name(key.q_type);777 778 switch (key.dst_type) {779 case GGML_TYPE_F32:780 defines.push_back("DST_F32");781 break;782 case GGML_TYPE_F16:783 defines.push_back("DST_F16");784 break;785 default:786 GGML_ABORT("Unsupported dst type for flash attention shader");787 }788 variant += std::string("_dst") + ggml_type_name(key.dst_type);789 790 if (key.has_mask) {791 defines.push_back("MASK");792 variant += "_mask";793 }794 if (key.has_sinks) {795 defines.push_back("SINKS");796 variant += "_sinks";797 }798 if (key.uses_logit_softcap) {799 defines.push_back("LOGIT_SOFTCAP");800 variant += "_lgsc";801 }802 if (key.k_direct) {803 defines.push_back("K_DIRECT");804 variant += "_k_direct";805 }806 if (key.v_direct) {807 defines.push_back("V_DIRECT");808 variant += "_v_direct";809 }810 if (key.kv_overlap) {811 defines.push_back("KV_OVERLAP");812 variant += "_kv_overlap";813 }814 815 defines.push_back(std::string("HEAD_DIM_QK=") + std::to_string(key.head_dim_qk));816 variant += std::string("_hsqk") + std::to_string(key.head_dim_qk);817 818 defines.push_back(std::string("HEAD_DIM_V=") + std::to_string(key.head_dim_v));819 variant += std::string("_hsv") + std::to_string(key.head_dim_v);820 821 defines.push_back(std::string("Q_TILE=") + std::to_string(q_tile));822 defines.push_back(std::string("KV_TILE=") + std::to_string(kv_tile));823 defines.push_back(std::string("WG_SIZE=") + std::to_string(wg_size));824 825 if (ggml_is_quantized(key.k_type) || ggml_is_quantized(key.v_type)) {826 defines.push_back("U32_DEQUANT_HELPERS");827 if (ggml_is_quantized(key.k_type)) {828 defines.push_back("LOADERS_QUANTIZED_K");829 }830 if (ggml_is_quantized(key.v_type)) {831 defines.push_back("LOADERS_QUANTIZED_V");832 }833 }834 835 return defines;836}837 838struct ggml_webgpu_flash_attn_vec_reduce_pipeline_key {839 uint32_t head_dim_v;840 uint32_t wg_size;841 ggml_type dst_type;842};843 844struct ggml_webgpu_flash_attn_vec_reduce_pipeline_key_hash {845 size_t operator()(const ggml_webgpu_flash_attn_vec_reduce_pipeline_key & key) const {846 size_t seed = 0;847 ggml_webgpu_hash_combine(seed, key.head_dim_v);848 ggml_webgpu_hash_combine(seed, key.wg_size);849 ggml_webgpu_hash_combine(seed, key.dst_type);850 return seed;851 }852};853 854inline bool operator==(const ggml_webgpu_flash_attn_vec_reduce_pipeline_key & lhs,855 const ggml_webgpu_flash_attn_vec_reduce_pipeline_key & rhs) {856 return lhs.head_dim_v == rhs.head_dim_v && lhs.wg_size == rhs.wg_size && lhs.dst_type == rhs.dst_type;857}858 859struct ggml_webgpu_flash_attn_blk_pipeline_key {860 uint32_t kv_tile;861 862 bool operator==(const ggml_webgpu_flash_attn_blk_pipeline_key & other) const { return kv_tile == other.kv_tile; }863};864 865struct ggml_webgpu_flash_attn_blk_pipeline_key_hash {866 size_t operator()(const ggml_webgpu_flash_attn_blk_pipeline_key & key) const {867 size_t seed = 0;868 ggml_webgpu_hash_combine(seed, key.kv_tile);869 return seed;870 }871};872 873// Note: this will slightly overestimate memory usage for vec path874// since row_max and exp_sum shmem are not needed.875inline size_t ggml_webgpu_flash_attn_wg_mem_bytes(uint32_t q_tile,876 uint32_t kv_tile,877 uint32_t head_dim_qk,878 uint32_t head_dim_v,879 bool has_mask,880 bool kv_direct) {881 const uint32_t max_head_dim = std::max(head_dim_qk, head_dim_v);882 size_t f16_elems = 0;883 size_t f32_elems = 0;884 885 f32_elems += q_tile * head_dim_qk; // q_shmem886 if (!kv_direct) {887 f32_elems += kv_tile * max_head_dim; // kv_shmem888 }889 f32_elems += q_tile * head_dim_v; // o_shmem890 if (has_mask) {891 f32_elems += q_tile * kv_tile; // mask_shmem892 }893 f32_elems += q_tile * kv_tile; // inter_shmem894 f32_elems += q_tile; // row_max_shmem895 f32_elems += q_tile; // exp_sum_shmem896 return f16_elems * GGML_WEBGPU_F16_SIZE_BYTES + f32_elems * GGML_WEBGPU_F32_SIZE_BYTES;897}898 899inline uint32_t ggml_webgpu_flash_attn_max_kv_tile(size_t limit_bytes,900 uint32_t q_tile,901 uint32_t kv_granularity,902 uint32_t head_dim_qk,903 uint32_t head_dim_v,904 bool has_mask,905 bool kv_direct) {906 const size_t base_q_bytes =907 ggml_webgpu_flash_attn_wg_mem_bytes(q_tile, 0, head_dim_qk, head_dim_v, has_mask, kv_direct);908 if (limit_bytes <= base_q_bytes) {909 return 0;910 }911 const size_t one_kv_bytes =912 ggml_webgpu_flash_attn_wg_mem_bytes(q_tile, 1, head_dim_qk, head_dim_v, has_mask, kv_direct);913 const size_t bytes_per_kv = one_kv_bytes - base_q_bytes;914 if (bytes_per_kv == 0) {915 return 0;916 }917 const size_t max_kv_tile = (limit_bytes - base_q_bytes) / bytes_per_kv;918 return (uint32_t) ((max_kv_tile / kv_granularity) * kv_granularity);919}920 921inline uint32_t ggml_webgpu_flash_attn_get_vec_kv_tile(size_t wg_mem_limit_bytes,922 uint32_t head_dim_qk,923 uint32_t head_dim_v,924 bool has_mask,925 bool kv_direct) {926 const uint32_t max_kv_tile =927 ggml_webgpu_flash_attn_max_kv_tile(wg_mem_limit_bytes, 1u, 1u, head_dim_qk, head_dim_v, has_mask, kv_direct);928 GGML_ASSERT(max_kv_tile > 0);929 930 uint32_t kv_tile = std::min(GGML_WEBGPU_FLASH_ATTN_VEC_MAX_KV_TILE, max_kv_tile);931 if (kv_direct) {932 kv_tile = std::min(kv_tile, GGML_WEBGPU_KV_SEQ_PAD);933 while (GGML_WEBGPU_KV_SEQ_PAD % kv_tile != 0) {934 kv_tile -= 1u;935 }936 }937 938 return kv_tile;939}940 941inline bool ggml_webgpu_flash_attn_can_use_subgroup_matrix_path(bool supports_subgroup_matrix,942 uint32_t sg_mat_k,943 uint32_t sg_mat_n,944 const ggml_tensor * Q,945 const ggml_tensor * V) {946 return supports_subgroup_matrix && Q->ne[0] % sg_mat_k == 0 && V->ne[0] % sg_mat_n == 0;947}948 949/** Matrix Multiplication **/950 951struct ggml_webgpu_mul_mat_vec_pipeline_key {952 ggml_type src0_type;953 ggml_type src1_type;954 int vectorized;955 uint32_t num_cols;956 bool use_mmvq;957 958 bool operator==(const ggml_webgpu_mul_mat_vec_pipeline_key & other) const {959 return src0_type == other.src0_type && src1_type == other.src1_type && vectorized == other.vectorized &&960 num_cols == other.num_cols && use_mmvq == other.use_mmvq;961 }962};963 964struct ggml_webgpu_mul_mat_vec_pipeline_key_hash {965 size_t operator()(const ggml_webgpu_mul_mat_vec_pipeline_key & key) const {966 size_t seed = 0;967 ggml_webgpu_hash_combine(seed, key.src0_type);968 ggml_webgpu_hash_combine(seed, key.src1_type);969 ggml_webgpu_hash_combine(seed, key.vectorized);970 ggml_webgpu_hash_combine(seed, key.num_cols);971 ggml_webgpu_hash_combine(seed, key.use_mmvq);972 return seed;973 }974};975 976struct ggml_webgpu_mul_mat_vec_shader_decisions {977 uint32_t wg_size;978 uint32_t outputs_per_wg;979 uint32_t vec_size;980};981 982struct ggml_webgpu_quantize_q8_pipeline_key {983 ggml_type src0_type;984 985 bool operator==(const ggml_webgpu_quantize_q8_pipeline_key & other) const { return src0_type == other.src0_type; }986};987 988struct ggml_webgpu_quantize_q8_pipeline_key_hash {989 size_t operator()(const ggml_webgpu_quantize_q8_pipeline_key & key) const {990 size_t seed = 0;991 ggml_webgpu_hash_combine(seed, key.src0_type);992 return seed;993 }994};995 996struct ggml_webgpu_mul_mat_pipeline_key {997 ggml_type src0_type;998 ggml_type src1_type;999 int vectorized;1000 int use_subgroup_matrix;1001 1002 bool operator==(const ggml_webgpu_mul_mat_pipeline_key & other) const {1003 return src0_type == other.src0_type && src1_type == other.src1_type && vectorized == other.vectorized &&1004 use_subgroup_matrix == other.use_subgroup_matrix;1005 }1006};1007 1008struct ggml_webgpu_mul_mat_pipeline_key_hash {1009 size_t operator()(const ggml_webgpu_mul_mat_pipeline_key & key) const {1010 size_t seed = 0;1011 ggml_webgpu_hash_combine(seed, key.src0_type);1012 ggml_webgpu_hash_combine(seed, key.src1_type);1013 ggml_webgpu_hash_combine(seed, key.vectorized);1014 ggml_webgpu_hash_combine(seed, key.use_subgroup_matrix);1015 return seed;1016 }1017};1018 1019struct ggml_webgpu_mul_mat_shader_decisions {1020 uint32_t tile_k;1021 uint32_t wg_size_m;1022 uint32_t wg_size_n;1023 uint32_t wg_size;1024 uint32_t outputs_per_wg;1025 int use_subgroup_matrix;1026 1027 uint32_t tile_m;1028 uint32_t tile_n;1029 1030 // Subgroup matrix parameters1031 uint32_t subgroup_m;1032 uint32_t subgroup_n;1033 uint32_t subgroup_matrix_m;1034 uint32_t subgroup_matrix_n;1035 1036 uint32_t mul_mat_wg_size;1037};1038 1039/** MUL_MAT_ID **/1040 1041struct ggml_webgpu_mul_mat_id_pipeline_key {1042 ggml_type src0_type;1043 ggml_type src1_type;1044 uint32_t n_experts;1045 uint32_t num_cols;1046 int vectorized;1047 1048 bool operator==(const ggml_webgpu_mul_mat_id_pipeline_key & other) const {1049 return src0_type == other.src0_type && src1_type == other.src1_type && n_experts == other.n_experts &&1050 num_cols == other.num_cols && vectorized == other.vectorized;1051 }1052};1053 1054struct ggml_webgpu_mul_mat_id_pipeline_key_hash {1055 size_t operator()(const ggml_webgpu_mul_mat_id_pipeline_key & key) const {1056 size_t seed = 0;1057 ggml_webgpu_hash_combine(seed, key.src0_type);1058 ggml_webgpu_hash_combine(seed, key.src1_type);1059 ggml_webgpu_hash_combine(seed, key.n_experts);1060 ggml_webgpu_hash_combine(seed, key.num_cols);1061 ggml_webgpu_hash_combine(seed, key.vectorized);1062 return seed;1063 }1064};1065 1066/** Cpy **/1067 1068struct ggml_webgpu_cpy_pipeline_key {1069 ggml_type src_type;1070 ggml_type dst_type;1071 1072 bool operator==(const ggml_webgpu_cpy_pipeline_key & other) const {1073 return src_type == other.src_type && dst_type == other.dst_type;1074 }1075};1076 1077struct ggml_webgpu_cpy_pipeline_key_hash {1078 size_t operator()(const ggml_webgpu_cpy_pipeline_key & key) const {1079 size_t seed = 0;1080 ggml_webgpu_hash_combine(seed, key.src_type);1081 ggml_webgpu_hash_combine(seed, key.dst_type);1082 return seed;1083 }1084};1085 1086/** Glu **/1087 1088struct ggml_webgpu_glu_pipeline_key {1089 ggml_glu_op glu_op;1090 ggml_type type;1091 bool split;1092 bool src_overlap;1093 1094 bool operator==(const ggml_webgpu_glu_pipeline_key & other) const {1095 return glu_op == other.glu_op && type == other.type && split == other.split && src_overlap == other.src_overlap;1096 }1097};1098 1099struct ggml_webgpu_glu_pipeline_key_hash {1100 size_t operator()(const ggml_webgpu_glu_pipeline_key & key) const {1101 size_t seed = 0;1102 ggml_webgpu_hash_combine(seed, key.glu_op);1103 ggml_webgpu_hash_combine(seed, key.type);1104 ggml_webgpu_hash_combine(seed, key.split);1105 ggml_webgpu_hash_combine(seed, key.src_overlap);1106 return seed;1107 }1108};1109 1110/** Rope **/1111 1112struct ggml_webgpu_rope_pipeline_key {1113 ggml_type type;1114 bool inplace;1115 bool has_ff;1116 1117 bool operator==(const ggml_webgpu_rope_pipeline_key & other) const {1118 return type == other.type && inplace == other.inplace && has_ff == other.has_ff;1119 }1120};1121 1122struct ggml_webgpu_rope_pipeline_key_hash {1123 size_t operator()(const ggml_webgpu_rope_pipeline_key & key) const {1124 size_t seed = 0;1125 ggml_webgpu_hash_combine(seed, key.type);1126 ggml_webgpu_hash_combine(seed, key.inplace);1127 ggml_webgpu_hash_combine(seed, key.has_ff);1128 return seed;1129 }1130};1131 1132/** SoftMax **/1133 1134struct ggml_webgpu_soft_max_pipeline_key {1135 ggml_type mask_type;1136 bool has_mask;1137 bool has_sink;1138 bool inplace;1139 1140 bool operator==(const ggml_webgpu_soft_max_pipeline_key & other) const {1141 return mask_type == other.mask_type && has_mask == other.has_mask && has_sink == other.has_sink &&1142 inplace == other.inplace;1143 }1144};1145 1146struct ggml_webgpu_soft_max_pipeline_key_hash {1147 size_t operator()(const ggml_webgpu_soft_max_pipeline_key & key) const {1148 size_t seed = 0;1149 ggml_webgpu_hash_combine(seed, key.mask_type);1150 ggml_webgpu_hash_combine(seed, key.has_mask);1151 ggml_webgpu_hash_combine(seed, key.has_sink);1152 ggml_webgpu_hash_combine(seed, key.inplace);1153 return seed;1154 }1155};1156 1157/** MMVQ **/1158 1159inline bool ggml_webgpu_can_use_mmvq(const ggml_tensor * src0,1160 const ggml_tensor * src1,1161 bool supports_dot_product,1162 const std::string & vendor) {1163 if (src1->ne[1] <= 4) {1164 bool supports_dp4a = vendor == "amd" || vendor == "intel" || vendor == "nvidia";1165 if (supports_dp4a && supports_dot_product) {1166 switch (src1->type) {1167 case GGML_TYPE_F32:1168 switch (src0->type) {1169 case GGML_TYPE_Q4_0:1170 case GGML_TYPE_Q4_1:1171 case GGML_TYPE_Q8_0:1172 case GGML_TYPE_Q2_K:1173 case GGML_TYPE_Q4_K:1174 return src0->ne[0] % 4 == 0;1175 default:1176 break;1177 }1178 break;1179 default:1180 break;1181 }1182 }1183 }1184 return false;1185}1186 1187class ggml_webgpu_shader_lib {1188 wgpu::Device device;1189 pre_wgsl::Preprocessor preprocessor;1190 1191 std::unordered_map<int, webgpu_pipeline> sum_rows_pipelines; // key is fixed, no variants yet1192 std::unordered_map<int, webgpu_pipeline> argmax_pipelines; // key is vec41193 std::unordered_map<int, webgpu_pipeline> argsort_pipelines; // key is order1194 std::unordered_map<int, webgpu_pipeline> argsort_merge_pipelines; // key is order1195 std::unordered_map<int, webgpu_pipeline> cumsum_pipelines; // key is fixed, no variants yet1196 std::unordered_map<ggml_webgpu_row_norm_pipeline_key, webgpu_pipeline, ggml_webgpu_row_norm_pipeline_key_hash>1197 row_norm_pipelines; // op/inplace1198 1199 std::unordered_map<ggml_webgpu_get_rows_pipeline_key, webgpu_pipeline, ggml_webgpu_get_rows_pipeline_key_hash>1200 get_rows_pipelines; // src_type, vectorized