diff --git a/backends/webgpu/runtime/WebGPUBackend.cpp b/backends/webgpu/runtime/WebGPUBackend.cpp index ba7f6f13e64..e1497bc4ed8 100644 --- a/backends/webgpu/runtime/WebGPUBackend.cpp +++ b/backends/webgpu/runtime/WebGPUBackend.cpp @@ -79,8 +79,32 @@ Result WebGPUBackend::init( return Error::DelegateInvalidCompatibility; } + // Load-time backend option (BackendOption / LoadBackendOptionsMap), keyed by + // the registered backend name; default false. Mirrors the CoreML/XNNPACK + // runtime-spec pattern -- no compile flag and no .pte re-export needed. + bool enable_f16_kv_cache = false; + { + Result spec = context.get_runtime_spec("enable_f16_kv_cache"); + if (spec.ok()) { + enable_f16_kv_cache = spec.get(); + } + } + bool enable_f16_accumulate_gemm = false; + { + Result spec = + context.get_runtime_spec("enable_f16_accumulate_gemm"); + if (spec.ok()) { + enable_f16_accumulate_gemm = spec.get(); + } + } + try { - graph->build(flatbuffer_data, constant_data, context.get_named_data_map()); + graph->build( + flatbuffer_data, + constant_data, + context.get_named_data_map(), + enable_f16_kv_cache, + enable_f16_accumulate_gemm); } catch (const std::exception& e) { ET_LOG(Error, "WebGPU graph build failed: %s", e.what()); graph->~WebGPUGraph(); diff --git a/backends/webgpu/runtime/WebGPUDevice.cpp b/backends/webgpu/runtime/WebGPUDevice.cpp index 9f48347c16b..d4b148cda5f 100644 --- a/backends/webgpu/runtime/WebGPUDevice.cpp +++ b/backends/webgpu/runtime/WebGPUDevice.cpp @@ -13,9 +13,7 @@ #include #include #include -#ifdef WGPU_BACKEND_ENABLE_PROFILING #include -#endif // WGPU_BACKEND_ENABLE_PROFILING namespace executorch { namespace backends { @@ -143,16 +141,28 @@ WebGPUContext create_webgpu_context() { device_desc.requiredLimits = &supported_limits; } + // Request optional features when the adapter advertises them. A feature the + // adapter lacks is skipped (its fast path stays disabled). A feature the + // adapter advertises becomes required of the device, so if the device were + // to reject it, context creation fails below rather than falling back. The + // vector must outlive wgpuAdapterRequestDevice below (device_desc points + // into it). + std::vector required_features; + if (wgpuAdapterHasFeature(ctx.adapter, WGPUFeatureName_ShaderF16)) { + required_features.push_back(WGPUFeatureName_ShaderF16); + ctx.shader_f16_supported = true; + } #ifdef WGPU_BACKEND_ENABLE_PROFILING // Bench: enable TimestampQuery if available; fail-open (skip timing if not). - std::vector required_features; if (wgpuAdapterHasFeature(ctx.adapter, WGPUFeatureName_TimestampQuery)) { required_features.push_back(WGPUFeatureName_TimestampQuery); - device_desc.requiredFeatureCount = required_features.size(); - device_desc.requiredFeatures = required_features.data(); ctx.timestamp_supported = true; } #endif // WGPU_BACKEND_ENABLE_PROFILING + if (!required_features.empty()) { + device_desc.requiredFeatureCount = required_features.size(); + device_desc.requiredFeatures = required_features.data(); + } device_desc.uncapturedErrorCallbackInfo.callback = on_device_error; diff --git a/backends/webgpu/runtime/WebGPUDevice.h b/backends/webgpu/runtime/WebGPUDevice.h index a332edef443..12f73c969a7 100644 --- a/backends/webgpu/runtime/WebGPUDevice.h +++ b/backends/webgpu/runtime/WebGPUDevice.h @@ -25,6 +25,9 @@ struct WebGPUContext { WGPUAdapter adapter = nullptr; WGPUDevice device = nullptr; WGPUQueue queue = nullptr; + // True if the device was created with the ShaderF16 feature; reserved for a + // future fp16 storage/compute path (fp32 is used when false or unset). + bool shader_f16_supported = false; #ifdef WGPU_BACKEND_ENABLE_PROFILING // True if the device was created with the TimestampQuery feature (bench). bool timestamp_supported = false; diff --git a/backends/webgpu/runtime/WebGPUGraph.cpp b/backends/webgpu/runtime/WebGPUGraph.cpp index eb10fb44e28..d9ab6467882 100644 --- a/backends/webgpu/runtime/WebGPUGraph.cpp +++ b/backends/webgpu/runtime/WebGPUGraph.cpp @@ -90,6 +90,49 @@ WGPUBuffer WebGPUGraph::create_scratch_buffer(size_t nbytes) { return buffer; } +WGPUBuffer WebGPUGraph::acquire_scratch(size_t nbytes) { + nbytes = nbytes > 0 ? nbytes : 4; + // Best-fit reuse: smallest free slot with size in [nbytes, 2*nbytes] -- the + // 2x cap stops a large Cmax-sized buffer from backing a tiny request. Never + // reuse an in_use slot (co-live safety). + ScratchSlot* best = nullptr; + for (auto& s : scratch_pool_) { + // s.size - nbytes (safe: s.size >= nbytes) avoids overflowing 2 * nbytes. + if (!s.in_use && s.size >= nbytes && s.size - nbytes <= nbytes) { + if (best == nullptr || s.size < best->size) { + best = &s; + } + } + } + if (best != nullptr) { + best->in_use = true; + return best->buffer; + } + // None reusable -> create a new slot (freed in the dtor, like + // scratch_buffers_). + WGPUBufferDescriptor buf_desc = {}; + buf_desc.size = nbytes; + buf_desc.usage = WGPUBufferUsage_Storage | WGPUBufferUsage_CopyDst | + WGPUBufferUsage_CopySrc; + buf_desc.mappedAtCreation = false; + WGPUBuffer buffer = wgpuDeviceCreateBuffer(device_, &buf_desc); + scratch_pool_.push_back({buffer, nbytes, true}); + return buffer; +} + +void WebGPUGraph::release_scratch(WGPUBuffer buffer) { + if (!buffer) { + return; + } + for (auto& s : scratch_pool_) { + if (s.buffer == buffer) { + s.in_use = false; + return; + } + } + // Not a pooled buffer -> no-op; the dtor frees it via scratch_buffers_. +} + WGPUBuffer WebGPUGraph::make_uniform_buffer(const void* data, size_t size) { WGPUBufferDescriptor desc = {}; desc.size = size; @@ -267,6 +310,11 @@ WebGPUGraph::~WebGPUGraph() { wgpuBufferRelease(buf); } } + for (auto& s : scratch_pool_) { + if (s.buffer) { + wgpuBufferRelease(s.buffer); + } + } for (auto& buf : owned_uniform_buffers_) { if (buf) { wgpuBufferRelease(buf); @@ -310,7 +358,9 @@ WebGPUGraph::~WebGPUGraph() { void WebGPUGraph::build( const void* flatbuffer_data, const uint8_t* constant_data, - const executorch::runtime::NamedDataMap* named_data_map) { + const executorch::runtime::NamedDataMap* named_data_map, + bool f16_kv_cache, + bool f16_accumulate_gemm) { if (!device_) { auto* ctx = get_default_webgpu_context(); if (ctx) { @@ -331,6 +381,15 @@ void WebGPUGraph::build( constant_data_ = constant_data; named_data_map_ = named_data_map; + // f16 KV cache (runtime opt-in): store K/V caches as f16 iff the opt-in is + // set AND the device negotiated shader-f16 (fail-closed). + const WebGPUContext* kv_ctx = get_default_webgpu_context(); + kv_f16_ = f16_kv_cache && (kv_ctx != nullptr && kv_ctx->shader_f16_supported); + + // f16-accumulate q4gsw steel prefill GEMM (runtime opt-in). QuantizedLinear + // additionally gates the kernel on the negotiated shader-f16 feature. + f16_accumulate_gemm_ = f16_accumulate_gemm; + // Phase 1: Create all values const auto* values = graph->values(); const int num_vals = values ? values->size() : 0; @@ -358,6 +417,13 @@ void WebGPUGraph::build( if (!a) { continue; } + // f16 KV: tag sdpa K/V cache values (args[3],[4]) for half-size alloc. + // Inert unless kv_f16_ (runtime opt-in) is set. + if (kv_f16_ && a->size() > 4 && + oc->name()->str() == "sdpa_with_kv_cache.default") { + kv_cache_ids_.insert(static_cast(a->Get(3))); + kv_cache_ids_.insert(static_cast(a->Get(4))); + } for (unsigned j = 0; j < a->size(); j++) { int id = static_cast(a->Get(j)); if (is_prepack && j == 0) { @@ -378,6 +444,29 @@ void WebGPUGraph::build( } } + // f16 KV defensive guard: fail loud if a non-sdpa op reads an f16 cache. + // Inert unless kv_f16_ (runtime opt-in) is set. + if (kv_f16_ && !kv_cache_ids_.empty() && chain_prescan) { + for (unsigned ci = 0; ci < chain_prescan->size(); ci++) { + const auto* oc = chain_prescan->Get(ci); + const std::string nm = oc->name()->str(); + if (nm == "sdpa_with_kv_cache.default" || nm == kPrepackOpName) { + continue; + } + const auto* a = oc->args(); + if (!a) { + continue; + } + for (unsigned j = 0; j < a->size(); j++) { + if (kv_cache_ids_.count(static_cast(a->Get(j))) != 0) { + throw std::runtime_error( + "WebGPU f16 KV: cache tensor consumed by non-sdpa op '" + nm + + "' would misread the f16 buffer"); + } + } + } + } + for (int i = 0; i < num_vals; i++) { const auto* val = values->Get(i); if (!val || val->value_type() == vkgraph::GraphTypes::NONE) { @@ -407,6 +496,23 @@ void WebGPUGraph::build( tensor.cur_dims = tensor.dims; tensor.cur_nbytes = tensor.nbytes; + // f16 KV cache: dedicated half-size array buffer. WebGPU + // zero-initializes freshly-created buffers, so no explicit clear is + // needed. Inert unless kv_f16_ (runtime opt-in) is set. + if (kv_f16_ && kv_cache_ids_.count(i) != 0) { + tensor.elem_size = 2; + tensor.nbytes = numel * 2; + tensor.cur_nbytes = tensor.nbytes; + tensor_mem_obj_ids_[i] = -1; + WGPUBufferDescriptor buf_desc = {}; + buf_desc.size = std::max(tensor.nbytes, size_t(4)); + buf_desc.usage = WGPUBufferUsage_Storage | WGPUBufferUsage_CopyDst | + WGPUBufferUsage_CopySrc; + buf_desc.mappedAtCreation = false; + tensor.buffer = wgpuDeviceCreateBuffer(device_, &buf_desc); + break; + } + int constant_id = vk_tensor->constant_id(); int mem_obj_id = vk_tensor->mem_obj_id(); diff --git a/backends/webgpu/runtime/WebGPUGraph.h b/backends/webgpu/runtime/WebGPUGraph.h index 0ebcd8071f9..66f0e401de5 100644 --- a/backends/webgpu/runtime/WebGPUGraph.h +++ b/backends/webgpu/runtime/WebGPUGraph.h @@ -104,7 +104,9 @@ class WebGPUGraph { void build( const void* flatbuffer_data, const uint8_t* constant_data, - const executorch::runtime::NamedDataMap* named_data_map = nullptr); + const executorch::runtime::NamedDataMap* named_data_map = nullptr, + bool f16_kv_cache = false, + bool f16_accumulate_gemm = false); // Copy input tensor data from host pointers into GPU buffers. void copy_inputs(const std::vector& inputs); @@ -267,6 +269,35 @@ class WebGPUGraph { // Graph-owned scratch storage buffer for fused-op intermediates (e.g. SDPA). WGPUBuffer create_scratch_buffer(size_t nbytes); + // Reusable scratch pool for SINGLE-OP-LIFETIME fused-op scratch (SDPA + // attn_weights/softmax, FlashDecoding partials). acquire_scratch() reuses a + // free slot (best-fit, size in [n,2n]) or creates one; the caller RELEASES it + // at op-lowering scope exit (use ScopedScratch), so N layers' scratch reuses + // a small constant of buffers instead of N x held to graph teardown. + // Correctness: WebGPU/Dawn auto-inserts RAW hazard barriers between + // dispatches on a shared storage buffer regardless of pass structure -- the + // SAME guarantee mem_obj_id aliasing already relies on -- so reuse is + // bit-identical. Never hand a still-in_use slot to a co-live requester. + WGPUBuffer acquire_scratch(size_t nbytes); + void release_scratch(WGPUBuffer buffer); + // RAII: releases an acquired scratch slot when the op-lowering scope exits + // (leak-safe vs early returns). + struct ScopedScratch { + WebGPUGraph* g = nullptr; + WGPUBuffer buf = nullptr; + ScopedScratch(WebGPUGraph* graph, WGPUBuffer b) : g(graph), buf(b) {} + ~ScopedScratch() { + if (g && buf) { + g->release_scratch(buf); + } + } + ScopedScratch(const ScopedScratch&) = delete; + ScopedScratch& operator=(const ScopedScratch&) = delete; + operator WGPUBuffer() const { + return buf; + } + }; + // Create a mapped-at-creation uniform buffer from `size` bytes and track it // in the memory stats. Shared helper for ops needing a uniform Params buffer. WGPUBuffer make_uniform_buffer(const void* data, size_t size); @@ -314,6 +345,23 @@ class WebGPUGraph { return value_types_[id]; } + public: + // True when the sdpa K/V cache is stored f16-packed (runtime opt-in). + bool kv_f16() const { + return kv_f16_; + } + + // True when the q4gsw steel prefill GEMM uses the lossy f16-accumulate kernel + // (runtime opt-in; perplexity-gated, not bit-exact). + bool f16_accumulate_gemm() const { + return f16_accumulate_gemm_; + } + + private: + bool kv_f16_ = false; + std::unordered_set kv_cache_ids_; + bool f16_accumulate_gemm_ = false; + private: WGPUInstance instance_ = nullptr; WGPUDevice device_ = nullptr; @@ -366,6 +414,16 @@ class WebGPUGraph { // Long-lived scratch storage buffers for fused ops (e.g. SDPA temporaries). std::vector scratch_buffers_; + // Reusable scratch pool: single-op-lifetime buffers recycled across ops + // (acquire_scratch/release_scratch). Each slot is freed in the dtor. See + // acquire_scratch() for the reuse policy. + struct ScratchSlot { + WGPUBuffer buffer = nullptr; + size_t size = 0; + bool in_use = false; + }; + std::vector scratch_pool_; + // Uniform buffers owned for the graph's lifetime; released in the dtor. std::vector owned_uniform_buffers_; diff --git a/backends/webgpu/runtime/ops/quantized_linear/QuantizedLinear.cpp b/backends/webgpu/runtime/ops/quantized_linear/QuantizedLinear.cpp index 130d81e600f..b0728764310 100644 --- a/backends/webgpu/runtime/ops/quantized_linear/QuantizedLinear.cpp +++ b/backends/webgpu/runtime/ops/quantized_linear/QuantizedLinear.cpp @@ -6,11 +6,16 @@ * LICENSE file in the root directory of this source tree. */ +#include #include #include #include #include #include +#include +#include +#include +#include #include #include @@ -52,6 +57,62 @@ constexpr int64_t kQ4gswShmemTileN = 32; constexpr uint32_t kQ4gswShmemMinDim = 4096u; constexpr uint32_t kQ4gswShmemNMinDim = 2048u; +// steel GEMM: 64x64 tile, 256 threads (16x16), fixed wg (no override). +constexpr uint32_t kQ4gswSteelTile = 64u; +constexpr uint32_t kQ4gswSteelBK = 16u; +constexpr uint32_t kQ4gswSteelInvocations = 256u; + +// Max workgroups per 1D dispatch dimension: the device limit, or 65535 when the +// query fails / reports 0. +uint32_t max_workgroups_per_dim(WGPUDevice device) { + WGPULimits limits = {}; + return (wgpuDeviceGetLimits(device, &limits) == WGPUStatus_Success && + limits.maxComputeWorkgroupsPerDimension > 0) + ? limits.maxComputeWorkgroupsPerDimension + : 65535u; +} + +// One workgroup per (tile_m x tile_n) tile, no grid-stride: throw when the tile +// count would exceed the 1D dispatch limit. Shared by the steel + shmem GEMM +// routes; `kind` names the route in the error message. +uint32_t tiled_wg_count( + WGPUDevice device, + uint32_t m, + uint32_t n, + int64_t tile_m, + int64_t tile_n, + const char* op_name, + const char* kind) { + const int64_t total_wgs = + utils::div_up(m, tile_m) * utils::div_up(n, tile_n); + if (total_wgs > static_cast(max_workgroups_per_dim(device))) { + throw std::runtime_error( + std::string("WebGPU ") + op_name + ": " + kind + + " tile count exceeds the 1D dispatch limit"); + } + return static_cast(total_wgs); +} + +// steel needs 256-thread workgroups; fail-closed (query ok AND >=256). +bool steel_supported(WGPUDevice device) { + WGPULimits limits = {}; + return wgpuDeviceGetLimits(device, &limits) == WGPUStatus_Success && + limits.maxComputeInvocationsPerWorkgroup >= kQ4gswSteelInvocations; +} + +// Not grid-strided: 0 (fall back) when K%BK != 0 or over the 1D dispatch limit. +uint32_t +steel_workgroup_count(WGPUDevice device, uint32_t m, uint32_t n, uint32_t K) { + if (K % kQ4gswSteelBK != 0u) { + return 0u; + } + const uint64_t total = + static_cast((m + kQ4gswSteelTile - 1u) / kQ4gswSteelTile) * + static_cast((n + kQ4gswSteelTile - 1u) / kQ4gswSteelTile); + const uint32_t max_count = max_workgroups_per_dim(device); + return (total == 0u || total > max_count) ? 0u : static_cast(total); +} + // Workgroup count for a linear_q4gsw dispatch (bicol GEMV / shmem GEMM / tiled // GEMM), with the range/limit guards shared by the build-time path and the // resize hook. use_gemv/use_shmem_gemm are the build-time routing decision (the @@ -59,6 +120,7 @@ constexpr uint32_t kQ4gswShmemNMinDim = 2048u; uint32_t compute_q4gsw_workgroup_count( WGPUDevice device, bool use_gemv, + bool use_steel, bool use_shmem_gemm, uint32_t m, uint32_t n, @@ -80,22 +142,24 @@ uint32_t compute_q4gsw_workgroup_count( } return wgc; } + if (use_steel) { + // steel: one workgroup per 64x64 tile. Over-limit THROWS here -- unlike the + // build-time steel_workgroup_count, which returns 0 so the caller falls + // back to shmem/tiled. The routed kernel is baked into the pipeline at + // build, so the resize path cannot switch kernels for a larger live M. + return tiled_wg_count( + device, m, n, kQ4gswSteelTile, kQ4gswSteelTile, op_name, "steel GEMM"); + } if (use_shmem_gemm) { - // shmem GEMM: one workgroup per tile, no grid-stride -> throw over limit. - const int64_t total_wgs = utils::div_up(m, kQ4gswShmemTileM) * - utils::div_up(n, kQ4gswShmemTileN); - WGPULimits limits = {}; - const uint32_t max_wgs = - wgpuDeviceGetLimits(device, &limits) == WGPUStatus_Success && - limits.maxComputeWorkgroupsPerDimension > 0 - ? limits.maxComputeWorkgroupsPerDimension - : 65535u; - if (total_wgs > static_cast(max_wgs)) { - throw std::runtime_error( - std::string("WebGPU ") + op_name + - ": shmem GEMM tile count exceeds the 1D dispatch limit"); - } - return static_cast(total_wgs); + // shmem GEMM: one workgroup per tile. + return tiled_wg_count( + device, + m, + n, + kQ4gswShmemTileM, + kQ4gswShmemTileN, + op_name, + "shmem GEMM"); } const int64_t total_tiles = utils::div_up(m, kQ4gswTileM) * utils::div_up(n, kQ4gswTileN); @@ -186,17 +250,57 @@ void q4gsw_linear_impl(WebGPUGraph& graph, const std::vector& args) { "WebGPU linear_q4gsw: scales dims too small for K/N"); } - // M==1 -> bicol GEMV; M>1 -> shmem GEMM (large K/N) else tiled GEMM. + // M==1 -> bicol GEMV; M>1 -> steel GEMM (preferred) else shmem else tiled. const uint32_t wg_size = utils::clamp_workgroup_size(device, kQ4gswLinearWorkgroupSizeX); const bool use_gemv = (M == 1u && K % 8u == 0u && gs % 8u == 0u); - const bool use_shmem_gemm = - !use_gemv && (K >= kQ4gswShmemMinDim || N >= kQ4gswShmemNMinDim); + // steel (256-thread) is the preferred M>1 prefill GEMM; 0 count = ineligible. + const bool use_steel = !use_gemv && steel_supported(device) && + steel_workgroup_count(device, M, N, K) > 0u; + // shmem GEMM is now a FALLBACK, not dead: steel shadows it whenever eligible, + // so shmem only wins when steel is ineligible (K % 16 != 0, or a + // <256-invocation device such as SwiftShader) and the shape still hits the + // large K/N thresholds; otherwise the register-tiled path handles it. + const bool use_shmem_gemm = !use_gemv && !use_steel && + (K >= kQ4gswShmemMinDim || N >= kQ4gswShmemNMinDim); const char* shader_src = use_gemv ? kQ4gswLinearCoop4BicolWGSL + : use_steel ? kQ4gswLinearGemmSteelWGSL : use_shmem_gemm ? kQ4gswLinearGemmShmemWGSL : kQ4gswLinearWGSL; + // f16-multiply steel: only when the device negotiated shader-f16; else the + // f32 steel kernel runs (fail-closed). Same bindings and tile. + if (use_steel) { + const WebGPUContext* ctx = get_default_webgpu_context(); + if (ctx != nullptr && ctx->shader_f16_supported) { + // Packed-word dequant: bit-exact to the steel `half` kernel but loads + // each u32 weight word once + hoists the per-column scale (half re-reads + // them ~8x/~16x). Needs group_size % BK == 0 so the hoisted scale is + // constant across the BK tile; else the per-nibble `half` kernel. + shader_src = (gs % kQ4gswSteelBK == 0u) + ? kQ4gswLinearGemmSteelHalfPwdqWGSL + : kQ4gswLinearGemmSteelHalfWGSL; + } + } + // f16-accumulate: pwdq staging with an f16 register accumulator. + // Lossy (f16 accumulate over K) -> opt-in via the enable_f16_accumulate_gemm + // runtime spec (default off), gated on the negotiated shader-f16 feature and + // group_size % BK == 0 (same hoisted-scale requirement as pwdq). Overrides + // the f32-accumulate steel kernels. + if (use_steel && graph.f16_accumulate_gemm() && (gs % kQ4gswSteelBK == 0u)) { + const WebGPUContext* ctx = get_default_webgpu_context(); + if (ctx != nullptr && ctx->shader_f16_supported) { + shader_src = kQ4gswLinearGemmSteelHalfPwdqF16accWGSL; + } + } const uint32_t workgroup_count = compute_q4gsw_workgroup_count( - device, use_gemv, use_shmem_gemm, M, N, wg_size, "linear_q4gsw"); + device, + use_gemv, + use_steel, + use_shmem_gemm, + M, + N, + wg_size, + "linear_q4gsw"); // Optional bias: real buffer if present, else a dummy for the fixed layout. uint32_t has_bias = 0; @@ -276,8 +380,8 @@ void q4gsw_linear_impl(WebGPUGraph& graph, const std::vector& args) { pipeline_desc.layout = pipeline_layout; pipeline_desc.compute.module = shader; pipeline_desc.compute.entryPoint = {"main", WGPU_STRLEN}; - // Only the tiled GEMM has a wg_size override; GEMV + shmem are fixed 64. - const bool fixed_wg = use_gemv || use_shmem_gemm; + // Only tiled GEMM overrides wg_size; GEMV/shmem (64) + steel (256) are fixed. + const bool fixed_wg = use_gemv || use_steel || use_shmem_gemm; pipeline_desc.compute.constantCount = fixed_wg ? 0u : 1u; pipeline_desc.compute.constants = fixed_wg ? nullptr : &wg_size_constant; WGPUComputePipeline pipeline = @@ -328,6 +432,7 @@ void q4gsw_linear_impl(WebGPUGraph& graph, const std::vector& args) { has_bias, wg_size, use_gemv, + use_steel, use_shmem_gemm, dispatch_idx, uniform_buffer](WebGPUGraph& g) { @@ -356,6 +461,7 @@ void q4gsw_linear_impl(WebGPUGraph& graph, const std::vector& args) { const uint32_t wgc = compute_q4gsw_workgroup_count( g.device(), use_gemv, + use_steel, use_shmem_gemm, m, N, diff --git a/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel.wgsl b/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel.wgsl new file mode 100644 index 00000000000..4d7ab0b1d1e --- /dev/null +++ b/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel.wgsl @@ -0,0 +1,146 @@ +$if DTYPE == "half": + enable f16; +@group(0) @binding(0) var t_out: array; +@group(0) @binding(1) var t_input: array; +@group(0) @binding(2) var t_weight: array; +@group(0) @binding(3) var t_scales: array; +@group(0) @binding(4) var t_bias: array; + +struct Params { + M: u32, + N: u32, + K: u32, + K_packed: u32, + group_size: u32, + padded_N: u32, + has_bias: u32, + _pad: u32, +} +@group(0) @binding(5) var params: Params; + +// "steel" prefill GEMM (M>1): 64x64 tile, 256 threads; K%16==0 host-guarded. +// The "steel" name + register-tiled dequant-to-shared GEMM structure are +// inspired by MLX's steel GEMM kernels (github.com/ml-explore/mlx, +// mlx/backend/metal/kernels/steel). One template, four variants: +// DTYPE=float f32 storage/multiply, per-nibble weight staging. +// DTYPE=half f16 storage/multiply, per-nibble weight staging. +// PWDQ (half only) packed-word dequant: load each u32 weight word ONCE, +// unpack all 16 nibbles of a column + hoist the per-column scale to one read +// (the per-nibble path re-reads each word ~8x). Requires K%BK==0 (steel +// route guarantees it) and group_size%BK==0 (hoisted scale across the tile). +// ACC=half (PWDQ only) f16 accumulate with fma(), cast to f32 in the epilogue +// -- LOSSY, perplexity-gated, opt-in via a runtime spec. ACC=float is f32 +// accumulate -- BIT-EXACT to the per-nibble half kernel. +const BM: u32 = 64u; const BN: u32 = 64u; const BK: u32 = 16u; +var As: array<${buffer_scalar_type(DTYPE)}, 1024>; // BM*BK +var Bs: array<${buffer_scalar_type(DTYPE)}, 1024>; // BK*BN +@compute @workgroup_size(16, 16) +fn main(@builtin(workgroup_id) wid: vec3, + @builtin(local_invocation_id) lid: vec3) { + let nbN = (params.N + BN - 1u) / BN; + let bx = wid.x % nbN; // decode 2D tile id from 1D dispatch + let by = wid.x / nbN; + let row0 = by * BM; + let col0 = bx * BN; + let tid = lid.y * 16u + lid.x; + var acc: array, 4>; + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = ${"0.0h" if ACC == "half" else "0.0"}; } + } + // A staging coords: 256 threads load 64x16 = 1024 f32 -> 4 rows each (4 contiguous K). + let ar = tid / 4u; // 0..63 (row in tile) + let ac = (tid % 4u) * 4u; // 0,4,8,12 (K offset, 4 contiguous) + $if not PWDQ: + // B staging coords: 256 threads load 16x64 = 1024 dequant weights -> 4 cols each. + let br = tid / 16u; // 0..15 (K within BK) + let bc = (tid % 16u) * 4u; // 0,4,..60 (N offset, 4 contiguous) + + var k0: u32 = 0u; + loop { + if (k0 >= params.K) { break; } + // stage activations (edge-masked on M; K is a multiple of BK for our shapes) + let arow = row0 + ar; + if (arow < params.M) { + let base = arow * params.K + k0 + ac; + As[ar * BK + ac + 0u] = ${buffer_scalar_type(DTYPE)}(t_input[base]); + As[ar * BK + ac + 1u] = ${buffer_scalar_type(DTYPE)}(t_input[base + 1u]); + As[ar * BK + ac + 2u] = ${buffer_scalar_type(DTYPE)}(t_input[base + 2u]); + As[ar * BK + ac + 3u] = ${buffer_scalar_type(DTYPE)}(t_input[base + 3u]); + } else { + As[ar * BK + ac + 0u] = ${"0.0h" if PWDQ else "0.0"}; As[ar * BK + ac + 1u] = ${"0.0h" if PWDQ else "0.0"}; + As[ar * BK + ac + 2u] = ${"0.0h" if PWDQ else "0.0"}; As[ar * BK + ac + 3u] = ${"0.0h" if PWDQ else "0.0"}; + } + $if PWDQ: + // Packed-word dequant: threads [0,BN) each stage one full BK-column of Bs. + if (tid < BN) { + let c = tid; // Bs column within this tile + let n = col0 + c; // global output column + if (n < params.N) { + // Scale is constant across the BK tile (group_size % BK == 0 for all real + // group sizes; K%BK==0 on the steel route), so hoist it to one read. + let scale_row = (k0 / params.group_size) * params.padded_N; + let scale = f16(t_scales[scale_row + n]); + // Column n's 16-nibble K-slice for this tile = two consecutive words. + // K_packed multiple of 8 => base_word stays inside column n's own region. + let base_word = n * (params.K_packed >> 2u) + (k0 >> 3u); + let w0 = t_weight[base_word]; + let w1 = t_weight[base_word + 1u]; + for (var br: u32 = 0u; br < BK; br = br + 1u) { + let word = select(w1, w0, br < 8u); // word0 holds K-slice [0,8) + let nib = (word >> ((br & 7u) * 4u)) & 0x0Fu; + Bs[br * BN + c] = f16(i32(nib) - 8) * scale; + } + } else { + for (var br: u32 = 0u; br < BK; br = br + 1u) { Bs[br * BN + c] = 0.0h; } + } + } + $else: + // stage DEQUANTIZED weights into Bs[k][n]: 4 contiguous N per thread. + let kk = k0 + br; // K index for this shmem row + let scale_row = (kk / params.group_size) * params.padded_N; + for (var j: u32 = 0u; j < 4u; j = j + 1u) { + let n = col0 + bc + j; + var dqv: ${buffer_scalar_type(DTYPE)} = 0.0; + if (n < params.N) { + let byte_idx = n * params.K_packed + (kk >> 1u); + let word = t_weight[byte_idx >> 2u]; + let b = (word >> ((byte_idx & 3u) * 8u)) & 0xFFu; + var nib: u32; + if ((kk & 1u) == 0u) { nib = b & 0x0Fu; } else { nib = (b >> 4u) & 0x0Fu; } + $if DTYPE == "half": + dqv = f16(i32(nib) - 8) * f16(t_scales[scale_row + n]); + $else: + dqv = f32(i32(nib) - 8) * t_scales[scale_row + n]; + } + Bs[br * BN + bc + j] = dqv; + } + workgroupBarrier(); + for (var k: u32 = 0u; k < BK; k = k + 1u) { + var a: array<${buffer_scalar_type(DTYPE)}, 4>; + var bvec: array<${buffer_scalar_type(DTYPE)}, 4>; + for (var m: u32 = 0u; m < 4u; m = m + 1u) { a[m] = As[(lid.y * 4u + m) * BK + k]; } + for (var n: u32 = 0u; n < 4u; n = n + 1u) { bvec[n] = Bs[k * BN + lid.x * 4u + n]; } + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + $if ACC == "half": + for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = fma(a[m], bvec[n], acc[m][n]); } + $elif DTYPE == "half": + for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = acc[m][n] + f32(a[m] * bvec[n]); } + $else: + for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = acc[m][n] + a[m] * bvec[n]; } + } + } + workgroupBarrier(); + k0 = k0 + BK; + } + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { + let r = row0 + lid.y * 4u + m; + let c = col0 + lid.x * 4u + n; + if (r < params.M && c < params.N) { + var v = ${"f32(acc[m][n])" if ACC == "half" else "acc[m][n]"}; + if (params.has_bias != 0u) { v = v + t_bias[c]; } + t_out[r * params.N + c] = v; + } + } + } +} diff --git a/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel.yaml b/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel.yaml new file mode 100644 index 00000000000..5a2cae5e499 --- /dev/null +++ b/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel.yaml @@ -0,0 +1,22 @@ +q4gsw_linear_gemm_steel: + parameter_names_with_default_values: + DTYPE: float + PWDQ: false + ACC: float + shader_variants: + - NAME: q4gsw_linear_gemm_steel + DTYPE: float + PWDQ: false + ACC: float + - NAME: q4gsw_linear_gemm_steel_half + DTYPE: half + PWDQ: false + ACC: float + - NAME: q4gsw_linear_gemm_steel_half_pwdq + DTYPE: half + PWDQ: true + ACC: float + - NAME: q4gsw_linear_gemm_steel_half_pwdq_f16acc + DTYPE: half + PWDQ: true + ACC: half diff --git a/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_half_pwdq_f16acc_wgsl.h b/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_half_pwdq_f16acc_wgsl.h new file mode 100644 index 00000000000..efefd7edce1 --- /dev/null +++ b/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_half_pwdq_f16acc_wgsl.h @@ -0,0 +1,141 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#pragma once + +#include + +namespace executorch::backends::webgpu { + +// @generated from q4gsw_linear_gemm_steel.wgsl - DO NOT EDIT. +// wgsl-sha256: 36b3d3f9dd08a529909c13ec7d66cd0cf392c347ca047a4d38453b3c295f72ce +inline constexpr const char* kQ4gswLinearGemmSteelHalfPwdqF16accWGSL = R"( +enable f16; +@group(0) @binding(0) var t_out: array; +@group(0) @binding(1) var t_input: array; +@group(0) @binding(2) var t_weight: array; +@group(0) @binding(3) var t_scales: array; +@group(0) @binding(4) var t_bias: array; + +struct Params { + M: u32, + N: u32, + K: u32, + K_packed: u32, + group_size: u32, + padded_N: u32, + has_bias: u32, + _pad: u32, +} +@group(0) @binding(5) var params: Params; + +// "steel" prefill GEMM (M>1): 64x64 tile, 256 threads; K%16==0 host-guarded. +// The "steel" name + register-tiled dequant-to-shared GEMM structure are +// inspired by MLX's steel GEMM kernels (github.com/ml-explore/mlx, +// mlx/backend/metal/kernels/steel). One template, four variants: +// DTYPE=float f32 storage/multiply, per-nibble weight staging. +// DTYPE=half f16 storage/multiply, per-nibble weight staging. +// PWDQ (half only) packed-word dequant: load each u32 weight word ONCE, +// unpack all 16 nibbles of a column + hoist the per-column scale to one read +// (the per-nibble path re-reads each word ~8x). Requires K%BK==0 (steel +// route guarantees it) and group_size%BK==0 (hoisted scale across the tile). +// ACC=half (PWDQ only) f16 accumulate with fma(), cast to f32 in the epilogue +// -- LOSSY, perplexity-gated, opt-in via a runtime spec. ACC=float is f32 +// accumulate -- BIT-EXACT to the per-nibble half kernel. +const BM: u32 = 64u; const BN: u32 = 64u; const BK: u32 = 16u; +var As: array; // BM*BK +var Bs: array; // BK*BN +@compute @workgroup_size(16, 16) +fn main(@builtin(workgroup_id) wid: vec3, + @builtin(local_invocation_id) lid: vec3) { + let nbN = (params.N + BN - 1u) / BN; + let bx = wid.x % nbN; // decode 2D tile id from 1D dispatch + let by = wid.x / nbN; + let row0 = by * BM; + let col0 = bx * BN; + let tid = lid.y * 16u + lid.x; + var acc: array, 4>; + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = 0.0h; } + } + // A staging coords: 256 threads load 64x16 = 1024 f32 -> 4 rows each (4 contiguous K). + let ar = tid / 4u; // 0..63 (row in tile) + let ac = (tid % 4u) * 4u; // 0,4,8,12 (K offset, 4 contiguous) + + var k0: u32 = 0u; + loop { + if (k0 >= params.K) { break; } + // stage activations (edge-masked on M; K is a multiple of BK for our shapes) + let arow = row0 + ar; + if (arow < params.M) { + let base = arow * params.K + k0 + ac; + As[ar * BK + ac + 0u] = f16(t_input[base]); + As[ar * BK + ac + 1u] = f16(t_input[base + 1u]); + As[ar * BK + ac + 2u] = f16(t_input[base + 2u]); + As[ar * BK + ac + 3u] = f16(t_input[base + 3u]); + } else { + As[ar * BK + ac + 0u] = 0.0h; As[ar * BK + ac + 1u] = 0.0h; + As[ar * BK + ac + 2u] = 0.0h; As[ar * BK + ac + 3u] = 0.0h; + } + // Packed-word dequant: threads [0,BN) each stage one full BK-column of Bs. + if (tid < BN) { + let c = tid; // Bs column within this tile + let n = col0 + c; // global output column + if (n < params.N) { + // Scale is constant across the BK tile (group_size % BK == 0 for all real + // group sizes; K%BK==0 on the steel route), so hoist it to one read. + let scale_row = (k0 / params.group_size) * params.padded_N; + let scale = f16(t_scales[scale_row + n]); + // Column n's 16-nibble K-slice for this tile = two consecutive words. + // K_packed multiple of 8 => base_word stays inside column n's own region. + let base_word = n * (params.K_packed >> 2u) + (k0 >> 3u); + let w0 = t_weight[base_word]; + let w1 = t_weight[base_word + 1u]; + for (var br: u32 = 0u; br < BK; br = br + 1u) { + let word = select(w1, w0, br < 8u); // word0 holds K-slice [0,8) + let nib = (word >> ((br & 7u) * 4u)) & 0x0Fu; + Bs[br * BN + c] = f16(i32(nib) - 8) * scale; + } + } else { + for (var br: u32 = 0u; br < BK; br = br + 1u) { Bs[br * BN + c] = 0.0h; } + } + } + workgroupBarrier(); + for (var k: u32 = 0u; k < BK; k = k + 1u) { + var a: array; + var bvec: array; + for (var m: u32 = 0u; m < 4u; m = m + 1u) { a[m] = As[(lid.y * 4u + m) * BK + k]; } + for (var n: u32 = 0u; n < 4u; n = n + 1u) { bvec[n] = Bs[k * BN + lid.x * 4u + n]; } + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = fma(a[m], bvec[n], acc[m][n]); } + } + } + workgroupBarrier(); + k0 = k0 + BK; + } + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { + let r = row0 + lid.y * 4u + m; + let c = col0 + lid.x * 4u + n; + if (r < params.M && c < params.N) { + var v = f32(acc[m][n]); + if (params.has_bias != 0u) { v = v + t_bias[c]; } + t_out[r * params.N + c] = v; + } + } + } +} +)"; + +inline constexpr uint32_t kQ4gswLinearGemmSteelHalfPwdqF16accWorkgroupSizeX = + 16; +inline constexpr uint32_t kQ4gswLinearGemmSteelHalfPwdqF16accWorkgroupSizeY = + 16; +inline constexpr uint32_t kQ4gswLinearGemmSteelHalfPwdqF16accWorkgroupSizeZ = 1; + +} // namespace executorch::backends::webgpu diff --git a/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_half_pwdq_wgsl.h b/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_half_pwdq_wgsl.h new file mode 100644 index 00000000000..46057de2340 --- /dev/null +++ b/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_half_pwdq_wgsl.h @@ -0,0 +1,139 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#pragma once + +#include + +namespace executorch::backends::webgpu { + +// @generated from q4gsw_linear_gemm_steel.wgsl - DO NOT EDIT. +// wgsl-sha256: 1f916bcd30dbbbcc7eca37e795ecc26e3c72e645ccd2c361fa0ac4e66f1a174a +inline constexpr const char* kQ4gswLinearGemmSteelHalfPwdqWGSL = R"( +enable f16; +@group(0) @binding(0) var t_out: array; +@group(0) @binding(1) var t_input: array; +@group(0) @binding(2) var t_weight: array; +@group(0) @binding(3) var t_scales: array; +@group(0) @binding(4) var t_bias: array; + +struct Params { + M: u32, + N: u32, + K: u32, + K_packed: u32, + group_size: u32, + padded_N: u32, + has_bias: u32, + _pad: u32, +} +@group(0) @binding(5) var params: Params; + +// "steel" prefill GEMM (M>1): 64x64 tile, 256 threads; K%16==0 host-guarded. +// The "steel" name + register-tiled dequant-to-shared GEMM structure are +// inspired by MLX's steel GEMM kernels (github.com/ml-explore/mlx, +// mlx/backend/metal/kernels/steel). One template, four variants: +// DTYPE=float f32 storage/multiply, per-nibble weight staging. +// DTYPE=half f16 storage/multiply, per-nibble weight staging. +// PWDQ (half only) packed-word dequant: load each u32 weight word ONCE, +// unpack all 16 nibbles of a column + hoist the per-column scale to one read +// (the per-nibble path re-reads each word ~8x). Requires K%BK==0 (steel +// route guarantees it) and group_size%BK==0 (hoisted scale across the tile). +// ACC=half (PWDQ only) f16 accumulate with fma(), cast to f32 in the epilogue +// -- LOSSY, perplexity-gated, opt-in via a runtime spec. ACC=float is f32 +// accumulate -- BIT-EXACT to the per-nibble half kernel. +const BM: u32 = 64u; const BN: u32 = 64u; const BK: u32 = 16u; +var As: array; // BM*BK +var Bs: array; // BK*BN +@compute @workgroup_size(16, 16) +fn main(@builtin(workgroup_id) wid: vec3, + @builtin(local_invocation_id) lid: vec3) { + let nbN = (params.N + BN - 1u) / BN; + let bx = wid.x % nbN; // decode 2D tile id from 1D dispatch + let by = wid.x / nbN; + let row0 = by * BM; + let col0 = bx * BN; + let tid = lid.y * 16u + lid.x; + var acc: array, 4>; + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = 0.0; } + } + // A staging coords: 256 threads load 64x16 = 1024 f32 -> 4 rows each (4 contiguous K). + let ar = tid / 4u; // 0..63 (row in tile) + let ac = (tid % 4u) * 4u; // 0,4,8,12 (K offset, 4 contiguous) + + var k0: u32 = 0u; + loop { + if (k0 >= params.K) { break; } + // stage activations (edge-masked on M; K is a multiple of BK for our shapes) + let arow = row0 + ar; + if (arow < params.M) { + let base = arow * params.K + k0 + ac; + As[ar * BK + ac + 0u] = f16(t_input[base]); + As[ar * BK + ac + 1u] = f16(t_input[base + 1u]); + As[ar * BK + ac + 2u] = f16(t_input[base + 2u]); + As[ar * BK + ac + 3u] = f16(t_input[base + 3u]); + } else { + As[ar * BK + ac + 0u] = 0.0h; As[ar * BK + ac + 1u] = 0.0h; + As[ar * BK + ac + 2u] = 0.0h; As[ar * BK + ac + 3u] = 0.0h; + } + // Packed-word dequant: threads [0,BN) each stage one full BK-column of Bs. + if (tid < BN) { + let c = tid; // Bs column within this tile + let n = col0 + c; // global output column + if (n < params.N) { + // Scale is constant across the BK tile (group_size % BK == 0 for all real + // group sizes; K%BK==0 on the steel route), so hoist it to one read. + let scale_row = (k0 / params.group_size) * params.padded_N; + let scale = f16(t_scales[scale_row + n]); + // Column n's 16-nibble K-slice for this tile = two consecutive words. + // K_packed multiple of 8 => base_word stays inside column n's own region. + let base_word = n * (params.K_packed >> 2u) + (k0 >> 3u); + let w0 = t_weight[base_word]; + let w1 = t_weight[base_word + 1u]; + for (var br: u32 = 0u; br < BK; br = br + 1u) { + let word = select(w1, w0, br < 8u); // word0 holds K-slice [0,8) + let nib = (word >> ((br & 7u) * 4u)) & 0x0Fu; + Bs[br * BN + c] = f16(i32(nib) - 8) * scale; + } + } else { + for (var br: u32 = 0u; br < BK; br = br + 1u) { Bs[br * BN + c] = 0.0h; } + } + } + workgroupBarrier(); + for (var k: u32 = 0u; k < BK; k = k + 1u) { + var a: array; + var bvec: array; + for (var m: u32 = 0u; m < 4u; m = m + 1u) { a[m] = As[(lid.y * 4u + m) * BK + k]; } + for (var n: u32 = 0u; n < 4u; n = n + 1u) { bvec[n] = Bs[k * BN + lid.x * 4u + n]; } + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = acc[m][n] + f32(a[m] * bvec[n]); } + } + } + workgroupBarrier(); + k0 = k0 + BK; + } + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { + let r = row0 + lid.y * 4u + m; + let c = col0 + lid.x * 4u + n; + if (r < params.M && c < params.N) { + var v = acc[m][n]; + if (params.has_bias != 0u) { v = v + t_bias[c]; } + t_out[r * params.N + c] = v; + } + } + } +} +)"; + +inline constexpr uint32_t kQ4gswLinearGemmSteelHalfPwdqWorkgroupSizeX = 16; +inline constexpr uint32_t kQ4gswLinearGemmSteelHalfPwdqWorkgroupSizeY = 16; +inline constexpr uint32_t kQ4gswLinearGemmSteelHalfPwdqWorkgroupSizeZ = 1; + +} // namespace executorch::backends::webgpu diff --git a/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_half_wgsl.h b/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_half_wgsl.h new file mode 100644 index 00000000000..ae03f63fb5c --- /dev/null +++ b/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_half_wgsl.h @@ -0,0 +1,135 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#pragma once + +#include + +namespace executorch::backends::webgpu { + +// @generated from q4gsw_linear_gemm_steel.wgsl - DO NOT EDIT. +// wgsl-sha256: 00cdd2f2fb98a5c7343d16fdf7e59f1b840e180cec3f82bf9b569513c0a45396 +inline constexpr const char* kQ4gswLinearGemmSteelHalfWGSL = R"( +enable f16; +@group(0) @binding(0) var t_out: array; +@group(0) @binding(1) var t_input: array; +@group(0) @binding(2) var t_weight: array; +@group(0) @binding(3) var t_scales: array; +@group(0) @binding(4) var t_bias: array; + +struct Params { + M: u32, + N: u32, + K: u32, + K_packed: u32, + group_size: u32, + padded_N: u32, + has_bias: u32, + _pad: u32, +} +@group(0) @binding(5) var params: Params; + +// "steel" prefill GEMM (M>1): 64x64 tile, 256 threads; K%16==0 host-guarded. +// The "steel" name + register-tiled dequant-to-shared GEMM structure are +// inspired by MLX's steel GEMM kernels (github.com/ml-explore/mlx, +// mlx/backend/metal/kernels/steel). One template, four variants: +// DTYPE=float f32 storage/multiply, per-nibble weight staging. +// DTYPE=half f16 storage/multiply, per-nibble weight staging. +// PWDQ (half only) packed-word dequant: load each u32 weight word ONCE, +// unpack all 16 nibbles of a column + hoist the per-column scale to one read +// (the per-nibble path re-reads each word ~8x). Requires K%BK==0 (steel +// route guarantees it) and group_size%BK==0 (hoisted scale across the tile). +// ACC=half (PWDQ only) f16 accumulate with fma(), cast to f32 in the epilogue +// -- LOSSY, perplexity-gated, opt-in via a runtime spec. ACC=float is f32 +// accumulate -- BIT-EXACT to the per-nibble half kernel. +const BM: u32 = 64u; const BN: u32 = 64u; const BK: u32 = 16u; +var As: array; // BM*BK +var Bs: array; // BK*BN +@compute @workgroup_size(16, 16) +fn main(@builtin(workgroup_id) wid: vec3, + @builtin(local_invocation_id) lid: vec3) { + let nbN = (params.N + BN - 1u) / BN; + let bx = wid.x % nbN; // decode 2D tile id from 1D dispatch + let by = wid.x / nbN; + let row0 = by * BM; + let col0 = bx * BN; + let tid = lid.y * 16u + lid.x; + var acc: array, 4>; + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = 0.0; } + } + // A staging coords: 256 threads load 64x16 = 1024 f32 -> 4 rows each (4 contiguous K). + let ar = tid / 4u; // 0..63 (row in tile) + let ac = (tid % 4u) * 4u; // 0,4,8,12 (K offset, 4 contiguous) + // B staging coords: 256 threads load 16x64 = 1024 dequant weights -> 4 cols each. + let br = tid / 16u; // 0..15 (K within BK) + let bc = (tid % 16u) * 4u; // 0,4,..60 (N offset, 4 contiguous) + + var k0: u32 = 0u; + loop { + if (k0 >= params.K) { break; } + // stage activations (edge-masked on M; K is a multiple of BK for our shapes) + let arow = row0 + ar; + if (arow < params.M) { + let base = arow * params.K + k0 + ac; + As[ar * BK + ac + 0u] = f16(t_input[base]); + As[ar * BK + ac + 1u] = f16(t_input[base + 1u]); + As[ar * BK + ac + 2u] = f16(t_input[base + 2u]); + As[ar * BK + ac + 3u] = f16(t_input[base + 3u]); + } else { + As[ar * BK + ac + 0u] = 0.0; As[ar * BK + ac + 1u] = 0.0; + As[ar * BK + ac + 2u] = 0.0; As[ar * BK + ac + 3u] = 0.0; + } + // stage DEQUANTIZED weights into Bs[k][n]: 4 contiguous N per thread. + let kk = k0 + br; // K index for this shmem row + let scale_row = (kk / params.group_size) * params.padded_N; + for (var j: u32 = 0u; j < 4u; j = j + 1u) { + let n = col0 + bc + j; + var dqv: f16 = 0.0; + if (n < params.N) { + let byte_idx = n * params.K_packed + (kk >> 1u); + let word = t_weight[byte_idx >> 2u]; + let b = (word >> ((byte_idx & 3u) * 8u)) & 0xFFu; + var nib: u32; + if ((kk & 1u) == 0u) { nib = b & 0x0Fu; } else { nib = (b >> 4u) & 0x0Fu; } + dqv = f16(i32(nib) - 8) * f16(t_scales[scale_row + n]); + } + Bs[br * BN + bc + j] = dqv; + } + workgroupBarrier(); + for (var k: u32 = 0u; k < BK; k = k + 1u) { + var a: array; + var bvec: array; + for (var m: u32 = 0u; m < 4u; m = m + 1u) { a[m] = As[(lid.y * 4u + m) * BK + k]; } + for (var n: u32 = 0u; n < 4u; n = n + 1u) { bvec[n] = Bs[k * BN + lid.x * 4u + n]; } + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = acc[m][n] + f32(a[m] * bvec[n]); } + } + } + workgroupBarrier(); + k0 = k0 + BK; + } + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { + let r = row0 + lid.y * 4u + m; + let c = col0 + lid.x * 4u + n; + if (r < params.M && c < params.N) { + var v = acc[m][n]; + if (params.has_bias != 0u) { v = v + t_bias[c]; } + t_out[r * params.N + c] = v; + } + } + } +} +)"; + +inline constexpr uint32_t kQ4gswLinearGemmSteelHalfWorkgroupSizeX = 16; +inline constexpr uint32_t kQ4gswLinearGemmSteelHalfWorkgroupSizeY = 16; +inline constexpr uint32_t kQ4gswLinearGemmSteelHalfWorkgroupSizeZ = 1; + +} // namespace executorch::backends::webgpu diff --git a/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_wgsl.h b/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_wgsl.h new file mode 100644 index 00000000000..71b0f45bdf7 --- /dev/null +++ b/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_wgsl.h @@ -0,0 +1,134 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#pragma once + +#include + +namespace executorch::backends::webgpu { + +// @generated from q4gsw_linear_gemm_steel.wgsl - DO NOT EDIT. +// wgsl-sha256: dd771b9ab096410f3ad0d9259bef7816e41330a434325ea28baa5abbfb2841d2 +inline constexpr const char* kQ4gswLinearGemmSteelWGSL = R"( +@group(0) @binding(0) var t_out: array; +@group(0) @binding(1) var t_input: array; +@group(0) @binding(2) var t_weight: array; +@group(0) @binding(3) var t_scales: array; +@group(0) @binding(4) var t_bias: array; + +struct Params { + M: u32, + N: u32, + K: u32, + K_packed: u32, + group_size: u32, + padded_N: u32, + has_bias: u32, + _pad: u32, +} +@group(0) @binding(5) var params: Params; + +// "steel" prefill GEMM (M>1): 64x64 tile, 256 threads; K%16==0 host-guarded. +// The "steel" name + register-tiled dequant-to-shared GEMM structure are +// inspired by MLX's steel GEMM kernels (github.com/ml-explore/mlx, +// mlx/backend/metal/kernels/steel). One template, four variants: +// DTYPE=float f32 storage/multiply, per-nibble weight staging. +// DTYPE=half f16 storage/multiply, per-nibble weight staging. +// PWDQ (half only) packed-word dequant: load each u32 weight word ONCE, +// unpack all 16 nibbles of a column + hoist the per-column scale to one read +// (the per-nibble path re-reads each word ~8x). Requires K%BK==0 (steel +// route guarantees it) and group_size%BK==0 (hoisted scale across the tile). +// ACC=half (PWDQ only) f16 accumulate with fma(), cast to f32 in the epilogue +// -- LOSSY, perplexity-gated, opt-in via a runtime spec. ACC=float is f32 +// accumulate -- BIT-EXACT to the per-nibble half kernel. +const BM: u32 = 64u; const BN: u32 = 64u; const BK: u32 = 16u; +var As: array; // BM*BK +var Bs: array; // BK*BN +@compute @workgroup_size(16, 16) +fn main(@builtin(workgroup_id) wid: vec3, + @builtin(local_invocation_id) lid: vec3) { + let nbN = (params.N + BN - 1u) / BN; + let bx = wid.x % nbN; // decode 2D tile id from 1D dispatch + let by = wid.x / nbN; + let row0 = by * BM; + let col0 = bx * BN; + let tid = lid.y * 16u + lid.x; + var acc: array, 4>; + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = 0.0; } + } + // A staging coords: 256 threads load 64x16 = 1024 f32 -> 4 rows each (4 contiguous K). + let ar = tid / 4u; // 0..63 (row in tile) + let ac = (tid % 4u) * 4u; // 0,4,8,12 (K offset, 4 contiguous) + // B staging coords: 256 threads load 16x64 = 1024 dequant weights -> 4 cols each. + let br = tid / 16u; // 0..15 (K within BK) + let bc = (tid % 16u) * 4u; // 0,4,..60 (N offset, 4 contiguous) + + var k0: u32 = 0u; + loop { + if (k0 >= params.K) { break; } + // stage activations (edge-masked on M; K is a multiple of BK for our shapes) + let arow = row0 + ar; + if (arow < params.M) { + let base = arow * params.K + k0 + ac; + As[ar * BK + ac + 0u] = f32(t_input[base]); + As[ar * BK + ac + 1u] = f32(t_input[base + 1u]); + As[ar * BK + ac + 2u] = f32(t_input[base + 2u]); + As[ar * BK + ac + 3u] = f32(t_input[base + 3u]); + } else { + As[ar * BK + ac + 0u] = 0.0; As[ar * BK + ac + 1u] = 0.0; + As[ar * BK + ac + 2u] = 0.0; As[ar * BK + ac + 3u] = 0.0; + } + // stage DEQUANTIZED weights into Bs[k][n]: 4 contiguous N per thread. + let kk = k0 + br; // K index for this shmem row + let scale_row = (kk / params.group_size) * params.padded_N; + for (var j: u32 = 0u; j < 4u; j = j + 1u) { + let n = col0 + bc + j; + var dqv: f32 = 0.0; + if (n < params.N) { + let byte_idx = n * params.K_packed + (kk >> 1u); + let word = t_weight[byte_idx >> 2u]; + let b = (word >> ((byte_idx & 3u) * 8u)) & 0xFFu; + var nib: u32; + if ((kk & 1u) == 0u) { nib = b & 0x0Fu; } else { nib = (b >> 4u) & 0x0Fu; } + dqv = f32(i32(nib) - 8) * t_scales[scale_row + n]; + } + Bs[br * BN + bc + j] = dqv; + } + workgroupBarrier(); + for (var k: u32 = 0u; k < BK; k = k + 1u) { + var a: array; + var bvec: array; + for (var m: u32 = 0u; m < 4u; m = m + 1u) { a[m] = As[(lid.y * 4u + m) * BK + k]; } + for (var n: u32 = 0u; n < 4u; n = n + 1u) { bvec[n] = Bs[k * BN + lid.x * 4u + n]; } + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = acc[m][n] + a[m] * bvec[n]; } + } + } + workgroupBarrier(); + k0 = k0 + BK; + } + for (var m: u32 = 0u; m < 4u; m = m + 1u) { + for (var n: u32 = 0u; n < 4u; n = n + 1u) { + let r = row0 + lid.y * 4u + m; + let c = col0 + lid.x * 4u + n; + if (r < params.M && c < params.N) { + var v = acc[m][n]; + if (params.has_bias != 0u) { v = v + t_bias[c]; } + t_out[r * params.N + c] = v; + } + } + } +} +)"; + +inline constexpr uint32_t kQ4gswLinearGemmSteelWorkgroupSizeX = 16; +inline constexpr uint32_t kQ4gswLinearGemmSteelWorkgroupSizeY = 16; +inline constexpr uint32_t kQ4gswLinearGemmSteelWorkgroupSizeZ = 1; + +} // namespace executorch::backends::webgpu diff --git a/backends/webgpu/runtime/ops/rms_norm/rms_norm.wgsl b/backends/webgpu/runtime/ops/rms_norm/rms_norm.wgsl index 4bd5618596f..1234174981e 100644 --- a/backends/webgpu/runtime/ops/rms_norm/rms_norm.wgsl +++ b/backends/webgpu/runtime/ops/rms_norm/rms_norm.wgsl @@ -1,6 +1,8 @@ -@group(0) @binding(0) var t_out: array; -@group(0) @binding(1) var t_in: array; -@group(0) @binding(2) var t_weight: array; +$if DTYPE == "half": + enable f16; +@group(0) @binding(0) var t_out: array<${buffer_gvec_type(DTYPE, VEC)}>; +@group(0) @binding(1) var t_in: array<${buffer_gvec_type(DTYPE, VEC)}>; +@group(0) @binding(2) var t_weight: array<${buffer_gvec_type(DTYPE, VEC)}>; struct Params { num_rows: u32, @@ -12,7 +14,7 @@ struct Params { const WG_SIZE: u32 = 64u; -var shared_sum: array; +var shared_sum: array<${accum_scalar_type(DTYPE)}, WG_SIZE>; fn reduce_shared(worker_id: u32) { workgroupBarrier(); @@ -29,6 +31,10 @@ fn reduce_shared(worker_id: u32) { } } +$if VEC == 4: + // vec4 variant of rms_norm: each lane strides by WG_SIZE over rw4 = row_width/4 + // texels and accumulates dot(v, v). row_width is the ELEMENT count, so mean_sq + // divides by it (not rw4). The host selects this only when row_width % 4 == 0. @compute @workgroup_size(64, 1, 1) fn main( @builtin(workgroup_id) wid: vec3, @@ -40,18 +46,39 @@ fn main( return; } - let base = row_idx * params.row_width; + $if VEC == 4: + let rw4 = params.row_width / 4u; + let base4 = row_idx * rw4; + $else: + let base = row_idx * params.row_width; - var local_sq_sum: f32 = 0.0; - var x: u32 = worker_id; - loop { - if (x >= params.row_width) { - break; + var local_sq_sum: ${accum_scalar_type(DTYPE)} = 0.0; + $if VEC == 4: + var x4: u32 = worker_id; + loop { + if (x4 >= rw4) { + break; + } + let v = t_in[base4 + x4]; + $if DTYPE == "half": + local_sq_sum = local_sq_sum + dot(vec4(v), vec4(v)); + $else: + local_sq_sum = local_sq_sum + dot(v, v); + x4 = x4 + WG_SIZE; + } + $else: + var x: u32 = worker_id; + loop { + if (x >= params.row_width) { + break; + } + let v = t_in[base + x]; + $if DTYPE == "half": + local_sq_sum = local_sq_sum + f32(v) * f32(v); + $else: + local_sq_sum = local_sq_sum + v * v; + x = x + WG_SIZE; } - let v = t_in[base + x]; - local_sq_sum = local_sq_sum + v * v; - x = x + WG_SIZE; - } shared_sum[worker_id] = local_sq_sum; reduce_shared(worker_id); @@ -59,14 +86,30 @@ fn main( let mean_sq = shared_sum[0] / f32(params.row_width); let rstd = inverseSqrt(mean_sq + params.epsilon); - x = worker_id; - loop { - if (x >= params.row_width) { - break; + $if VEC == 4: + x4 = worker_id; + loop { + if (x4 >= rw4) { + break; + } + $if DTYPE == "half": + t_out[base4 + x4] = vec4(vec4(t_in[base4 + x4]) * rstd * vec4(t_weight[x4])); + $else: + t_out[base4 + x4] = t_in[base4 + x4] * rstd * t_weight[x4]; + x4 = x4 + WG_SIZE; + } + $else: + x = worker_id; + loop { + if (x >= params.row_width) { + break; + } + let v = t_in[base + x]; + let w = t_weight[x]; + $if DTYPE == "half": + t_out[base + x] = f16(f32(v) * rstd * f32(w)); + $else: + t_out[base + x] = v * rstd * w; + x = x + WG_SIZE; } - let v = t_in[base + x]; - let w = t_weight[x]; - t_out[base + x] = v * rstd * w; - x = x + WG_SIZE; - } } diff --git a/backends/webgpu/runtime/ops/rms_norm/rms_norm.yaml b/backends/webgpu/runtime/ops/rms_norm/rms_norm.yaml new file mode 100644 index 00000000000..6bfeeb012c3 --- /dev/null +++ b/backends/webgpu/runtime/ops/rms_norm/rms_norm.yaml @@ -0,0 +1,12 @@ +rms_norm: + parameter_names_with_default_values: + DTYPE: float + VEC: 1 + generate_variant_forall: + VEC: + - VALUE: 1 + SUFFIX: "" + - VALUE: 4 + SUFFIX: vec4 + shader_variants: + - NAME: rms_norm diff --git a/backends/webgpu/runtime/ops/rms_norm/rms_norm_vec4.wgsl b/backends/webgpu/runtime/ops/rms_norm/rms_norm_vec4.wgsl deleted file mode 100644 index c2f731e5f60..00000000000 --- a/backends/webgpu/runtime/ops/rms_norm/rms_norm_vec4.wgsl +++ /dev/null @@ -1,74 +0,0 @@ -@group(0) @binding(0) var t_out: array>; -@group(0) @binding(1) var t_in: array>; -@group(0) @binding(2) var t_weight: array>; - -struct Params { - num_rows: u32, - row_width: u32, - epsilon: f32, - _pad: u32, -} -@group(0) @binding(3) var params: Params; - -const WG_SIZE: u32 = 64u; - -var shared_sum: array; - -fn reduce_shared(worker_id: u32) { - workgroupBarrier(); - var stride: u32 = WG_SIZE / 2u; - loop { - if (stride == 0u) { - break; - } - if (worker_id < stride) { - shared_sum[worker_id] = shared_sum[worker_id] + shared_sum[worker_id + stride]; - } - workgroupBarrier(); - stride = stride >> 1u; - } -} - -// vec4 variant of rms_norm: each lane strides by WG_SIZE over rw4 = row_width/4 -// texels and accumulates dot(v, v). row_width is the ELEMENT count, so mean_sq -// divides by it (not rw4). The host selects this only when row_width % 4 == 0. -@compute @workgroup_size(64, 1, 1) -fn main( - @builtin(workgroup_id) wid: vec3, - @builtin(local_invocation_id) lid: vec3) { - let row_idx = wid.x; - let worker_id = lid.x; - - if (row_idx >= params.num_rows) { - return; - } - - let rw4 = params.row_width / 4u; - let base4 = row_idx * rw4; - - var local_sq_sum: f32 = 0.0; - var x4: u32 = worker_id; - loop { - if (x4 >= rw4) { - break; - } - let v = t_in[base4 + x4]; - local_sq_sum = local_sq_sum + dot(v, v); - x4 = x4 + WG_SIZE; - } - - shared_sum[worker_id] = local_sq_sum; - reduce_shared(worker_id); - - let mean_sq = shared_sum[0] / f32(params.row_width); - let rstd = inverseSqrt(mean_sq + params.epsilon); - - x4 = worker_id; - loop { - if (x4 >= rw4) { - break; - } - t_out[base4 + x4] = t_in[base4 + x4] * rstd * t_weight[x4]; - x4 = x4 + WG_SIZE; - } -} diff --git a/backends/webgpu/runtime/ops/rms_norm/rms_norm_vec4_wgsl.h b/backends/webgpu/runtime/ops/rms_norm/rms_norm_vec4_wgsl.h index 633bf3adfc0..02213e3d7b0 100644 --- a/backends/webgpu/runtime/ops/rms_norm/rms_norm_vec4_wgsl.h +++ b/backends/webgpu/runtime/ops/rms_norm/rms_norm_vec4_wgsl.h @@ -12,7 +12,7 @@ namespace executorch::backends::webgpu { -// @generated from rms_norm_vec4.wgsl - DO NOT EDIT. +// @generated from rms_norm.wgsl - DO NOT EDIT. // wgsl-sha256: 4c0ba56708bf125a7ec6ea3c51d1288e05ac00a8e2cfa10e38e9a208e230b8df inline constexpr const char* kRmsNormVec4WGSL = R"( @group(0) @binding(0) var t_out: array>; diff --git a/backends/webgpu/runtime/ops/sdpa/Sdpa.cpp b/backends/webgpu/runtime/ops/sdpa/Sdpa.cpp index 17918863a6e..50321ba4bdf 100644 --- a/backends/webgpu/runtime/ops/sdpa/Sdpa.cpp +++ b/backends/webgpu/runtime/ops/sdpa/Sdpa.cpp @@ -9,10 +9,13 @@ #include #include #include +#include #include +#include #include #include #include +#include #include #include @@ -255,9 +258,13 @@ static WGPUBuffer record_update_cache_dispatch( WGPUBuffer ubuf = graph.make_uniform_buffer(&uc, sizeof(uc)); BufferBinding bindings[2] = { {cache.buffer, cache.nbytes}, {src.buffer, src.nbytes}}; + const char* uc_src = kUpdateCacheWGSL; + if (graph.kv_f16()) { + uc_src = kUpdateCacheHalfWGSL; + } build_dispatch( graph, - kUpdateCacheWGSL, + uc_src, bindings, 2, ubuf, @@ -474,8 +481,11 @@ void sdpa_with_kv_cache_impl(WebGPUGraph& graph, const std::vector& args) { } // QK/softmax scratch — allocated only on the non-FD path (Hq*S*Cmax prefill). - WGPUBuffer attn_weights = graph.create_scratch_buffer(aw_bytes); - WGPUBuffer attn_weights_softmax = graph.create_scratch_buffer(aw_bytes); + WGPUBuffer attn_weights = graph.acquire_scratch(aw_bytes); + WebGPUGraph::ScopedScratch attn_weights_guard(&graph, attn_weights); + WGPUBuffer attn_weights_softmax = graph.acquire_scratch(aw_bytes); + WebGPUGraph::ScopedScratch attn_weights_softmax_guard( + &graph, attn_weights_softmax); // --- Dispatch 3: QK -> attn_weights. One thread per TM x TN tile. { @@ -494,9 +504,13 @@ void sdpa_with_kv_cache_impl(WebGPUGraph& graph, const std::vector& args) { {attn_weights, aw_bytes}, {q.buffer, q.nbytes}, {k_cache.buffer, k_cache.nbytes}}; + const char* qk_src = kSdpaComputeAttnWeightsWGSL; + if (graph.kv_f16()) { + qk_src = kSdpaComputeAttnWeightsHalfWGSL; + } build_dispatch( graph, - kSdpaComputeAttnWeightsWGSL, + qk_src, bindings, 3, ubuf, @@ -547,9 +561,13 @@ void sdpa_with_kv_cache_impl(WebGPUGraph& graph, const std::vector& args) { {out.buffer, out.nbytes}, {attn_weights_softmax, aw_bytes}, {v_cache.buffer, v_cache.nbytes}}; + const char* av_src = kSdpaComputeOutWGSL; + if (graph.kv_f16()) { + av_src = kSdpaComputeOutHalfWGSL; + } build_dispatch( graph, - kSdpaComputeOutWGSL, + av_src, bindings, 3, ubuf, diff --git a/backends/webgpu/runtime/ops/sdpa/sdpa_compute_attn_weights.wgsl b/backends/webgpu/runtime/ops/sdpa/sdpa_compute_attn_weights.wgsl index 014f0039048..097d87ecce0 100644 --- a/backends/webgpu/runtime/ops/sdpa/sdpa_compute_attn_weights.wgsl +++ b/backends/webgpu/runtime/ops/sdpa/sdpa_compute_attn_weights.wgsl @@ -1,6 +1,8 @@ +$if DTYPE == "half": + enable f16; @group(0) @binding(0) var t_attn_weights: array; @group(0) @binding(1) var t_q: array>; -@group(0) @binding(2) var t_k_cache: array>; +@group(0) @binding(2) var t_k_cache: array<${buffer_gvec_type(DTYPE, 4)}>; struct Params { S: u32, @@ -36,7 +38,10 @@ fn load_k_vec4(c: u32, kvh: u32, d4: u32) -> vec4 { return vec4(0.0, 0.0, 0.0, 0.0); } let base = c * params.Hkv * params.D + kvh * params.D + d4; - return t_k_cache[base / 4u]; + $if DTYPE == "half": + return vec4(t_k_cache[base / 4u]); + $else: + return t_k_cache[base / 4u]; } fn store_qk(s: u32, c: u32, h: u32, raw: f32) { diff --git a/backends/webgpu/runtime/ops/sdpa/sdpa_compute_attn_weights.yaml b/backends/webgpu/runtime/ops/sdpa/sdpa_compute_attn_weights.yaml new file mode 100644 index 00000000000..d031e2864f4 --- /dev/null +++ b/backends/webgpu/runtime/ops/sdpa/sdpa_compute_attn_weights.yaml @@ -0,0 +1,11 @@ +sdpa_compute_attn_weights: + parameter_names_with_default_values: + DTYPE: float + generate_variant_forall: + DTYPE: + - VALUE: float + SUFFIX: "" + - VALUE: half + SUFFIX: half + shader_variants: + - NAME: sdpa_compute_attn_weights diff --git a/backends/webgpu/runtime/ops/sdpa/sdpa_compute_attn_weights_half_wgsl.h b/backends/webgpu/runtime/ops/sdpa/sdpa_compute_attn_weights_half_wgsl.h new file mode 100644 index 00000000000..dc6f5858c7d --- /dev/null +++ b/backends/webgpu/runtime/ops/sdpa/sdpa_compute_attn_weights_half_wgsl.h @@ -0,0 +1,145 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#pragma once + +#include + +namespace executorch::backends::webgpu { + +// @generated from sdpa_compute_attn_weights.wgsl - DO NOT EDIT. +// wgsl-sha256: c8795d66b9b51516795fb0113b21fd55086b4a54d9a4b81c7f394ffd96d117b3 +inline constexpr const char* kSdpaComputeAttnWeightsHalfWGSL = R"( +enable f16; +@group(0) @binding(0) var t_attn_weights: array; +@group(0) @binding(1) var t_q: array>; +@group(0) @binding(2) var t_k_cache: array>; + +struct Params { + S: u32, + Hq: u32, + Hkv: u32, + D: u32, + context_len: u32, + input_pos: u32, + g: u32, + scale: f32, +} +@group(0) @binding(3) var params: Params; + +// WGSL forbids literal -inf; large finite negative is a WGSL-safe stand-in. +const NEG_INF: f32 = -1.0e30; + +override wg_size: u32 = 64; + +const TM: u32 = 4u; +const TN: u32 = 4u; + +// D is a multiple of 4 (host-guarded), so a d4 chunk is fully in-bounds — no per-lane check. +fn load_q_vec4(s: u32, h: u32, d4: u32) -> vec4 { + if (s >= params.S) { + return vec4(0.0, 0.0, 0.0, 0.0); + } + let base = s * params.Hq * params.D + h * params.D + d4; + return t_q[base / 4u]; +} + +fn load_k_vec4(c: u32, kvh: u32, d4: u32) -> vec4 { + if (c >= params.context_len) { + return vec4(0.0, 0.0, 0.0, 0.0); + } + let base = c * params.Hkv * params.D + kvh * params.D + d4; + return vec4(t_k_cache[base / 4u]); +} + +fn store_qk(s: u32, c: u32, h: u32, raw: f32) { + if (s >= params.S || c >= params.context_len) { + return; + } + var val = raw * params.scale; + // Causal mask: position c may not attend beyond s + input_pos. + if (c > s + params.input_pos) { + val = NEG_INF; + } + let idx = h * params.S * params.context_len + s * params.context_len + c; + t_attn_weights[idx] = val; +} + +@compute @workgroup_size(wg_size, 1, 1) +fn main( + @builtin(global_invocation_id) gid: vec3, + @builtin(num_workgroups) num_workgroups: vec3) { + let nrt = (params.S + TM - 1u) / TM; + let nct = (params.context_len + TN - 1u) / TN; + let tiles = nrt * nct; + let total = tiles * params.Hq; + // 2D dispatch fold: recover the linear tile index across x/y. + let idx = gid.x + gid.y * (num_workgroups.x * wg_size); + if (idx >= total) { + return; + } + + let h = idx / tiles; + let rem = idx % tiles; + let row_tile = rem / nct; + let col_tile = rem % nct; + let kvh = h / params.g; + let s0 = row_tile * TM; + let c0 = col_tile * TN; + + var acc: array, 4>; + acc[0] = vec4(0.0, 0.0, 0.0, 0.0); + acc[1] = vec4(0.0, 0.0, 0.0, 0.0); + acc[2] = vec4(0.0, 0.0, 0.0, 0.0); + acc[3] = vec4(0.0, 0.0, 0.0, 0.0); + + // Skip fully-masked causal tiles; mirrors Vulkan attn_weights_tiled.glsl. + let skip_tile = c0 > s0 + (TM - 1u) + params.input_pos; + var d4: u32 = 0u; + loop { + if (d4 >= params.D || skip_tile) { + break; + } + var q: array, TM>; + var k: array, TN>; + for (var i: u32 = 0u; i < TM; i = i + 1u) { + q[i] = load_q_vec4(s0 + i, h, d4); + } + for (var j: u32 = 0u; j < TN; j = j + 1u) { + k[j] = load_k_vec4(c0 + j, kvh, d4); + } + for (var i: u32 = 0u; i < TM; i = i + 1u) { + acc[i] += vec4( + dot(q[i], k[0]), + dot(q[i], k[1]), + dot(q[i], k[2]), + dot(q[i], k[3])); + } + d4 = d4 + 4u; + } + + var m: u32 = 0u; + loop { + if (m >= TM) { + break; + } + let av = acc[m]; + store_qk(s0 + m, c0 + 0u, h, av.x); + store_qk(s0 + m, c0 + 1u, h, av.y); + store_qk(s0 + m, c0 + 2u, h, av.z); + store_qk(s0 + m, c0 + 3u, h, av.w); + m = m + 1u; + } +} +)"; + +inline constexpr uint32_t kSdpaComputeAttnWeightsHalfWorkgroupSizeX = 64; +inline constexpr uint32_t kSdpaComputeAttnWeightsHalfWorkgroupSizeY = 1; +inline constexpr uint32_t kSdpaComputeAttnWeightsHalfWorkgroupSizeZ = 1; + +} // namespace executorch::backends::webgpu diff --git a/backends/webgpu/runtime/ops/sdpa/sdpa_compute_out.wgsl b/backends/webgpu/runtime/ops/sdpa/sdpa_compute_out.wgsl index 713345c0afa..53067447832 100644 --- a/backends/webgpu/runtime/ops/sdpa/sdpa_compute_out.wgsl +++ b/backends/webgpu/runtime/ops/sdpa/sdpa_compute_out.wgsl @@ -1,6 +1,8 @@ +$if DTYPE == "half": + enable f16; @group(0) @binding(0) var t_out: array>; @group(0) @binding(1) var t_attn_weights_softmax: array; -@group(0) @binding(2) var t_v_cache: array>; +@group(0) @binding(2) var t_v_cache: array<${buffer_gvec_type(DTYPE, 4)}>; struct Params { S: u32, @@ -38,7 +40,10 @@ fn load_v_d4(c: u32, kvh: u32, d0: u32) -> vec4 { return vec4(0.0, 0.0, 0.0, 0.0); } let base = c * params.Hkv * params.D + kvh * params.D + d0; - return t_v_cache[base / 4u]; + $if DTYPE == "half": + return vec4(t_v_cache[base / 4u]); + $else: + return t_v_cache[base / 4u]; } // Branch-free loaders for the aligned body: caller guarantees c4..c4+3 < context_len. @@ -52,7 +57,10 @@ fn load_a_vec4_nc(s: u32, h: u32, c4: u32) -> vec4 { fn load_v_d4_nc(c: u32, kvh: u32, d0: u32) -> vec4 { let base = c * params.Hkv * params.D + kvh * params.D + d0; - return t_v_cache[base / 4u]; + $if DTYPE == "half": + return vec4(t_v_cache[base / 4u]); + $else: + return t_v_cache[base / 4u]; } fn store_out_vec4(s: u32, d0: u32, h: u32, val: vec4) { diff --git a/backends/webgpu/runtime/ops/sdpa/sdpa_compute_out.yaml b/backends/webgpu/runtime/ops/sdpa/sdpa_compute_out.yaml new file mode 100644 index 00000000000..14b418f8a60 --- /dev/null +++ b/backends/webgpu/runtime/ops/sdpa/sdpa_compute_out.yaml @@ -0,0 +1,11 @@ +sdpa_compute_out: + parameter_names_with_default_values: + DTYPE: float + generate_variant_forall: + DTYPE: + - VALUE: float + SUFFIX: "" + - VALUE: half + SUFFIX: half + shader_variants: + - NAME: sdpa_compute_out diff --git a/backends/webgpu/runtime/ops/sdpa/sdpa_compute_out_half_wgsl.h b/backends/webgpu/runtime/ops/sdpa/sdpa_compute_out_half_wgsl.h new file mode 100644 index 00000000000..228921af58b --- /dev/null +++ b/backends/webgpu/runtime/ops/sdpa/sdpa_compute_out_half_wgsl.h @@ -0,0 +1,163 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#pragma once + +#include + +namespace executorch::backends::webgpu { + +// @generated from sdpa_compute_out.wgsl - DO NOT EDIT. +// wgsl-sha256: ed9709c966538edf2cbc6be97c284b89a9d921b6a4dbf115c6cbd76af301a1be +inline constexpr const char* kSdpaComputeOutHalfWGSL = R"( +enable f16; +@group(0) @binding(0) var t_out: array>; +@group(0) @binding(1) var t_attn_weights_softmax: array; +@group(0) @binding(2) var t_v_cache: array>; + +struct Params { + S: u32, + Hq: u32, + Hkv: u32, + D: u32, + context_len: u32, + g: u32, + _pad0: u32, + _pad1: u32, +} +@group(0) @binding(3) var params: Params; + +override wg_size: u32 = 64; + +const TM: u32 = 4u; +const TN: u32 = 4u; + +// Checked loaders mask context lanes past context_len (D%4==0, host-guarded). +fn load_a_vec4(s: u32, h: u32, c4: u32) -> vec4 { + var r = vec4(0.0, 0.0, 0.0, 0.0); + if (s >= params.S) { + return r; + } + let base = h * params.S * params.context_len + s * params.context_len; + if (c4 + 0u < params.context_len) { r.x = t_attn_weights_softmax[base + c4 + 0u]; } + if (c4 + 1u < params.context_len) { r.y = t_attn_weights_softmax[base + c4 + 1u]; } + if (c4 + 2u < params.context_len) { r.z = t_attn_weights_softmax[base + c4 + 2u]; } + if (c4 + 3u < params.context_len) { r.w = t_attn_weights_softmax[base + c4 + 3u]; } + return r; +} + +fn load_v_d4(c: u32, kvh: u32, d0: u32) -> vec4 { + if (c >= params.context_len) { + return vec4(0.0, 0.0, 0.0, 0.0); + } + let base = c * params.Hkv * params.D + kvh * params.D + d0; + return vec4(t_v_cache[base / 4u]); +} + +// Branch-free loaders for the aligned body: caller guarantees c4..c4+3 < context_len. +fn load_a_vec4_nc(s: u32, h: u32, c4: u32) -> vec4 { + if (s >= params.S) { + return vec4(0.0, 0.0, 0.0, 0.0); + } + let base = h * params.S * params.context_len + s * params.context_len + c4; + return vec4(t_attn_weights_softmax[base], t_attn_weights_softmax[base + 1u], t_attn_weights_softmax[base + 2u], t_attn_weights_softmax[base + 3u]); +} + +fn load_v_d4_nc(c: u32, kvh: u32, d0: u32) -> vec4 { + let base = c * params.Hkv * params.D + kvh * params.D + d0; + return vec4(t_v_cache[base / 4u]); +} + +fn store_out_vec4(s: u32, d0: u32, h: u32, val: vec4) { + if (s >= params.S) { + return; + } + let idx = s * params.Hq * params.D + h * params.D + d0; + t_out[idx / 4u] = val; +} + +@compute @workgroup_size(wg_size, 1, 1) +fn main( + @builtin(global_invocation_id) gid: vec3, + @builtin(num_workgroups) num_workgroups: vec3) { + let nrt = (params.S + TM - 1u) / TM; + let nct = (params.D + TN - 1u) / TN; + let tiles = nrt * nct; + let total = tiles * params.Hq; + // 2D dispatch fold: recover the linear tile index across x/y. + let idx = gid.x + gid.y * (num_workgroups.x * wg_size); + if (idx >= total) { + return; + } + + let h = idx / tiles; + let rem = idx % tiles; + let row_tile = rem / nct; + let col_tile = rem % nct; + let kvh = h / params.g; + let s0 = row_tile * TM; + let d0 = col_tile * TN; + + var acc: array, 4>; + acc[0] = vec4(0.0, 0.0, 0.0, 0.0); + acc[1] = vec4(0.0, 0.0, 0.0, 0.0); + acc[2] = vec4(0.0, 0.0, 0.0, 0.0); + acc[3] = vec4(0.0, 0.0, 0.0, 0.0); + + // Branch-free aligned body + checked tail; mirrors Vulkan out_tiled.glsl. + let ctx_aligned = params.context_len - (params.context_len & 3u); + var c4: u32 = 0u; + loop { + if (c4 >= ctx_aligned) { + break; + } + let a0 = load_a_vec4_nc(s0 + 0u, h, c4); + let a1 = load_a_vec4_nc(s0 + 1u, h, c4); + let a2 = load_a_vec4_nc(s0 + 2u, h, c4); + let a3 = load_a_vec4_nc(s0 + 3u, h, c4); + let v0 = load_v_d4_nc(c4 + 0u, kvh, d0); + let v1 = load_v_d4_nc(c4 + 1u, kvh, d0); + let v2 = load_v_d4_nc(c4 + 2u, kvh, d0); + let v3 = load_v_d4_nc(c4 + 3u, kvh, d0); + acc[0] += a0.x * v0 + a0.y * v1 + a0.z * v2 + a0.w * v3; + acc[1] += a1.x * v0 + a1.y * v1 + a1.z * v2 + a1.w * v3; + acc[2] += a2.x * v0 + a2.y * v1 + a2.z * v2 + a2.w * v3; + acc[3] += a3.x * v0 + a3.y * v1 + a3.z * v2 + a3.w * v3; + c4 = c4 + 4u; + } + if (c4 < params.context_len) { + let a0 = load_a_vec4(s0 + 0u, h, c4); + let a1 = load_a_vec4(s0 + 1u, h, c4); + let a2 = load_a_vec4(s0 + 2u, h, c4); + let a3 = load_a_vec4(s0 + 3u, h, c4); + let v0 = load_v_d4(c4 + 0u, kvh, d0); + let v1 = load_v_d4(c4 + 1u, kvh, d0); + let v2 = load_v_d4(c4 + 2u, kvh, d0); + let v3 = load_v_d4(c4 + 3u, kvh, d0); + acc[0] += a0.x * v0 + a0.y * v1 + a0.z * v2 + a0.w * v3; + acc[1] += a1.x * v0 + a1.y * v1 + a1.z * v2 + a1.w * v3; + acc[2] += a2.x * v0 + a2.y * v1 + a2.z * v2 + a2.w * v3; + acc[3] += a3.x * v0 + a3.y * v1 + a3.z * v2 + a3.w * v3; + } + + var m: u32 = 0u; + loop { + if (m >= TM) { + break; + } + store_out_vec4(s0 + m, d0, h, acc[m]); + m = m + 1u; + } +} +)"; + +inline constexpr uint32_t kSdpaComputeOutHalfWorkgroupSizeX = 64; +inline constexpr uint32_t kSdpaComputeOutHalfWorkgroupSizeY = 1; +inline constexpr uint32_t kSdpaComputeOutHalfWorkgroupSizeZ = 1; + +} // namespace executorch::backends::webgpu diff --git a/backends/webgpu/runtime/ops/sdpa_fd_decode/SdpaFdDecode.cpp b/backends/webgpu/runtime/ops/sdpa_fd_decode/SdpaFdDecode.cpp index 70108beb892..ffd3b24dce3 100644 --- a/backends/webgpu/runtime/ops/sdpa_fd_decode/SdpaFdDecode.cpp +++ b/backends/webgpu/runtime/ops/sdpa_fd_decode/SdpaFdDecode.cpp @@ -12,6 +12,7 @@ #include #include #include +#include #include #include @@ -185,8 +186,10 @@ void sdpa_fd_decode_dispatch( static_cast(kSdpaFdMaxSplits) * static_cast(D); const uint64_t pml_floats = static_cast(Hq) * static_cast(kSdpaFdMaxSplits) * 2ull; - WGPUBuffer part_o = graph.create_scratch_buffer(po_floats * sizeof(float)); - WGPUBuffer part_ml = graph.create_scratch_buffer(pml_floats * sizeof(float)); + WGPUBuffer part_o = graph.acquire_scratch(po_floats * sizeof(float)); + WebGPUGraph::ScopedScratch part_o_guard(&graph, part_o); + WGPUBuffer part_ml = graph.acquire_scratch(pml_floats * sizeof(float)); + WebGPUGraph::ScopedScratch part_ml_guard(&graph, part_ml); // Pass 1: split (Hq*num_splits WGs) -> writes part_o, part_ml. FdSplitParams sp = {}; @@ -218,9 +221,13 @@ void sdpa_fd_decode_dispatch( static_cast(split_threads), kSdpaFdSplitWorkgroupSizeX, "fd_split"); + const char* split_shader = kSdpaFdSplitWGSL; + if (graph.kv_f16()) { + split_shader = kSdpaFdSplitHalfWGSL; + } build_dispatch( graph, - kSdpaFdSplitWGSL, + split_shader, split_bindings, 5, 2, diff --git a/backends/webgpu/runtime/ops/sdpa_fd_decode/sdpa_fd_split.wgsl b/backends/webgpu/runtime/ops/sdpa_fd_decode/sdpa_fd_split.wgsl index c14c6bd07bd..da67489cc4d 100644 --- a/backends/webgpu/runtime/ops/sdpa_fd_decode/sdpa_fd_split.wgsl +++ b/backends/webgpu/runtime/ops/sdpa_fd_decode/sdpa_fd_split.wgsl @@ -1,8 +1,10 @@ +$if DTYPE == "half": + enable f16; @group(0) @binding(0) var t_part_o: array; @group(0) @binding(1) var t_part_ml: array; @group(0) @binding(2) var t_q: array; -@group(0) @binding(3) var t_k_cache: array; -@group(0) @binding(4) var t_v_cache: array; +@group(0) @binding(3) var t_k_cache: array<${buffer_scalar_type(DTYPE)}>; +@group(0) @binding(4) var t_v_cache: array<${buffer_scalar_type(DTYPE)}>; struct Params { _pad0: u32, @@ -66,9 +68,14 @@ fn main( let qi = q_base + i4 * 4u; let ki = kvbase + i4 * 4u; let qv = vec4(t_q[qi], t_q[qi + 1u], t_q[qi + 2u], t_q[qi + 3u]); - let kvv = vec4( - t_k_cache[ki], t_k_cache[ki + 1u], - t_k_cache[ki + 2u], t_k_cache[ki + 3u]); + $if DTYPE == "half": + let kvv = vec4( + f32(t_k_cache[ki]), f32(t_k_cache[ki + 1u]), + f32(t_k_cache[ki + 2u]), f32(t_k_cache[ki + 3u])); + $else: + let kvv = vec4( + t_k_cache[ki], t_k_cache[ki + 1u], + t_k_cache[ki + 2u], t_k_cache[ki + 3u]); acc4 = acc4 + qv * kvv; } s = (acc4.x + acc4.y + acc4.z + acc4.w) * params.scale; @@ -107,7 +114,10 @@ fn main( var acc: f32 = rescale * o_acc[nd]; for (var j: u32 = 0u; j < n; j = j + 1u) { let vbase = (block + j) * kv_row_stride + kv * D; - acc = acc + sh_s[j] * t_v_cache[vbase + d]; + $if DTYPE == "half": + acc = acc + sh_s[j] * f32(t_v_cache[vbase + d]); + $else: + acc = acc + sh_s[j] * t_v_cache[vbase + d]; } o_acc[nd] = acc; } diff --git a/backends/webgpu/runtime/ops/sdpa_fd_decode/sdpa_fd_split.yaml b/backends/webgpu/runtime/ops/sdpa_fd_decode/sdpa_fd_split.yaml new file mode 100644 index 00000000000..283e6516996 --- /dev/null +++ b/backends/webgpu/runtime/ops/sdpa_fd_decode/sdpa_fd_split.yaml @@ -0,0 +1,11 @@ +sdpa_fd_split: + parameter_names_with_default_values: + DTYPE: float + generate_variant_forall: + DTYPE: + - VALUE: float + SUFFIX: "" + - VALUE: half + SUFFIX: half + shader_variants: + - NAME: sdpa_fd_split diff --git a/backends/webgpu/runtime/ops/sdpa_fd_decode/sdpa_fd_split_half_wgsl.h b/backends/webgpu/runtime/ops/sdpa_fd_decode/sdpa_fd_split_half_wgsl.h new file mode 100644 index 00000000000..bc69a444edb --- /dev/null +++ b/backends/webgpu/runtime/ops/sdpa_fd_decode/sdpa_fd_split_half_wgsl.h @@ -0,0 +1,156 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#pragma once + +#include + +namespace executorch::backends::webgpu { + +// @generated from sdpa_fd_split.wgsl - DO NOT EDIT. +// wgsl-sha256: 147fc5775b76f3626dc934b2df5e34148bef12c345789332f098d13143bb646e +inline constexpr const char* kSdpaFdSplitHalfWGSL = R"( +enable f16; +@group(0) @binding(0) var t_part_o: array; +@group(0) @binding(1) var t_part_ml: array; +@group(0) @binding(2) var t_q: array; +@group(0) @binding(3) var t_k_cache: array; +@group(0) @binding(4) var t_v_cache: array; + +struct Params { + _pad0: u32, + Hkv: u32, + D: u32, + context_len: u32, + g: u32, + num_splits: u32, + split_len: u32, + scale: f32, +} +@group(0) @binding(5) var params: Params; + +const WG_SIZE: u32 = 64u; +const MAX_SPLITS: u32 = 128u; +const MAX_D_PER_LANE: u32 = 2u; +const NEG_INF: f32 = -1.0e30; + +// sh_s: block scores then softmax weights; sh_red: max/sum reduction scratch. +var sh_s: array; +var sh_red: array; + +// FlashDecoding pass 1: per-(head,split) unnormalized softmax partial. +@compute @workgroup_size(64, 1, 1) +fn main( + @builtin(workgroup_id) wid: vec3, + @builtin(local_invocation_id) lid: vec3) { + let h = wid.x / params.num_splits; + let split_i = wid.x % params.num_splits; + let t = lid.x; + let D = params.D; + let D4 = D / 4u; // D is a multiple of 4 (guarded host-side); vec4 QK dot + let ctx = params.context_len; + let kv = h / params.g; + let q_base = h * D; + let kv_row_stride = params.Hkv * D; + + let c0 = split_i * params.split_len; + var c1 = c0 + params.split_len; + if (c1 > ctx) { c1 = ctx; } + + var m: f32 = NEG_INF; + var l: f32 = 0.0; + var o_acc: array; + for (var nd: u32 = 0u; nd < MAX_D_PER_LANE; nd = nd + 1u) { o_acc[nd] = 0.0; } + + // Stream the split in blocks of WG_SIZE KV positions. + var block: u32 = c0; + loop { + if (block >= c1) { break; } + var n: u32 = c1 - block; + if (n > WG_SIZE) { n = WG_SIZE; } + + // Phase 1: lane t computes the full QK dot for position block+t (vec4), one + // K row read once. Out-of-block lanes hold NEG_INF (safe for the max). + var s: f32 = NEG_INF; + if (t < n) { + let kvbase = (block + t) * kv_row_stride + kv * D; + var acc4 = vec4(0.0, 0.0, 0.0, 0.0); + for (var i4: u32 = 0u; i4 < D4; i4 = i4 + 1u) { + let qi = q_base + i4 * 4u; + let ki = kvbase + i4 * 4u; + let qv = vec4(t_q[qi], t_q[qi + 1u], t_q[qi + 2u], t_q[qi + 3u]); + let kvv = vec4( + f32(t_k_cache[ki]), f32(t_k_cache[ki + 1u]), + f32(t_k_cache[ki + 2u]), f32(t_k_cache[ki + 3u])); + acc4 = acc4 + qv * kvv; + } + s = (acc4.x + acc4.y + acc4.z + acc4.w) * params.scale; + } + sh_s[t] = s; + + // Phase 2a: block max via tree reduction (sh_red written from register s). + sh_red[t] = s; + workgroupBarrier(); + for (var stride: u32 = WG_SIZE / 2u; stride > 0u; stride = stride >> 1u) { + if (t < stride) { sh_red[t] = max(sh_red[t], sh_red[t + stride]); } + workgroupBarrier(); + } + let m_new = max(m, sh_red[0]); + let rescale = exp(m - m_new); + + // Phase 2b: each lane exponentiates ITS position once -> p (reuse sh_s), + // and reduce the block sum of p. + var p_t: f32 = 0.0; + if (t < n) { p_t = exp(sh_s[t] - m_new); } + workgroupBarrier(); // all reads of sh_s (the scores) done before overwrite + sh_s[t] = p_t; + sh_red[t] = p_t; + workgroupBarrier(); + for (var stride: u32 = WG_SIZE / 2u; stride > 0u; stride = stride >> 1u) { + if (t < stride) { sh_red[t] = sh_red[t] + sh_red[t + stride]; } + workgroupBarrier(); + } + l = rescale * l + sh_red[0]; + + // Phase 2c: each lane accumulates V for its own output dims over the block, + // reading the shared per-position weights (no exp in this loop). + for (var nd: u32 = 0u; nd < MAX_D_PER_LANE; nd = nd + 1u) { + let d = t + nd * WG_SIZE; + if (d < D) { + var acc: f32 = rescale * o_acc[nd]; + for (var j: u32 = 0u; j < n; j = j + 1u) { + let vbase = (block + j) * kv_row_stride + kv * D; + acc = acc + sh_s[j] * f32(t_v_cache[vbase + d]); + } + o_acc[nd] = acc; + } + } + m = m_new; + workgroupBarrier(); // before the next block overwrites sh_s / sh_red + block = block + WG_SIZE; + } + + let part = h * MAX_SPLITS + split_i; + for (var nd: u32 = 0u; nd < MAX_D_PER_LANE; nd = nd + 1u) { + let d = t + nd * WG_SIZE; + if (d < D) { + t_part_o[part * D + d] = o_acc[nd]; + } + } + if (t == 0u) { + t_part_ml[part * 2u + 0u] = m; + t_part_ml[part * 2u + 1u] = l; + } +} +)"; + +inline constexpr uint32_t kSdpaFdSplitHalfWorkgroupSizeX = 64; +inline constexpr uint32_t kSdpaFdSplitHalfWorkgroupSizeY = 1; +inline constexpr uint32_t kSdpaFdSplitHalfWorkgroupSizeZ = 1; + +} // namespace executorch::backends::webgpu diff --git a/backends/webgpu/runtime/ops/update_cache/update_cache.wgsl b/backends/webgpu/runtime/ops/update_cache/update_cache.wgsl index 62f882ad547..cdda4aac9cd 100644 --- a/backends/webgpu/runtime/ops/update_cache/update_cache.wgsl +++ b/backends/webgpu/runtime/ops/update_cache/update_cache.wgsl @@ -1,4 +1,6 @@ -@group(0) @binding(0) var t_cache: array; +$if DTYPE == "half": + enable f16; +@group(0) @binding(0) var t_cache: array<${buffer_scalar_type(DTYPE)}>; @group(0) @binding(1) var t_value: array; struct Params { @@ -20,5 +22,8 @@ fn main(@builtin(global_invocation_id) gid: vec3) { if (params.dst_offset + i >= params.cache_numel) { return; } - t_cache[params.dst_offset + i] = t_value[i]; + $if DTYPE == "half": + t_cache[params.dst_offset + i] = f16(t_value[i]); + $else: + t_cache[params.dst_offset + i] = t_value[i]; } diff --git a/backends/webgpu/runtime/ops/update_cache/update_cache.yaml b/backends/webgpu/runtime/ops/update_cache/update_cache.yaml new file mode 100644 index 00000000000..4df425c763e --- /dev/null +++ b/backends/webgpu/runtime/ops/update_cache/update_cache.yaml @@ -0,0 +1,11 @@ +update_cache: + parameter_names_with_default_values: + DTYPE: float + generate_variant_forall: + DTYPE: + - VALUE: float + SUFFIX: "" + - VALUE: half + SUFFIX: half + shader_variants: + - NAME: update_cache diff --git a/backends/webgpu/runtime/ops/update_cache/update_cache_half_wgsl.h b/backends/webgpu/runtime/ops/update_cache/update_cache_half_wgsl.h new file mode 100644 index 00000000000..d31728178a8 --- /dev/null +++ b/backends/webgpu/runtime/ops/update_cache/update_cache_half_wgsl.h @@ -0,0 +1,49 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#pragma once + +#include + +namespace executorch::backends::webgpu { + +// @generated from update_cache.wgsl - DO NOT EDIT. +// wgsl-sha256: 390daabe0d4545311dd5c6768d427fc9d133125bb1143869184c7d7631a88954 +inline constexpr const char* kUpdateCacheHalfWGSL = R"( +enable f16; +@group(0) @binding(0) var t_cache: array; +@group(0) @binding(1) var t_value: array; + +struct Params { + numel: u32, + dst_offset: u32, + cache_numel: u32, + _pad0: u32, +} +@group(0) @binding(2) var params: Params; + +override wg_size: u32 = 256; + +@compute @workgroup_size(wg_size, 1, 1) +fn main(@builtin(global_invocation_id) gid: vec3) { + let i = gid.x; + if (i >= params.numel) { + return; + } + if (params.dst_offset + i >= params.cache_numel) { + return; + } + t_cache[params.dst_offset + i] = f16(t_value[i]); +} +)"; + +inline constexpr uint32_t kUpdateCacheHalfWorkgroupSizeX = 256; +inline constexpr uint32_t kUpdateCacheHalfWorkgroupSizeY = 1; +inline constexpr uint32_t kUpdateCacheHalfWorkgroupSizeZ = 1; + +} // namespace executorch::backends::webgpu diff --git a/backends/webgpu/scripts/gen_wgsl_headers.py b/backends/webgpu/scripts/gen_wgsl_headers.py index 90293fc6cfe..c66aff2039a 100644 --- a/backends/webgpu/scripts/gen_wgsl_headers.py +++ b/backends/webgpu/scripts/gen_wgsl_headers.py @@ -7,21 +7,39 @@ """Generate runtime/ops//_wgsl.h from each .wgsl. -Each header embeds the shader verbatim as `inline constexpr const char* +Each header embeds the shader text unchanged as `inline constexpr const char* kWGSL` plus `kWorkgroupSize` (parsed from @workgroup_size). Usage: gen_wgsl_headers.py # (re)write all _wgsl.h gen_wgsl_headers.py --check # exit 1 if any committed header is stale -Stdlib only (the devserver has no third-party pip). +A shader is treated as a template iff a sibling .yaml spec exists; the +$-block engine (preprocess/escape/generate_variant_combinations) expands one +template + a DTYPE/VEC variant matrix into the concrete per-variant headers. + +Spec parsing uses PyYAML (a declared ExecuTorch codegen dependency, mirroring +backends/vulkan/runtime/gen_vulkan_spv.py); run under the ExecuTorch dev env. """ import argparse +import copy import hashlib +import io import re import sys +from itertools import product from pathlib import Path +from typing import Any, Dict, List, Optional, Set + +import yaml +from yaml.constructor import ConstructorError +from yaml.nodes import MappingNode + +try: + from yaml import CLoader as Loader +except ImportError: + from yaml import Loader # type: ignore[assignment, misc] BACKEND_ROOT = Path(__file__).resolve().parents[1] @@ -37,6 +55,328 @@ */""" +######################################################################## +# WGSL template engine +# +# A $-block transpiler (extract_leading_whitespace / escape / preprocess) +# plus a DTYPE/VEC variant matrix (generate_variant_combinations / +# parse_template_spec) expand one template + its YAML sidecar into the +# per-variant WGSL headers. +######################################################################## + + +# WGSL type-helpers injected into preprocess's exec globals so ${...} template +# expressions can spell WGSL types (f32/f16, vec4); @group/@binding layout is +# written directly in the templates. Names mirror gen_vulkan_spv.py's +# buffer_scalar_type / buffer_gvec_type / accum_scalar_type. +def buffer_scalar_type(dtype: str) -> str: + if dtype == "half": + return "f16" + elif dtype == "float": + return "f32" + return dtype + + +def buffer_gvec_type(dtype: str, n: int) -> str: + if n == 1: + return buffer_scalar_type(dtype) + return f"vec{n}<{buffer_scalar_type(dtype)}>" + + +def accum_scalar_type(dtype: str) -> str: + # The float family (incl. half) accumulates in f32 -- f16 accumulation is + # numerically unsafe on target GPUs. Mirrors gen_vulkan_spv.py's + # accum_scalar_type (half -> rgba16f -> "float" there is the same intent). + if dtype in ("half", "float"): + return "f32" + return buffer_scalar_type(dtype) + + +WGSL_HELPERS: Dict[str, Any] = { + "buffer_scalar_type": buffer_scalar_type, + "buffer_gvec_type": buffer_gvec_type, + "accum_scalar_type": accum_scalar_type, +} + + +# https://github.com/google/XNNPACK/blob/master/tools/xngen.py +def extract_leading_whitespace(line: str) -> str: + match = re.match(r"\s*", line) + return match.group(0) if match else "" + + +# https://github.com/google/XNNPACK/blob/master/tools/xngen.py +def escape(line: str) -> str: + output_parts = [] + while "${" in line: + start_pos = line.index("${") + end_pos = line.index("}", start_pos + 2) + if start_pos != 0: + output_parts.append('"' + line[:start_pos].replace('"', '\\"') + '"') + output_parts.append("str(" + line[start_pos + 2 : end_pos] + ")") + line = line[end_pos + 1 :] + if line: + output_parts.append('"' + line.replace('"', '\\"') + '"') + return " + ".join(output_parts) + + +# https://github.com/google/XNNPACK/blob/master/tools/xngen.py +def preprocess( + input_text: str, variables: Dict[str, Any], input_path: str = "codegen" +) -> str: + # Workaround to handle source files using \ to extend mecros to a new line + input_text = re.sub(r"\\$", r"\\\\", input_text, flags=re.MULTILINE) + + input_lines = input_text.splitlines() + python_lines = [] + + blank_lines = 0 + + last_indent = "" + + # List of tuples (total_index, python_indent) + indent_stack = [("", "")] + + # Indicates whether this is the first line inside Python + # code block (i.e. for, while, if, elif, else) + python_block_start = True + for input_line in input_lines: + if input_line == "": + blank_lines += 1 + continue + # Skip lint markers. + if "LINT" in input_line: + continue + + input_indent = extract_leading_whitespace(input_line) + if python_block_start: + assert input_indent.startswith(last_indent) + extra_python_indent = input_indent[len(last_indent) :] + python_indent = indent_stack[-1][1] + extra_python_indent + indent_stack.append((input_indent, python_indent)) + assert input_indent.startswith(indent_stack[-1][0]) + else: + while not input_indent.startswith(indent_stack[-1][0]): + del indent_stack[-1] + python_block_start = False + + python_indent = indent_stack[-1][1] + stripped_input_line = input_line.strip() + if stripped_input_line.startswith("$") and not stripped_input_line.startswith( + "${" + ): + if stripped_input_line.endswith(":"): + python_block_start = True + while blank_lines != 0: + python_lines.append(python_indent + "print(file=OUT_STREAM)") + blank_lines -= 1 + python_lines.append(python_indent + stripped_input_line.replace("$", "")) + else: + assert input_line.startswith(python_indent) + while blank_lines != 0: + python_lines.append(python_indent + "print(file=OUT_STREAM)") + blank_lines -= 1 + python_lines.append( + python_indent + + "print(%s, file=OUT_STREAM)" + % escape(input_line[len(python_indent) :]) + ) + last_indent = input_indent + + while blank_lines != 0: + python_lines.append(python_indent + "print(file=OUT_STREAM)") + blank_lines -= 1 + + exec_globals = dict(variables) + output_stream = io.StringIO() + exec_globals["OUT_STREAM"] = output_stream + + python_bytecode = compile("\n".join(python_lines), input_path, "exec") + exec(python_bytecode, exec_globals) + + return output_stream.getvalue() + + +# https://gist.github.com/pypt/94d747fe5180851196eb +class UniqueKeyLoader(Loader): + def construct_mapping(self, node, deep=False): # type: ignore[no-untyped-def] + if not isinstance(node, MappingNode): + raise ConstructorError( + None, + None, + f"expected a mapping node, but found {node.id}", + node.start_mark, + ) + mapping = {} + for key_node, value_node in node.value: + key = self.construct_object(key_node, deep=deep) # type: ignore[no-untyped-call] + try: + hash(key) + except TypeError as e: + raise ConstructorError( + "while constructing a mapping", + node.start_mark, + "found unacceptable key ", + key_node.start_mark, + ) from e + # check for duplicate keys + if key in mapping: + raise ConstructorError( + "while constructing a mapping", + node.start_mark, + "found duplicate key", + key_node.start_mark, + ) + value = self.construct_object(value_node, deep=deep) # type: ignore[no-untyped-call] + mapping[key] = value + return mapping + + +def generate_variant_combinations( # noqa: C901 + iterated_params: Dict[str, Any], + exclude_params: Optional[Set[str]] = None, +) -> List[Any]: + if exclude_params is None: + exclude_params = set() + all_iterated_params = [] + for param_name, value_list in iterated_params.items(): + if re.match(r"^combination\d*$", param_name): + param_values = [] + param_names = value_list["parameter_names"] + combos = value_list["combos"] + for combo in combos: + parameter_values = combo["parameter_values"] + if "suffix" in combo: + suffix = combo["suffix"] + else: + suffix = "" + for param_value in parameter_values: + if len(str(param_value)) > 0: + suffix += "_" + str(param_value) + suffix = suffix[1:] + param_values.append((param_names, suffix, parameter_values)) + + all_iterated_params.append(param_values) + + elif param_name not in exclude_params: + param_values = [] + for value in value_list: + if "RANGE" in value: + value_range = value["RANGE"] + suffix = value.get("SUFFIX", "") + if isinstance(value_range, list) and len(value_range) == 2: + for i in range(value_range[0], value_range[1] + 1): + curr_suffix = suffix + "_" + str(i) if suffix else str(i) + param_values.append((param_name, curr_suffix, i)) + else: + raise ValueError( + f"{value['RANGE']} is not a valid range. Must be in format [start, end] (inclusive)." + ) + + elif "VALUE" in value: + suffix = value.get("SUFFIX", value["VALUE"]) + if value["VALUE"] in ["int", "uint"]: + raise ValueError( + f"Use int32 or uint32 instead of {value['VALUE']}" + ) + param_values.append((param_name, suffix, value["VALUE"])) + + else: + raise KeyError( + "Parameter must be 'VALUE: string' or 'RANGE: [a, b]'" + ) + + all_iterated_params.append(param_values) + + return list(product(*all_iterated_params)) + + +def parse_template_spec(yaml_path) -> Dict[str, List[Dict[str, Any]]]: # noqa: C901 + """Parse a .yaml variant spec into {template_name: [expanded + per-variant param dicts]}. PyYAML with a dup-key-rejecting UniqueKeyLoader + (mirrors gen_vulkan_spv.py).""" + shader_template_params: Dict[str, List[Dict[str, Any]]] = {} + with open(yaml_path) as f: + contents = yaml.load(f, Loader=UniqueKeyLoader) + for template_name, params_dict in contents.items(): + if template_name in shader_template_params: + raise KeyError(f"{template_name} params file is defined twice") + + default_params = params_dict["parameter_names_with_default_values"] + params_names = set(default_params.keys()).union({"NAME"}) + + shader_template_params[template_name] = [] + + default_iterated_params = params_dict.get("generate_variant_forall", None) + + reserved_keys = { + "generate_variant_forall", + } + + for variant in params_dict["shader_variants"]: + default_iterated_params_names = set( + default_iterated_params.keys() + if default_iterated_params is not None + else {} + ) + variant_params_names = set(variant.keys()) + + invalid_keys = ( + variant_params_names + - default_iterated_params_names + - params_names + - reserved_keys + ) + if invalid_keys: + raise ValueError(f"unknown variant key(s): {sorted(invalid_keys)}") + + iterated_params = variant.get( + "generate_variant_forall", default_iterated_params + ) + + if iterated_params is not None: + variant_combinations = generate_variant_combinations( + iterated_params, variant_params_names + ) + + for combination in variant_combinations: + default_params_copy = copy.deepcopy(default_params) + for key in variant: + if key not in reserved_keys: + default_params_copy[key] = variant[key] + + variant_name = variant["NAME"] + + for setting in combination: + param_names = setting[0] + suffix = setting[1] + param_values = setting[2] + if isinstance(param_names, list): + for param_name, param_value in zip( + param_names, param_values + ): + default_params_copy[param_name] = param_value + else: + default_params_copy[param_names] = param_values + + if len(str(suffix)) > 0: + variant_name = f"{variant_name}_{suffix}" + + default_params_copy["NAME"] = variant_name + default_params_copy["VARIANT_NAME"] = variant["NAME"] + + shader_template_params[template_name].append(default_params_copy) + else: + default_params_copy = copy.deepcopy(default_params) + for key in variant: + if key not in reserved_keys: + default_params_copy[key] = variant[key] + + shader_template_params[template_name].append(default_params_copy) + + return shader_template_params + + def symbol_base(stem: str) -> str: """snake_case shader stem -> PascalCase symbol base (binary_add -> BinaryAdd).""" return "".join(part.capitalize() for part in stem.split("_")) @@ -88,12 +428,40 @@ def embedded_sha256(header_text: str) -> str: return m.group(1) if m else "" -def render_header(wgsl_path, wgsl_text: str) -> str: - """Render the full _wgsl.h text for a shader (shader embedded verbatim).""" +def _wg_size_const(base: str, axis: str, val: int) -> str: + """One WorkgroupSize constant; wrap to <=80 cols so CLANGFORMAT accepts it. + + Long shader names push the single-line form past the 80-col limit (clang-format + then breaks after '=' with a 4-space continuation indent); emit that wrapped + form up front so the generated header matches lintrunner's CLANGFORMAT. + """ + decl = f"inline constexpr uint32_t k{base}WorkgroupSize{axis} =" + if len(decl) + len(f" {val};") > 80: + return f"{decl}\n {val};\n" + return f"{decl} {val};\n" + + +def render_header( + name_or_path, wgsl_text: str, provenance_stem: Optional[str] = None +) -> str: + """Render the full _wgsl.h text for a shader (shader embedded unchanged). + + Two call forms: + - render_header(wgsl_path, wgsl_text): the plain, non-templated shaders -- + the symbol base and the `// @generated from` filename both derive from + Path(wgsl_path).stem. + - render_header(name, wgsl_text, provenance_stem): `name` is an expanded + variant name that drives the emitted symbols; `provenance_stem` is the + template stem cited in the `// @generated from` line. + """ + if provenance_stem is None: + name = Path(name_or_path).stem + provenance_stem = name + else: + name = name_or_path if ')"' in wgsl_text: raise ValueError('shader contains )" which would close the R"( literal') - stem = Path(wgsl_path).stem - base = symbol_base(stem) + base = symbol_base(name) x, y, z = parse_workgroup_size(wgsl_text) head = [ @@ -105,7 +473,7 @@ def render_header(wgsl_path, wgsl_text: str) -> str: "", "namespace executorch::backends::webgpu {", "", - f"// @generated from {stem}.wgsl - DO NOT EDIT.", + f"// @generated from {provenance_stem}.wgsl - DO NOT EDIT.", f"// wgsl-sha256: {wgsl_sha256(wgsl_text)}", f'inline constexpr const char* k{base}WGSL = R"(', ] @@ -115,9 +483,10 @@ def render_header(wgsl_path, wgsl_text: str) -> str: + wgsl_text + ')";' + "\n\n" - + f"inline constexpr uint32_t k{base}WorkgroupSizeX = {x};\n" - + f"inline constexpr uint32_t k{base}WorkgroupSizeY = {y};\n" - + f"inline constexpr uint32_t k{base}WorkgroupSizeZ = {z};\n\n" + + _wg_size_const(base, "X", x) + + _wg_size_const(base, "Y", y) + + _wg_size_const(base, "Z", z) + + "\n" + "} // namespace executorch::backends::webgpu\n" ) @@ -127,6 +496,36 @@ def discover(): return sorted((BACKEND_ROOT / "runtime/ops").glob("**/*.wgsl")) +def headers_for_shader(wgsl): + """Yield (header_path, rendered_text) pairs for one shader source. + + A shader is a template iff a sibling .yaml spec exists: each expanded + variant emits its own _wgsl.h (the provenance line cites the template + stem). Otherwise the shader is embedded unchanged into _wgsl.h. + """ + stem = wgsl.stem + text = wgsl.read_text() + spec_path = wgsl.with_name(stem + ".yaml") + if spec_path.exists(): + spec = parse_template_spec(spec_path) + if list(spec.keys()) != [stem]: + raise ValueError( + f"{spec_path.name}: top-level key must be '{stem}', got {list(spec.keys())}" + ) + for variant_params in spec[stem]: + name = variant_params["NAME"] + expanded = preprocess(text, {**WGSL_HELPERS, **variant_params}) + header = wgsl.with_name(name + "_wgsl.h") + yield header, render_header(name, expanded, stem) + else: + if "$if " in text or "${" in text: + raise ValueError( + f"shader uses $if/${{ templating but has no sibling {stem}.yaml spec" + ) + header = wgsl.with_name(stem + "_wgsl.h") + yield header, render_header(stem, text, stem) + + def _report_drift(missing, stale) -> None: """Print the --check report for missing/stale committed headers.""" if missing: @@ -152,20 +551,23 @@ def main(argv=None) -> int: missing = [] errors = [] for wgsl in discover(): - wgsl_text = wgsl.read_text() try: - want = render_header(wgsl, wgsl_text) - except ValueError as e: + rendered = list(headers_for_shader(wgsl)) + # A malformed spec raises yaml.YAMLError (incl. UniqueKeyLoader's + # ConstructorError) / ValueError / KeyError from parse_template_spec, and + # a malformed template raises AssertionError from preprocess; catch them + # all so a bad shader is a clean --check report, not a traceback. + except (ValueError, KeyError, AssertionError, yaml.YAMLError) as e: errors.append(f"{wgsl.relative_to(BACKEND_ROOT)}: {e}") continue - header = wgsl.with_name(wgsl.stem + "_wgsl.h") - # Full-content compare (not just the sha) catches generator-logic drift too. - if header.exists() and header.read_text() == want: - continue - if args.check: - (missing if not header.exists() else stale).append(header) - else: - header.write_text(want) + for header, want in rendered: + # Full-content compare (not just the sha) catches generator-logic drift too. + if header.exists() and header.read_text() == want: + continue + if args.check: + (missing if not header.exists() else stale).append(header) + else: + header.write_text(want) if errors: print("Cannot generate header (malformed shader):") diff --git a/backends/webgpu/test/native/test_scratch_buffer.cpp b/backends/webgpu/test/native/test_scratch_buffer.cpp index 98cf3648c6b..1a8f4fcd96f 100644 --- a/backends/webgpu/test/native/test_scratch_buffer.cpp +++ b/backends/webgpu/test/native/test_scratch_buffer.cpp @@ -6,7 +6,8 @@ * LICENSE file in the root directory of this source tree. */ -// White-box unit tests for WebGPUGraph::create_scratch_buffer. +// White-box unit tests for WebGPUGraph scratch buffers: +// create_scratch_buffer and the acquire_scratch/release_scratch reuse pool. #include #include @@ -19,6 +20,7 @@ #include #include #include +#include #include #include #include @@ -222,6 +224,59 @@ TEST(ScratchBuffer, Tier3Lifecycle) { } // each graph's dtor releases its 256 buffers here } +// Tier 4: reuse-pool semantics (acquire_scratch / release_scratch / +// ScopedScratch). The pool recycles single-op-lifetime scratch across ops so N +// layers reuse a small constant of buffers instead of N x. + +// A released slot is handed back on the next same-size acquire (the reuse win). +TEST(ScratchPool, ReuseAfterRelease) { + WebGPUGraph g; + g.set_device(g_device); + WGPUBuffer a = g.acquire_scratch(64 * sizeof(float)); + g.release_scratch(a); + WGPUBuffer b = g.acquire_scratch(64 * sizeof(float)); + EXPECT_EQ(a, b) << "released slot should be reused for a same-size request"; +} + +// A still-in_use slot is never handed to a co-live requester (RAW-safety). +TEST(ScratchPool, NoReuseWhileInUse) { + WebGPUGraph g; + g.set_device(g_device); + WGPUBuffer a = g.acquire_scratch(64 * sizeof(float)); + WGPUBuffer b = g.acquire_scratch(64 * sizeof(float)); // a not released + EXPECT_TRUE(a && b && a != b) << "co-live acquires must be distinct buffers"; +} + +// Best-fit 2x cap: a large free slot must not back a much smaller request, but +// a request it does fit (size in [n, 2n]) reuses it. +TEST(ScratchPool, BestFitSizeCap) { + WebGPUGraph g; + g.set_device(g_device); + WGPUBuffer big = g.acquire_scratch(1024 * sizeof(float)); + g.release_scratch(big); + // 1024*4 bytes is outside [4, 8], so the big slot is ineligible for 4 bytes. + WGPUBuffer tiny = g.acquire_scratch(4); + EXPECT_NE(big, tiny) + << "oversized slot must not back a tiny request (2x cap)"; + g.release_scratch(tiny); + WGPUBuffer same = g.acquire_scratch(1024 * sizeof(float)); + EXPECT_EQ(big, same) << "an in-range request should reuse the big slot"; +} + +// ScopedScratch releases its slot at scope exit, so the next acquire reuses it. +TEST(ScratchPool, ScopedScratchReleasesOnScopeExit) { + WebGPUGraph g; + g.set_device(g_device); + WGPUBuffer first = nullptr; + { + WebGPUGraph::ScopedScratch s(&g, g.acquire_scratch(64 * sizeof(float))); + first = s; // operator WGPUBuffer + EXPECT_NE(first, nullptr); + } // s releases the slot here + WGPUBuffer second = g.acquire_scratch(64 * sizeof(float)); + EXPECT_EQ(first, second) << "slot freed by ScopedScratch should be reused"; +} + int main(int argc, char** argv) { ::testing::InitGoogleTest(&argc, argv); diff --git a/backends/webgpu/test/ops/test_quantized_linear.py b/backends/webgpu/test/ops/test_quantized_linear.py index 5958ac6e2d5..72945c37d6d 100644 --- a/backends/webgpu/test/ops/test_quantized_linear.py +++ b/backends/webgpu/test/ops/test_quantized_linear.py @@ -59,7 +59,30 @@ class Q4gswConfig: # requires N % 8 == 0 (torchao pads N for the scale layout), so odd-N / N=1 are # not exportable -- bicol's has1 odd-N guard is defensive (mirrors coop4's # general-N robustness) and unreachable through this op. - # Prefill shapes routing to the shmem GEMM (K>=4096 or N>=2048); M=128. + # M>1 prefill: prefer the steel GEMM (K%16==0) on a >=256-invocation device + # (e.g. lvp); else shmem (K>=4096 or N>=2048) or register-tiled (SwiftShader + # caps at 128). Same fp64 golden regardless of which kernel runs. + Q4gswConfig("steel", 96, 2048, 256), # steel-isolating (K<4096, N<2048) + # Same shape as "steel"; the .pte is dtype-independent, so this fixture feeds + # the f16-multiply steel kernel (selected at runtime when the device reports + # shader-f16; goldened at a looser f16 tol in the native test). + Q4gswConfig("steel_f16", 96, 2048, 256), # f16-multiply steel (shader-f16) + # Partial M and N steel tiles under the f16 kernel; exercises f16 boundary + # masking (the exact-N "steel_f16" shape does not). N%8==0, steel-isolating. + Q4gswConfig("steel_f16_edge", 70, 1024, 136), # f16 partial-tile + # pwdq (packed-word dequant) backs the f16 steel path at group_size % BK(16) + # == 0 (bit-exact to steel_half; steel_f16 above runs it at gs=32). These lock + # the gs gate at group sizes those omit: gs=64 stays on pwdq; gs=8 (< BK) falls + # back to the per-nibble steel_half kernel (its hoisted-per-BK scale is invalid + # there). Same fp64 golden regardless of which kernel runs. + Q4gswConfig("pwdq_gs64", 96, 2048, 256, group_size=64), # pwdq, non-32 group + Q4gswConfig("pwdq_gs8", 96, 2048, 256, group_size=8), # steel_half fallback + # pwdqf16acc (f16-accumulate) runs when the enable_f16_accumulate_gemm runtime + # spec is set and gs % BK == 0 (perplexity-gated; see the kernel diff). Same + # .pte as the f32 configs -- only the accumulator dtype differs -- goldened at a + # looser f16-accumulate tol in the native test; deep-K stresses the worst case. + Q4gswConfig("pwdqf16acc", 96, 2048, 256), # f16-accumulate steel (runtime) + Q4gswConfig("pwdqf16acc_down", 128, 8192, 2048), # deep-K f16-accum worst case Q4gswConfig("gate_proj_pf", 128, 2048, 8192), # gate/up prefill (shmem via N) Q4gswConfig("down_proj_pf", 128, 8192, 2048), # down prefill (shmem via K) Q4gswConfig("shmem_edge", 130, 4096, 2056), # partial 32-tile bounds diff --git a/backends/webgpu/test/test_webgpu_native.cpp b/backends/webgpu/test/test_webgpu_native.cpp index 556eb0127b4..fbdfbd09076 100644 --- a/backends/webgpu/test/test_webgpu_native.cpp +++ b/backends/webgpu/test/test_webgpu_native.cpp @@ -11,6 +11,8 @@ #include #include #include +#include +#include #include @@ -229,6 +231,15 @@ bool sdpa_within_tol( int n, float* ma, float* mr) { + float atol = 1e-4f, rtol = 1e-3f; + // f16 KV (runtime opt-in) reads K/V at reduced precision; loosen the tol on a + // shader-f16 device to cover that rounding. Harmless for f32 KV (looser + // gate). + const WebGPUContext* kv_ctx = get_default_webgpu_context(); + if (kv_ctx != nullptr && kv_ctx->shader_f16_supported) { + atol = 2e-3f; + rtol = 1e-2f; + } float max_abs = 0.0f, max_rel = 0.0f; bool ok = true; for (int i = 0; i < n; i++) { @@ -236,7 +247,7 @@ bool sdpa_within_tol( const float re = ae / std::max(std::abs(golden[i]), 1e-6f); max_abs = std::max(max_abs, ae); max_rel = std::max(max_rel, re); - if (ae > 1e-4f && re > 1e-3f) { + if (ae > atol && re > rtol) { ok = false; } } @@ -275,6 +286,30 @@ const Q4gswConfig kQ4gswConfigs[] = { // scale over 64-256 K-groups). q4gsw requires N % 8 == 0, so odd-N is not // exportable; bicol's has1 odd-N guard is defensive (mirrors coop4 // general-N robustness). + // M>1: steel GEMM on a >=256-invocation device (K%16==0), else shmem/tiled. + {"steel", 96, 2048, 256, 1e-4f, 1e-3f, true, false}, // steel-isolating + // Same shape as "steel" run under the f16-multiply steel kernel; the f16 + // rounding floor (~2.3e-4, uniform in K -- not an accumulate bug) needs a + // looser abs gate than the strict f32 1e-4. Runs whenever the device + // negotiated shader-f16 (else the f32 steel kernel; the looser gate holds). + {"steel_f16", 96, 2048, 256, 2.3e-4f, 1e-3f, true, false}, + // Partial M and N steel tiles under the f16 kernel (f16 boundary masking). + {"steel_f16_edge", 70, 1024, 136, 2.3e-4f, 1e-3f, true, false}, + // pwdq (packed-word dequant) backs the f16 steel path at group_size % BK == + // 0 + // (bit-exact to steel_half; the steel_f16 configs above run it at gs=32). + // These lock the gs gate at group sizes those omit: gs=64 stays on pwdq; + // gs=8 (< BK=16) falls back to the per-nibble steel_half kernel. + {"pwdq_gs64", 96, 2048, 256, 2.3e-4f, 1e-3f, true, false}, + {"pwdq_gs8", 96, 2048, 256, 2.3e-4f, 1e-3f, true, false}, + // f16-ACCUMULATE steel (pwdqf16acc): lossy, so a wider gate than the + // f16-multiply steel_f16 (2.3e-4). f16 accumulation error grows with K, so + // the deep-K down shape (K=8192) gets the loosest tol. Perplexity is the + // primary quality gate (see the kernel diff); this catches gross bit/index + // bugs. gs=32 (% BK == 0) selects pwdqf16acc; the sweep loads these rows + // with the enable_f16_accumulate_gemm runtime spec set. + {"pwdqf16acc", 96, 2048, 256, 2e-2f, 3e-2f, true, false}, + {"pwdqf16acc_down", 128, 8192, 2048, 5e-2f, 8e-2f, true, false}, {"gate_proj_pf", 128, 2048, 8192, 1e-4f, 1e-3f, true, false}, // shmem via N {"down_proj_pf", 128, 8192, 2048, 1e-3f, 1e-2f, true, false}, // shmem via K {"shmem_edge", 130, 4096, 2056, 1e-4f, 1e-3f, true, false}, // partial tiles @@ -538,7 +573,18 @@ void test_q4gsw_config( cfg.n); Module module(pte); - ASSERT_EQ(module.load_forward(), Error::Ok) << "could not load " << pte; + // pwdqf16acc rows exercise the lossy f16-accumulate kernel, a runtime opt-in + // (default off); enable it via the backend option keyed by the registered id. + if (std::string(cfg.name).rfind("pwdqf16acc", 0) == 0) { + BackendOptions<1> opts; + opts.set_option("enable_f16_accumulate_gemm", true); + LoadBackendOptionsMap map; + ASSERT_EQ(map.set_options("VulkanBackend", opts.view()), Error::Ok); + ASSERT_EQ(module.load_forward(nullptr, nullptr, &map), Error::Ok) + << "could not load " << pte; + } else { + ASSERT_EQ(module.load_forward(), Error::Ok) << "could not load " << pte; + } const int in_numel = cfg.m * cfg.k; const int out_numel = cfg.m * cfg.n; diff --git a/backends/webgpu/test/test_wgsl_codegen.py b/backends/webgpu/test/test_wgsl_codegen.py index 283279e4fb5..46d285aa60b 100644 --- a/backends/webgpu/test/test_wgsl_codegen.py +++ b/backends/webgpu/test/test_wgsl_codegen.py @@ -11,15 +11,65 @@ import hashlib import importlib.util +import re import tempfile import unittest from pathlib import Path +import yaml + _GEN = Path(__file__).resolve().parents[1] / "scripts" / "gen_wgsl_headers.py" _spec = importlib.util.spec_from_file_location("gen_wgsl_headers", _GEN) g = importlib.util.module_from_spec(_spec) _spec.loader.exec_module(g) +# gen_wgsl_headers.py and backends/vulkan/runtime/gen_vulkan_spv.py share the +# same $-block transpiler helpers + the UniqueKeyLoader; the test below keeps +# them in sync with that source of truth. Resolve the path relative to the repo +# root (both backends/vulkan and backends/webgpu exist in pytorch/executorch) +# and compare the bodies as TEXT -- gen_vulkan_spv.py is the Vulkan backend's +# script and is not imported here. +_REPO_ROOT = g.BACKEND_ROOT.parents[1] +_VULKAN_SPV = _REPO_ROOT / "backends" / "vulkan" / "runtime" / "gen_vulkan_spv.py" +_SHARED_TRANSPILER_NAMES = ( + "extract_leading_whitespace", + "escape", + "preprocess", + "UniqueKeyLoader", +) + + +def _function_source(text: str, name: str) -> str: + """Return a top-level function/class's source: the `def`/`class ` line + through the last line before the next column-0 construct (to next dedent). + + Two-phase so a multi-line signature -- whose closing `) -> str:` sits at + column 0 -- is not mistaken for the next top-level construct. + """ + lines = text.splitlines() + start = next( + ( + i + for i, ln in enumerate(lines) + if re.match(rf"^(?:def|class) {re.escape(name)}\b", ln) + ), + None, + ) + if start is None: + raise AssertionError(f"def/class {name} not found") + # Advance past the (possibly multi-line) signature to the line ending in ':'. + sig = start + while not lines[sig].rstrip().endswith(":"): + sig += 1 + # The body ends at the next non-blank column-0 line. + end = len(lines) + for k in range(sig + 1, len(lines)): + head = lines[k][:1] + if head != "" and not head.isspace(): + end = k + break + return "\n".join(lines[start:end]).rstrip() + class WgslCodegenTest(unittest.TestCase): def test_symbol_base(self) -> None: @@ -95,11 +145,13 @@ def test_committed_headers_match_generator(self) -> None: wgsls = g.discover() self.assertGreater(len(wgsls), 0, "no .wgsl shaders discovered") for wgsl in wgsls: - want = g.render_header(wgsl, wgsl.read_text()) - got = wgsl.with_name(wgsl.stem + "_wgsl.h").read_text() - self.assertEqual( - got, want, f"{wgsl.stem}_wgsl.h stale; run scripts/gen_wgsl_headers.py" - ) + # headers_for_shader handles both verbatim shaders and templates + # (a template emits one header per expanded variant). + for header, want in g.headers_for_shader(wgsl): + got = header.read_text() + self.assertEqual( + got, want, f"{header.name} stale; run scripts/gen_wgsl_headers.py" + ) def test_parse_workgroup_allows_space(self) -> None: # @workgroup_size (64) — the spec-legal spaced form must still parse. @@ -187,5 +239,235 @@ def test_render_header_3d_emits_xyz(self) -> None: self.assertIn("inline constexpr uint32_t kFooWorkgroupSizeZ = 2;", h) +class WgslTemplateEngineTest(unittest.TestCase): + """Coverage for the $-block template engine + DTYPE/VEC variant matrix.""" + + # --- transpiler helpers stay in sync with their source --- + + @unittest.skipUnless( + _VULKAN_SPV.exists(), f"source of truth not present at {_VULKAN_SPV}" + ) + def test_transpiler_helpers_stay_in_sync(self) -> None: + # The shared $-block transpiler helpers must stay character-identical to + # their source of truth so they cannot silently drift. Read both files as + # TEXT (the source of truth cannot be imported -- it top-level + # `import yaml`s). + src_text = _VULKAN_SPV.read_text() + gen_text = _GEN.read_text() + for name in _SHARED_TRANSPILER_NAMES: + self.assertEqual( + _function_source(src_text, name), + _function_source(gen_text, name), + f"{name} has drifted from its source of truth " + f"({_VULKAN_SPV}) -- re-sync the shared transpiler helpers", + ) + + # --- preprocess ------------------------------------------------------- + + def test_preprocess_if_else_selects_branch(self) -> None: + tmpl = 'fn main() {\n $if MODE == "a":\n let x = 1;\n $else:\n let x = 2;\n}\n' + self.assertEqual( + g.preprocess(tmpl, {"MODE": "a"}), "fn main() {\n let x = 1;\n}\n" + ) + self.assertEqual( + g.preprocess(tmpl, {"MODE": "b"}), "fn main() {\n let x = 2;\n}\n" + ) + + def test_preprocess_inline_substitution_uses_helper(self) -> None: + tmpl = "type: ${buffer_gvec_type(DTYPE, VEC)};\n" + out = g.preprocess(tmpl, {**g.WGSL_HELPERS, "DTYPE": "float", "VEC": 4}) + self.assertEqual(out, "type: vec4;\n") + + def test_preprocess_guarded_body_indent_matches_control_column(self) -> None: + # $if authored at column 2 with its body one 2-space level deeper -> the + # guarded output line lands at column 2 (the control-line's column). + tmpl = "fn main() {\n $if VEC == 4:\n let a = 1;\n $else:\n let b = 2;\n}\n" + self.assertEqual( + g.preprocess(tmpl, {"VEC": 4}), "fn main() {\n let a = 1;\n}\n" + ) + self.assertEqual( + g.preprocess(tmpl, {"VEC": 1}), "fn main() {\n let b = 2;\n}\n" + ) + + def test_preprocess_enable_f16_only_for_half(self) -> None: + # DD-009: `enable f16;` is a literal line behind `$if DTYPE == "half":`, + # NOT an inline ${} (which would print a stray blank line for float and + # break byte-identity of the fp32 base). + tmpl = '$if DTYPE == "half":\n enable f16;\nfn main() {}\n' + self.assertEqual( + g.preprocess(tmpl, {"DTYPE": "half"}), "enable f16;\nfn main() {}\n" + ) + self.assertEqual(g.preprocess(tmpl, {"DTYPE": "float"}), "fn main() {}\n") + + # --- generate_variant_combinations ----------------------------------- + + def test_generate_variant_combinations_product(self) -> None: + iterated = { + "DTYPE": [{"VALUE": "float"}, {"VALUE": "half", "SUFFIX": "half"}], + "VEC": [{"VALUE": 1, "SUFFIX": ""}, {"VALUE": 4, "SUFFIX": "vec4"}], + } + combos = g.generate_variant_combinations(iterated) + self.assertEqual(len(combos), 4) + flat = [tuple((s[0], s[1], s[2]) for s in combo) for combo in combos] + self.assertIn((("DTYPE", "float", "float"), ("VEC", "", 1)), flat) + self.assertIn((("DTYPE", "half", "half"), ("VEC", "vec4", 4)), flat) + + def test_generate_variant_combinations_suffix_empty_suppresses(self) -> None: + combos = g.generate_variant_combinations({"VEC": [{"VALUE": 1, "SUFFIX": ""}]}) + self.assertEqual(combos, [(("VEC", "", 1),)]) + + def test_generate_variant_combinations_suffix_defaults_to_value(self) -> None: + # SUFFIX absent -> the suffix defaults to the VALUE (stringified in names). + combos = g.generate_variant_combinations({"VEC": [{"VALUE": 4}]}) + self.assertEqual(len(combos), 1) + ((name, suffix, value),) = combos[0] + self.assertEqual(name, "VEC") + self.assertEqual(value, 4) + self.assertEqual(str(suffix), "4") + + def test_generate_variant_combinations_excludes_param(self) -> None: + # A param already fixed by the variant is excluded from the forall product. + combos = g.generate_variant_combinations( + {"VEC": [{"VALUE": 1}, {"VALUE": 4}]}, {"VEC"} + ) + self.assertEqual(combos, [()]) + + # --- parse_template_spec --------------------------------------------- + + def _write_spec(self, tmp: str, name: str, spec_obj) -> Path: + p = Path(tmp) / f"{name}.yaml" + p.write_text(yaml.safe_dump(spec_obj)) + return p + + def test_parse_template_spec_minimal(self) -> None: + spec_obj = { + "op": { + "parameter_names_with_default_values": {"DTYPE": "float", "VEC": 1}, + "generate_variant_forall": { + "VEC": [ + {"VALUE": 1, "SUFFIX": ""}, + {"VALUE": 4, "SUFFIX": "vec4"}, + ] + }, + "shader_variants": [{"NAME": "op"}], + } + } + with tempfile.TemporaryDirectory() as tmp: + parsed = g.parse_template_spec(self._write_spec(tmp, "op", spec_obj)) + self.assertEqual(list(parsed.keys()), ["op"]) + v1, v4 = parsed["op"] + self.assertEqual((v1["NAME"], v1["VEC"], v1["DTYPE"]), ("op", 1, "float")) + self.assertEqual(v1["VARIANT_NAME"], "op") + self.assertEqual((v4["NAME"], v4["VEC"], v4["DTYPE"]), ("op_vec4", 4, "float")) + self.assertEqual(v4["VARIANT_NAME"], "op") + + def test_parse_template_spec_default_suffix_str_value_in_name(self) -> None: + # A forall value with no SUFFIX contributes str(VALUE) to the variant NAME. + spec_obj = { + "op": { + "parameter_names_with_default_values": {"VEC": 1}, + "generate_variant_forall": {"VEC": [{"VALUE": 4}]}, + "shader_variants": [{"NAME": "op"}], + } + } + with tempfile.TemporaryDirectory() as tmp: + parsed = g.parse_template_spec(self._write_spec(tmp, "op", spec_obj)) + self.assertEqual(parsed["op"][0]["NAME"], "op_4") + + def test_parse_template_spec_duplicate_key_raises(self) -> None: + # UniqueKeyLoader rejects a repeated key anywhere in the spec (this flow + # mapping is valid YAML with a duplicate key). + dup = '{"op": {"NAME": 1, "NAME": 2}}' + with tempfile.TemporaryDirectory() as tmp: + p = Path(tmp) / "op.yaml" + p.write_text(dup) + with self.assertRaises(yaml.YAMLError): + g.parse_template_spec(p) + + def test_headers_for_shader_top_level_key_must_match_stem(self) -> None: + with tempfile.TemporaryDirectory() as tmp: + op_dir = Path(tmp) / "runtime/ops/op" + op_dir.mkdir(parents=True) + (op_dir / "op.wgsl").write_text("@workgroup_size(64)\nfn main(){}\n") + # top-level key "WRONG" != stem "op" -> must raise. + (op_dir / "op.yaml").write_text( + '{"WRONG": {"parameter_names_with_default_values": {},' + ' "shader_variants": [{"NAME": "op"}]}}' + ) + with self.assertRaises(ValueError): + list(g.headers_for_shader(op_dir / "op.wgsl")) + + def test_headers_for_shader_templating_without_sidecar_raises(self) -> None: + # A $if/${ shader with no sibling .yaml spec is a hard error. + with tempfile.TemporaryDirectory() as tmp: + op_dir = Path(tmp) / "runtime/ops/op" + op_dir.mkdir(parents=True) + (op_dir / "op.wgsl").write_text( + "$if VEC == 4:\n x\n@workgroup_size(64)\nfn main(){}\n" + ) + with self.assertRaises(ValueError): + list(g.headers_for_shader(op_dir / "op.wgsl")) + + # --- WGSL type-helpers ----------------------------------------------- + + def test_buffer_scalar_type(self) -> None: + self.assertEqual(g.buffer_scalar_type("half"), "f16") + self.assertEqual(g.buffer_scalar_type("float"), "f32") + + def test_buffer_gvec_type(self) -> None: + self.assertEqual(g.buffer_gvec_type("float", 1), "f32") + self.assertEqual(g.buffer_gvec_type("float", 4), "vec4") + self.assertEqual(g.buffer_gvec_type("half", 4), "vec4") + + def test_accum_scalar_type(self) -> None: + # The float family (incl. half) accumulates in f32. + self.assertEqual(g.accum_scalar_type("float"), "f32") + self.assertEqual(g.accum_scalar_type("half"), "f32") + + # --- byte-identity round-trip ---------------------------------------- + + def test_rms_norm_template_roundtrip_byte_identical(self) -> None: + # Expanding the committed rms_norm.wgsl template + embedding it must + # reproduce the committed headers exactly (the dedup proof point). + rms_dir = g.BACKEND_ROOT / "runtime/ops/rms_norm" + template = (rms_dir / "rms_norm.wgsl").read_text() + for name, vec, header_name in [ + ("rms_norm", 1, "rms_norm_wgsl.h"), + ("rms_norm_vec4", 4, "rms_norm_vec4_wgsl.h"), + ]: + expanded = g.preprocess( + template, {**g.WGSL_HELPERS, "DTYPE": "float", "VEC": vec} + ) + want = g.render_header(name, expanded, "rms_norm") + got = (rms_dir / header_name).read_text() + self.assertEqual( + got, want, f"{header_name} not reproduced from rms_norm.wgsl template" + ) + + def test_rms_norm_half_variant_is_type_correct(self) -> None: + # A DTYPE=half expansion must emit compilable WGSL: `enable f16;`, an f32 + # accumulator, loads widened to f32 for the reduction, and the store + # narrowed back to f16 -- f16 storage with f32 compute, no type mismatch. + template = (g.BACKEND_ROOT / "runtime/ops/rms_norm/rms_norm.wgsl").read_text() + cases = { + 1: ("array", "f32(v) * f32(v)", "= f16(f32(v) * rstd * f32(w));"), + 4: ( + "array>", + "dot(vec4(v), vec4(v))", + "= vec4(vec4(t_in[base4 + x4]) * rstd" + " * vec4(t_weight[x4]));", + ), + } + for vec, (buf, widened_accum, narrowed_store) in cases.items(): + out = g.preprocess( + template, {**g.WGSL_HELPERS, "DTYPE": "half", "VEC": vec} + ) + self.assertTrue(out.startswith("enable f16;\n")) + self.assertIn(buf, out) + self.assertIn("local_sq_sum: f32", out) # f32 accumulator for both dtypes + self.assertIn(widened_accum, out) + self.assertIn(narrowed_store, out) + + if __name__ == "__main__": unittest.main()