15bff84bf5
* FlashAttention (#13) * Add inplace softmax * Move rms_norm to split row approach * Update debug for supports_op * clean up debug statements * neg f16xf32xip builds and runs, havent actually ran a model that uses neg kernel yet though * neg passes backend test * unary operators pass ggml tests * rms_norm double declaration bug atoned * abides by editor-config * removed vestigial files * fixed autoconfig * All operators (inlcluding xielu) working * removed unnecesarry checking if node->src[1] exists for unary operators * responded and dealt with PR comments * implemented REPL_Template support and removed bug in unary operators kernel * formatted embed wgsl and ggml-webgpu.cpp * Faster tensors (#8) Add fast matrix and matrix/vector multiplication. * Use map for shader replacements instead of pair of strings * Wasm (#9) * webgpu : fix build on emscripten * more debugging stuff * test-backend-ops: force single thread on wasm * fix single-thread case for init_tensor_uniform * use jspi * add pthread * test: remember to set n_thread for cpu backend * Add buffer label and enable dawn-specific toggles to turn off some checks * Intermediate state * Fast working f16/f32 vec4 * Working float fast mul mat * Clean up naming of mul_mat to match logical model, start work on q mul_mat * Setup for subgroup matrix mat mul * Basic working subgroup matrix * Working subgroup matrix tiling * Handle weirder sg matrix sizes (but still % sg matrix size) * Working start to gemv * working f16 accumulation with shared memory staging * Print out available subgroup matrix configurations * Vectorize dst stores for sg matrix shader * Gemv working scalar * Minor set_rows optimization (#4) * updated optimization, fixed errors * non vectorized version now dispatches one thread per element * Simplify * Change logic for set_rows pipelines --------- Co-authored-by: Neha Abbas <nehaabbas@macbookpro.lan> Co-authored-by: Neha Abbas <nehaabbas@ReeseLevines-MacBook-Pro.local> Co-authored-by: Reese Levine <reeselevine1@gmail.com> * Comment on dawn toggles * Working subgroup matrix code for (semi)generic sizes * Remove some comments * Cleanup code * Update dawn version and move to portable subgroup size * Try to fix new dawn release * Update subgroup size comment * Only check for subgroup matrix configs if they are supported * Add toggles for subgroup matrix/f16 support on nvidia+vulkan * Make row/col naming consistent * Refactor shared memory loading * Move sg matrix stores to correct file * Working q4_0 * Formatting * Work with emscripten builds * Fix test-backend-ops emscripten for f16/quantized types * Use emscripten memory64 to support get_memory * Add build flags and try ci --------- Co-authored-by: Xuan Son Nguyen <son@huggingface.co> * Remove extra whitespace * Move wasm single-thread logic out of test-backend-ops for cpu backend * Disable multiple threads for emscripten single-thread builds in ggml_graph_plan * Refactored pipelines and workgroup calculations (#10) * refactored pipelines * refactored workgroup calculation * removed commented out block of prior maps * Clean up ceiling division pattern --------- Co-authored-by: Neha Abbas <nehaabbas@eduroam-169-233-141-223.ucsc.edu> Co-authored-by: Reese Levine <reeselevine1@gmail.com> * Start work on flash attention * Shader structure set up (many bugs still) * debugging * Working first test * Working with head grouping, head sizes to 128, logit softcap, mask/sinks enabled, f32 * Generalize softmax to work with multiple subgroups, f16 accumulation, mask shared memory tiling * Start work on integrating pre-wgsl * Separate structs/initial shader compilation library into separate files * Work on compilation choices for flashattention * Work on subgroup matrix/tile size portability * subgroup size agnostic online softmax * Cleanups, quantization types * more cleanup * fix wasm build * Refactor flashattention to increase parallelism, use direct loads for KV in somce cases * Checkpoint * formatting * Update to account for default kv cache padding * formatting shader * Add workflow for ggml-ci webgpu * Try passing absolute path to dawn in ggml-ci * Avoid error on device destruction, add todos for proper cleanup * Fix unused warning * Forgot one parameter unused * Move some flashattn computation to f32 for correctness
170 lines
6.3 KiB
C++
170 lines
6.3 KiB
C++
#ifndef GGML_WEBGPU_SHADER_LIB_HPP
|
|
#define GGML_WEBGPU_SHADER_LIB_HPP
|
|
|
|
#include "ggml.h"
|
|
#include "pre_wgsl.hpp"
|
|
|
|
#include <string>
|
|
#include <vector>
|
|
|
|
#define GGML_WEBGPU_F16_SIZE_BYTES 2
|
|
#define GGML_WEBGPU_F32_SIZE_BYTES 4
|
|
#define GGML_WEBGPU_FLASH_ATTN_PREFERRED_KV_SG_TILES 8u
|
|
#define GGML_WEBGPU_FLASH_ATTN_PREFERRED_WG_SIZE 128u
|
|
// Matches GGML_PAD(..., 256) in src/llama-context.cpp for KV cache sizing.
|
|
#define GGML_WEBGPU_KV_SEQ_PAD 256u
|
|
|
|
struct ggml_webgpu_flash_attn_shader_lib_context {
|
|
ggml_type kv_type;
|
|
uint32_t head_dim_qk;
|
|
uint32_t head_dim_v;
|
|
bool kv_direct;
|
|
bool has_mask;
|
|
bool has_sinks;
|
|
bool uses_logit_softcap;
|
|
uint32_t sg_mat_m;
|
|
uint32_t sg_mat_n;
|
|
uint32_t sg_mat_k;
|
|
size_t wg_mem_limit_bytes;
|
|
uint32_t max_subgroup_size;
|
|
};
|
|
|
|
struct ggml_webgpu_flash_attn_shader_decisions {
|
|
uint32_t q_tile = 0;
|
|
uint32_t kv_tile = 0;
|
|
uint32_t wg_size = 0;
|
|
};
|
|
|
|
struct ggml_webgpu_processed_shader {
|
|
std::string wgsl;
|
|
std::string variant;
|
|
ggml_webgpu_flash_attn_shader_decisions decisions;
|
|
};
|
|
|
|
// This is exposed because it's necessary in supports_op
|
|
inline size_t ggml_webgpu_flash_attn_wg_mem_bytes(uint32_t q_tile,
|
|
uint32_t kv_tile,
|
|
uint32_t head_dim_qk,
|
|
uint32_t head_dim_v,
|
|
bool has_mask,
|
|
bool kv_direct) {
|
|
const uint32_t max_head_dim = std::max(head_dim_qk, head_dim_v);
|
|
size_t f16_elems = 0;
|
|
size_t f32_elems = 0;
|
|
f16_elems += q_tile * head_dim_qk; // q_shmem
|
|
if (!kv_direct) {
|
|
f16_elems += kv_tile * max_head_dim; // kv_shmem
|
|
}
|
|
f16_elems += q_tile * head_dim_v; // o_shmem
|
|
if (has_mask) {
|
|
f16_elems += q_tile * kv_tile; // mask_shmem
|
|
}
|
|
f16_elems += q_tile * kv_tile; // inter_shmem
|
|
f32_elems += q_tile; // row_max_shmem
|
|
f32_elems += q_tile; // exp_sum_shmem
|
|
return f16_elems * GGML_WEBGPU_F16_SIZE_BYTES + f32_elems * GGML_WEBGPU_F32_SIZE_BYTES;
|
|
}
|
|
|
|
static uint32_t ggml_webgpu_flash_attn_max_kv_tile(const ggml_webgpu_flash_attn_shader_lib_context & context) {
|
|
const size_t limit_bytes = context.wg_mem_limit_bytes;
|
|
const size_t q_tile = context.sg_mat_m;
|
|
const size_t base_q_bytes = (context.head_dim_qk + context.head_dim_v) * q_tile * GGML_WEBGPU_F16_SIZE_BYTES +
|
|
2 * q_tile * GGML_WEBGPU_F32_SIZE_BYTES;
|
|
size_t bytes_per_kv = 0;
|
|
if (!context.kv_direct) {
|
|
bytes_per_kv += std::max(context.head_dim_qk, context.head_dim_v);
|
|
}
|
|
if (context.has_mask) {
|
|
bytes_per_kv += q_tile;
|
|
}
|
|
bytes_per_kv += q_tile;
|
|
bytes_per_kv *= GGML_WEBGPU_F16_SIZE_BYTES;
|
|
const uint32_t max_kv_tile = (limit_bytes - base_q_bytes) / bytes_per_kv;
|
|
return (max_kv_tile / context.sg_mat_n) * context.sg_mat_n;
|
|
}
|
|
|
|
inline ggml_webgpu_processed_shader ggml_webgpu_preprocess_flash_attn_shader(
|
|
pre_wgsl::Preprocessor & preprocessor,
|
|
const char * shader_src,
|
|
const ggml_webgpu_flash_attn_shader_lib_context & context) {
|
|
std::vector<std::string> defines;
|
|
std::string variant = "flash_attn";
|
|
|
|
switch (context.kv_type) {
|
|
case GGML_TYPE_F32:
|
|
defines.push_back("KV_F32");
|
|
break;
|
|
case GGML_TYPE_F16:
|
|
defines.push_back("KV_F16");
|
|
break;
|
|
case GGML_TYPE_Q4_0:
|
|
defines.push_back("KV_Q4_0");
|
|
break;
|
|
case GGML_TYPE_Q8_0:
|
|
defines.push_back("KV_Q8_0");
|
|
break;
|
|
default:
|
|
GGML_ABORT("Unsupported KV type for flash attention shader");
|
|
}
|
|
variant += std::string("_") + ggml_type_name(context.kv_type);
|
|
|
|
if (context.has_mask) {
|
|
defines.push_back("MASK");
|
|
variant += "_mask";
|
|
}
|
|
if (context.has_sinks) {
|
|
defines.push_back("SINKS");
|
|
variant += "_sinks";
|
|
}
|
|
if (context.uses_logit_softcap) {
|
|
defines.push_back("LOGIT_SOFTCAP");
|
|
variant += "_lgsc";
|
|
}
|
|
|
|
if (context.kv_direct) {
|
|
defines.push_back("KV_DIRECT");
|
|
variant += "_kvdirect";
|
|
}
|
|
|
|
defines.push_back(std::string("HEAD_DIM_QK=") + std::to_string(context.head_dim_qk));
|
|
variant += std::string("_hsqk") + std::to_string(context.head_dim_qk);
|
|
|
|
defines.push_back(std::string("HEAD_DIM_V=") + std::to_string(context.head_dim_v));
|
|
variant += std::string("_hsv") + std::to_string(context.head_dim_v);
|
|
|
|
// For now these are not part of the variant name
|
|
defines.push_back(std::string("SG_MAT_M=") + std::to_string(context.sg_mat_m));
|
|
defines.push_back(std::string("SG_MAT_N=") + std::to_string(context.sg_mat_n));
|
|
defines.push_back(std::string("SG_MAT_K=") + std::to_string(context.sg_mat_k));
|
|
|
|
// Add chosen Q/KV tile sizes
|
|
uint32_t q_tile = context.sg_mat_m;
|
|
uint32_t kv_tile = std::min(ggml_webgpu_flash_attn_max_kv_tile(context),
|
|
context.sg_mat_n * GGML_WEBGPU_FLASH_ATTN_PREFERRED_KV_SG_TILES);
|
|
if (context.kv_direct) {
|
|
GGML_ASSERT(kv_tile <= GGML_WEBGPU_KV_SEQ_PAD);
|
|
// Avoids having to use bounds-checks and decreasing performance for direct KV loads
|
|
while (GGML_WEBGPU_KV_SEQ_PAD % kv_tile != 0) {
|
|
kv_tile -= context.sg_mat_n;
|
|
}
|
|
}
|
|
|
|
defines.push_back(std::string("Q_TILE=") + std::to_string(q_tile));
|
|
defines.push_back(std::string("KV_TILE=") + std::to_string(kv_tile));
|
|
|
|
// workgroup size
|
|
uint32_t wg_size = std::max(context.max_subgroup_size, GGML_WEBGPU_FLASH_ATTN_PREFERRED_WG_SIZE);
|
|
|
|
defines.push_back(std::string("WG_SIZE=") + std::to_string(wg_size));
|
|
|
|
ggml_webgpu_processed_shader result;
|
|
result.wgsl = preprocessor.preprocess(shader_src, defines);
|
|
result.variant = variant;
|
|
result.decisions.q_tile = q_tile;
|
|
result.decisions.kv_tile = kv_tile;
|
|
result.decisions.wg_size = wg_size;
|
|
return result;
|
|
}
|
|
|
|
#endif // GGML_WEBGPU_SHADER_LIB_HPP
|