Team Ai
Datasetpublic

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.

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes3.1kdownloads
ggml-webgpu-shader-lib.hpp3418 linesDownload Raw Back to ggml-webgpu
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

Showing the first 1,200 of 3418 lines. Download the file for the rest.