diff --git a/backends/webgpu/runtime/ops/q8ta_conv2d_transposed/Q8taConv2dTransposed.cpp b/backends/webgpu/runtime/ops/q8ta_conv2d_transposed/Q8taConv2dTransposed.cpp new file mode 100644 index 00000000000..0990b49bb9d --- /dev/null +++ b/backends/webgpu/runtime/ops/q8ta_conv2d_transposed/Q8taConv2dTransposed.cpp @@ -0,0 +1,309 @@ +/* + * 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. + */ + +#include +#include +#include +#include + +#include + +#include +#include +#include + +namespace executorch::backends::webgpu { + +namespace { + +struct Q8taConvTParams { + uint32_t N; + uint32_t IC; + uint32_t H_in; + uint32_t W_in; + uint32_t OC; + uint32_t H_out; + uint32_t W_out; + uint32_t Kh; + uint32_t Kw; + uint32_t stride_h; + uint32_t stride_w; + uint32_t pad_h; + uint32_t pad_w; + uint32_t dil_h; + uint32_t dil_w; + uint32_t weight_row_stride; + int32_t input_zero_point; + int32_t output_zero_point; + float input_scale; + float inv_output_scale; + uint32_t has_bias; + uint32_t pad0; + uint32_t pad1; + uint32_t pad2; +}; +static_assert( + sizeof(Q8taConvTParams) == 96, + "Q8taConvTParams must match the WGSL Params struct (96 bytes)"); + +std::pair pair_or_throw( + const std::vector& v, + const char* msg) { + if (v.size() != 2) { + throw std::runtime_error(msg); + } + return {v.at(0), v.at(1)}; +} + +// int8 transposed conv (groups==1, dilation==1); gather form; mirrors Vulkan. +void q8ta_conv2d_transposed_impl( + WebGPUGraph& graph, + const std::vector& args) { + const int in_id = args.at(0); + const int weight_id = args.at(3); + const int scales_id = args.at(5); + const int bias_id = args.at(8); + const int out_id = args.at(args.size() - 1); + + if (graph.get_value_type(in_id) != WebGPUGraph::ValueType::Tensor || + graph.get_value_type(weight_id) != WebGPUGraph::ValueType::Tensor || + graph.get_value_type(scales_id) != WebGPUGraph::ValueType::Tensor || + graph.get_value_type(out_id) != WebGPUGraph::ValueType::Tensor) { + throw std::runtime_error( + "q8ta_conv2d_transposed: in/weight/scales/out not tensor"); + } + const bool has_bias = + graph.get_value_type(bias_id) == WebGPUGraph::ValueType::Tensor; + + const int act_id = args.at(args.size() - 2); + if (graph.get_value_type(act_id) != WebGPUGraph::ValueType::String || + graph.get_string(act_id) != "none") { + throw std::runtime_error( + "q8ta_conv2d_transposed: only activation='none' supported"); + } + + WGPUDevice device = graph.device(); + const auto& in_tensor = graph.get_tensor(in_id); + const auto& weight_tensor = graph.get_tensor(weight_id); + const auto& scales_tensor = graph.get_tensor(scales_id); + const auto& out_tensor = graph.get_tensor(out_id); + if (in_tensor.buffer == nullptr || weight_tensor.buffer == nullptr || + scales_tensor.buffer == nullptr || out_tensor.buffer == nullptr) { + throw std::runtime_error("q8ta_conv2d_transposed: null buffer binding"); + } + if (in_tensor.dims.size() != 4 || out_tensor.dims.size() != 4 || + weight_tensor.dims.size() != 2) { + throw std::runtime_error( + "q8ta_conv2d_transposed: in/out must be 4D, weight 2D"); + } + + const double input_scale = graph.get_double(args.at(1)); + const int input_zero_point = graph.get_int(args.at(2)); + const double output_scale = graph.get_double(args.at(6)); + const int output_zero_point = graph.get_int(args.at(7)); + + // Transposed schema inserts output_padding at 12, shifting dilation to 13. + const auto [kernel_h, kernel_w] = pair_or_throw( + graph.get_int_list(args.at(9)), "q8ta_conv2d_transposed: kernel_size"); + const auto [stride_h, stride_w] = pair_or_throw( + graph.get_int_list(args.at(10)), "q8ta_conv2d_transposed: stride"); + const auto [pad_h, pad_w] = pair_or_throw( + graph.get_int_list(args.at(11)), "q8ta_conv2d_transposed: padding"); + const auto [dil_h, dil_w] = pair_or_throw( + graph.get_int_list(args.at(13)), "q8ta_conv2d_transposed: dilation"); + const int groups = graph.get_int(args.at(14)); + if (groups != 1) { + throw std::runtime_error( + "q8ta_conv2d_transposed: only groups==1 supported"); + } + if (dil_h != 1 || dil_w != 1) { + throw std::runtime_error( + "q8ta_conv2d_transposed: only dilation==1 supported"); + } + if (stride_h < 1 || stride_w < 1) { + throw std::runtime_error("q8ta_conv2d_transposed: stride must be >= 1"); + } + + const uint64_t N = static_cast(in_tensor.dims.at(0)); + const uint64_t IC = static_cast(in_tensor.dims.at(1)); + const uint64_t H_in = static_cast(in_tensor.dims.at(2)); + const uint64_t W_in = static_cast(in_tensor.dims.at(3)); + const uint64_t OC = static_cast(weight_tensor.dims.at(0)); + const uint64_t weight_row_stride = + static_cast(weight_tensor.dims.at(1)); + const uint64_t Kh = static_cast(kernel_h); + const uint64_t Kw = static_cast(kernel_w); + const uint64_t H_out = static_cast(out_tensor.dims.at(2)); + const uint64_t W_out = static_cast(out_tensor.dims.at(3)); + + if (static_cast(out_tensor.dims.at(0)) != N || + static_cast(out_tensor.dims.at(1)) != OC) { + throw std::runtime_error( + "q8ta_conv2d_transposed: output must be [N, OC, H_out, W_out]"); + } + if (weight_row_stride < Kh * Kw * IC) { + throw std::runtime_error( + "q8ta_conv2d_transposed: weight row stride < Kh*Kw*IC"); + } + const uint64_t out_numel = N * OC * H_out * W_out; + if (out_numel == 0 || W_out % 4 != 0) { + throw std::runtime_error( + "q8ta_conv2d_transposed: W_out must be a nonzero multiple of 4"); + } + const uint64_t in_numel = N * IC * H_in * W_in; + const uint64_t weight_numel = OC * weight_row_stride; + if (out_numel > UINT32_MAX || in_numel > UINT32_MAX || + weight_numel > UINT32_MAX) { + throw std::runtime_error("q8ta_conv2d_transposed: numel exceeds u32"); + } + if (!in_tensor.is_int8 || in_tensor.nbytes != in_numel || in_numel % 4 != 0 || + !weight_tensor.is_int8 || weight_tensor.nbytes != weight_numel || + weight_numel % 4 != 0 || !out_tensor.is_int8 || + out_tensor.nbytes != out_numel) { + throw std::runtime_error( + "q8ta_conv2d_transposed: int8 in/weight/out size mismatch"); + } + // scales/bias fp32 [OC]; AOT pads OC to mult-4; shader reads [0,OC) only. + if (scales_tensor.nbytes < OC * sizeof(float)) { + throw std::runtime_error( + "q8ta_conv2d_transposed: weight_scales must be fp32 [OC]"); + } + if (has_bias) { + const auto& b = graph.get_tensor(bias_id); + if (b.buffer == nullptr || b.nbytes < OC * sizeof(float)) { + throw std::runtime_error( + "q8ta_conv2d_transposed: bias must be fp32 [OC]"); + } + } + + Q8taConvTParams params = {}; + params.N = static_cast(N); + params.IC = static_cast(IC); + params.H_in = static_cast(H_in); + params.W_in = static_cast(W_in); + params.OC = static_cast(OC); + params.H_out = static_cast(H_out); + params.W_out = static_cast(W_out); + params.Kh = static_cast(Kh); + params.Kw = static_cast(Kw); + params.stride_h = static_cast(stride_h); + params.stride_w = static_cast(stride_w); + params.pad_h = static_cast(pad_h); + params.pad_w = static_cast(pad_w); + params.dil_h = static_cast(dil_h); + params.dil_w = static_cast(dil_w); + params.weight_row_stride = static_cast(weight_row_stride); + params.input_zero_point = static_cast(input_zero_point); + params.output_zero_point = static_cast(output_zero_point); + params.input_scale = static_cast(input_scale); + // Reciprocal in double then cast, matching torch's f32(1.0 / f64(scale)). + params.inv_output_scale = static_cast(1.0 / output_scale); + params.has_bias = has_bias ? 1u : 0u; + + const uint32_t num_words = static_cast(out_numel / 4); + uint32_t wg_size = + utils::clamp_workgroup_size(device, kQ8taConv2dTransposedWorkgroupSizeX); + utils::WgCount workgroup_count = utils::compute_2d_workgroup_count( + device, num_words, wg_size, "q8ta_conv2d_transposed"); + + WGPUConstantEntry wg_size_constant = {}; + wg_size_constant.key = {"wg_size", WGPU_STRLEN}; + wg_size_constant.value = static_cast(wg_size); + + WGPUBuffer params_buf = + utils::make_uniform(device, ¶ms, sizeof(Q8taConvTParams)); + graph.add_uniform_buffer_bytes(sizeof(Q8taConvTParams)); + + WGPUShaderSourceWGSL wgsl_desc = {}; + wgsl_desc.chain.sType = WGPUSType_ShaderSourceWGSL; + wgsl_desc.code = {kQ8taConv2dTransposedWGSL, WGPU_STRLEN}; + WGPUShaderModuleDescriptor shader_desc = {}; + shader_desc.nextInChain = &wgsl_desc.chain; + WGPUShaderModule shader = wgpuDeviceCreateShaderModule(device, &shader_desc); + + WGPUBindGroupLayoutEntry entries[6] = {}; + for (int i = 0; i < 6; i++) { + entries[i].binding = static_cast(i); + entries[i].visibility = WGPUShaderStage_Compute; + entries[i].buffer.type = (i == 0) ? WGPUBufferBindingType_Storage + : (i == 5) ? WGPUBufferBindingType_Uniform + : WGPUBufferBindingType_ReadOnlyStorage; + } + WGPUBindGroupLayoutDescriptor bgl_desc = {}; + bgl_desc.entryCount = 6; + bgl_desc.entries = entries; + WGPUBindGroupLayout bgl = wgpuDeviceCreateBindGroupLayout(device, &bgl_desc); + + WGPUPipelineLayoutDescriptor pl_desc = {}; + pl_desc.bindGroupLayoutCount = 1; + pl_desc.bindGroupLayouts = &bgl; + WGPUPipelineLayout pipeline_layout = + wgpuDeviceCreatePipelineLayout(device, &pl_desc); + + WGPUComputePipelineDescriptor pipeline_desc = {}; + pipeline_desc.layout = pipeline_layout; + pipeline_desc.compute.module = shader; + pipeline_desc.compute.entryPoint = {"main", WGPU_STRLEN}; + pipeline_desc.compute.constantCount = 1; + pipeline_desc.compute.constants = &wg_size_constant; + WGPUComputePipeline pipeline = + wgpuDeviceCreateComputePipeline(device, &pipeline_desc); + + // No-bias: bind scales as an unread placeholder (has_bias gates the read). + WGPUBuffer bias_buf = + has_bias ? graph.get_tensor(bias_id).buffer : scales_tensor.buffer; + const uint64_t bias_size = + has_bias ? graph.get_tensor(bias_id).nbytes : scales_tensor.nbytes; + + WGPUBindGroupEntry bg[6] = {}; + bg[0].binding = 0; + bg[0].buffer = out_tensor.buffer; + bg[0].size = out_tensor.nbytes; + bg[1].binding = 1; + bg[1].buffer = in_tensor.buffer; + bg[1].size = in_tensor.nbytes; + bg[2].binding = 2; + bg[2].buffer = weight_tensor.buffer; + bg[2].size = weight_tensor.nbytes; + bg[3].binding = 3; + bg[3].buffer = scales_tensor.buffer; + bg[3].size = scales_tensor.nbytes; + bg[4].binding = 4; + bg[4].buffer = bias_buf; + bg[4].size = bias_size; + bg[5].binding = 5; + bg[5].buffer = params_buf; + bg[5].size = sizeof(Q8taConvTParams); + + WGPUBindGroupDescriptor bg_desc = {}; + bg_desc.layout = bgl; + bg_desc.entryCount = 6; + bg_desc.entries = bg; + WGPUBindGroup bind_group = wgpuDeviceCreateBindGroup(device, &bg_desc); + + graph.add_dispatch( + {pipeline, + bind_group, + workgroup_count.x, + "q8ta_conv2d_transposed", + workgroup_count.y}); + + wgpuShaderModuleRelease(shader); + wgpuBindGroupLayoutRelease(bgl); + wgpuPipelineLayoutRelease(pipeline_layout); + graph.own_uniform_buffer(params_buf); +} + +} // namespace + +WEBGPU_REGISTER_OPERATORS { + WEBGPU_REGISTER_OP( + et_vk.q8ta_conv2d_transposed.default, q8ta_conv2d_transposed_impl); +} + +} // namespace executorch::backends::webgpu diff --git a/backends/webgpu/runtime/ops/q8ta_conv2d_transposed/q8ta_conv2d_transposed.wgsl b/backends/webgpu/runtime/ops/q8ta_conv2d_transposed/q8ta_conv2d_transposed.wgsl new file mode 100644 index 00000000000..6c7a959441a --- /dev/null +++ b/backends/webgpu/runtime/ops/q8ta_conv2d_transposed/q8ta_conv2d_transposed.wgsl @@ -0,0 +1,108 @@ +@group(0) @binding(0) var t_out: array; +@group(0) @binding(1) var t_x: 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 { + N: u32, + IC: u32, + H_in: u32, + W_in: u32, + OC: u32, + H_out: u32, + W_out: u32, + Kh: u32, + Kw: u32, + stride_h: u32, + stride_w: u32, + pad_h: u32, + pad_w: u32, + dil_h: u32, + dil_w: u32, + weight_row_stride: u32, + input_zero_point: i32, + output_zero_point: i32, + input_scale: f32, + inv_output_scale: f32, + has_bias: u32, + pad0: u32, + pad1: u32, + pad2: u32, +} +@group(0) @binding(5) var params: Params; + +override wg_size: u32 = 64u; + +fn unpack_i8(bi: u32, word: u32) -> i32 { + return i32(((word >> ((bi & 3u) * 8u)) & 0xFFu) ^ 0x80u) - 128; +} + +@compute @workgroup_size(wg_size, 1, 1) +fn main( + @builtin(global_invocation_id) gid: vec3, + @builtin(num_workgroups) num_workgroups: vec3) { + // One thread per output word = 4 W-positions of fixed (n,oc,oh); W_out%4==0. + let widx = gid.x + gid.y * (num_workgroups.x * wg_size); + let words = (params.N * params.OC * params.H_out * params.W_out) / 4u; + if (widx >= words) { + return; + } + let flat0 = widx * 4u; + let ow0 = flat0 % params.W_out; + var r = flat0 / params.W_out; + let oh = r % params.H_out; + r = r / params.H_out; + let oc = r % params.OC; + let n = r / params.OC; + + var acc: array; + for (var j: u32 = 0u; j < 4u; j = j + 1u) { + acc[j] = 0; + } + // Transposed conv gather: ih=(oh+pad-kh*dil)/stride (must divide), full IC. + let w_row = oc * params.weight_row_stride; + for (var ic: u32 = 0u; ic < params.IC; ic = ic + 1u) { + for (var kh: u32 = 0u; kh < params.Kh; kh = kh + 1u) { + let ih_num = i32(oh) + i32(params.pad_h) - i32(kh) * i32(params.dil_h); + if (ih_num < 0 || (ih_num % i32(params.stride_h)) != 0) { + continue; + } + let ih = ih_num / i32(params.stride_h); + if (ih >= i32(params.H_in)) { + continue; + } + let in_row = ((n * params.IC + ic) * params.H_in + u32(ih)) * params.W_in; + for (var kw: u32 = 0u; kw < params.Kw; kw = kw + 1u) { + let wbi = w_row + (kh * params.Kw + kw) * params.IC + ic; + let wv = unpack_i8(wbi, t_weight[wbi >> 2u]); + for (var j: u32 = 0u; j < 4u; j = j + 1u) { + let iw_num = + i32(ow0 + j) + i32(params.pad_w) - i32(kw) * i32(params.dil_w); + if (iw_num < 0 || (iw_num % i32(params.stride_w)) != 0) { + continue; + } + let iw = iw_num / i32(params.stride_w); + if (iw >= i32(params.W_in)) { + continue; + } + let xbi = in_row + u32(iw); + acc[j] = acc[j] + + (unpack_i8(xbi, t_x[xbi >> 2u]) - params.input_zero_point) * wv; + } + } + } + } + + var packed: u32 = 0u; + for (var j: u32 = 0u; j < 4u; j = j + 1u) { + var v = f32(acc[j]) * params.input_scale * t_scales[oc]; + if (params.has_bias != 0u) { + v = v + t_bias[oc]; + } + var q = i32(round(v * params.inv_output_scale)) + params.output_zero_point; + q = clamp(q, -128, 127); + packed = packed | ((bitcast(q) & 0xFFu) << (j * 8u)); + } + t_out[widx] = packed; +} diff --git a/backends/webgpu/runtime/ops/q8ta_conv2d_transposed/q8ta_conv2d_transposed_wgsl.h b/backends/webgpu/runtime/ops/q8ta_conv2d_transposed/q8ta_conv2d_transposed_wgsl.h new file mode 100644 index 00000000000..79d463f9ddd --- /dev/null +++ b/backends/webgpu/runtime/ops/q8ta_conv2d_transposed/q8ta_conv2d_transposed_wgsl.h @@ -0,0 +1,132 @@ +/* + * 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 q8ta_conv2d_transposed.wgsl - DO NOT EDIT. +// wgsl-sha256: 426daa41bd5d7202d9260462c0ef4bb4c1737f0708f13331b22c4ba1f7885c1b +inline constexpr const char* kQ8taConv2dTransposedWGSL = R"( +@group(0) @binding(0) var t_out: array; +@group(0) @binding(1) var t_x: 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 { + N: u32, + IC: u32, + H_in: u32, + W_in: u32, + OC: u32, + H_out: u32, + W_out: u32, + Kh: u32, + Kw: u32, + stride_h: u32, + stride_w: u32, + pad_h: u32, + pad_w: u32, + dil_h: u32, + dil_w: u32, + weight_row_stride: u32, + input_zero_point: i32, + output_zero_point: i32, + input_scale: f32, + inv_output_scale: f32, + has_bias: u32, + pad0: u32, + pad1: u32, + pad2: u32, +} +@group(0) @binding(5) var params: Params; + +override wg_size: u32 = 64u; + +fn unpack_i8(bi: u32, word: u32) -> i32 { + return i32(((word >> ((bi & 3u) * 8u)) & 0xFFu) ^ 0x80u) - 128; +} + +@compute @workgroup_size(wg_size, 1, 1) +fn main( + @builtin(global_invocation_id) gid: vec3, + @builtin(num_workgroups) num_workgroups: vec3) { + // One thread per output word = 4 W-positions of fixed (n,oc,oh); W_out%4==0. + let widx = gid.x + gid.y * (num_workgroups.x * wg_size); + let words = (params.N * params.OC * params.H_out * params.W_out) / 4u; + if (widx >= words) { + return; + } + let flat0 = widx * 4u; + let ow0 = flat0 % params.W_out; + var r = flat0 / params.W_out; + let oh = r % params.H_out; + r = r / params.H_out; + let oc = r % params.OC; + let n = r / params.OC; + + var acc: array; + for (var j: u32 = 0u; j < 4u; j = j + 1u) { + acc[j] = 0; + } + // Transposed conv gather: ih=(oh+pad-kh*dil)/stride (must divide), full IC. + let w_row = oc * params.weight_row_stride; + for (var ic: u32 = 0u; ic < params.IC; ic = ic + 1u) { + for (var kh: u32 = 0u; kh < params.Kh; kh = kh + 1u) { + let ih_num = i32(oh) + i32(params.pad_h) - i32(kh) * i32(params.dil_h); + if (ih_num < 0 || (ih_num % i32(params.stride_h)) != 0) { + continue; + } + let ih = ih_num / i32(params.stride_h); + if (ih >= i32(params.H_in)) { + continue; + } + let in_row = ((n * params.IC + ic) * params.H_in + u32(ih)) * params.W_in; + for (var kw: u32 = 0u; kw < params.Kw; kw = kw + 1u) { + let wbi = w_row + (kh * params.Kw + kw) * params.IC + ic; + let wv = unpack_i8(wbi, t_weight[wbi >> 2u]); + for (var j: u32 = 0u; j < 4u; j = j + 1u) { + let iw_num = + i32(ow0 + j) + i32(params.pad_w) - i32(kw) * i32(params.dil_w); + if (iw_num < 0 || (iw_num % i32(params.stride_w)) != 0) { + continue; + } + let iw = iw_num / i32(params.stride_w); + if (iw >= i32(params.W_in)) { + continue; + } + let xbi = in_row + u32(iw); + acc[j] = acc[j] + + (unpack_i8(xbi, t_x[xbi >> 2u]) - params.input_zero_point) * wv; + } + } + } + } + + var packed: u32 = 0u; + for (var j: u32 = 0u; j < 4u; j = j + 1u) { + var v = f32(acc[j]) * params.input_scale * t_scales[oc]; + if (params.has_bias != 0u) { + v = v + t_bias[oc]; + } + var q = i32(round(v * params.inv_output_scale)) + params.output_zero_point; + q = clamp(q, -128, 127); + packed = packed | ((bitcast(q) & 0xFFu) << (j * 8u)); + } + t_out[widx] = packed; +} +)"; + +inline constexpr uint32_t kQ8taConv2dTransposedWorkgroupSizeX = 64; +inline constexpr uint32_t kQ8taConv2dTransposedWorkgroupSizeY = 1; +inline constexpr uint32_t kQ8taConv2dTransposedWorkgroupSizeZ = 1; + +} // namespace executorch::backends::webgpu