Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
98 changes: 98 additions & 0 deletions backends/vulkan/custom_ops_lib.py
Original file line number Diff line number Diff line change
Expand Up @@ -900,6 +900,44 @@ def apply_rotary_emb_hf_meta(
lib.impl(name, apply_rotary_emb_hf_meta, "Meta")
apply_rotary_emb_hf_op = getattr(getattr(torch.ops, namespace), name)

################################
## apply_rotary_emb_hf_single ##
################################


def apply_rotary_emb_hf_single_impl(
x: torch.Tensor,
freqs_cos: torch.Tensor,
freqs_sin: torch.Tensor,
start_pos: int,
):
seq_len = x.shape[1]
freqs_cos = freqs_cos[start_pos : start_pos + seq_len]
freqs_sin = freqs_sin[start_pos : start_pos + seq_len]
pattern = vk_patterns.HfRotaryEmbeddingSinglePattern()
return pattern.forward(x, freqs_cos, freqs_sin)


def apply_rotary_emb_hf_single_meta(
x: torch.Tensor,
freqs_cos: torch.Tensor,
freqs_sin: torch.Tensor,
start_pos: int,
):
output_dtype = torch.promote_types(
torch.promote_types(x.dtype, freqs_cos.dtype), freqs_sin.dtype
)
return torch.empty_like(x, dtype=output_dtype)


name = "apply_rotary_emb_hf_single"
lib.define(
f"{name}(Tensor x, Tensor freqs_cos, Tensor freqs_sin, SymInt start_pos) -> Tensor"
)
lib.impl(name, apply_rotary_emb_hf_single_impl, "CompositeExplicitAutograd")
lib.impl(name, apply_rotary_emb_hf_single_meta, "Meta")
apply_rotary_emb_hf_single_op = getattr(getattr(torch.ops, namespace), name)

##################################
## apply_rotary_emb_interleaved ##
##################################
Expand Down Expand Up @@ -1149,6 +1187,66 @@ def sdpa_impl(
lib.impl(name, sdpa_impl, "CompositeExplicitAutograd")
sdpa_op = getattr(getattr(torch.ops, namespace), name)

#################
## gemma4_sdpa ##
#################


def gemma4_sdpa_impl(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
start_pos: int,
attn_mask: torch.Tensor,
dropout_p: float,
is_causal: bool,
scale: float,
) -> torch.Tensor:
del start_pos
if dropout_p != 0.0 or is_causal or scale != 1.0:
raise ValueError("gemma4_sdpa requires dropout=0, causal=false, scale=1")
if query.dim() != 4 or key.dim() != 4 or value.dim() != 4:
raise ValueError("gemma4_sdpa requires BSHD query, key, and value")
if key.shape != value.shape or query.shape[0] != key.shape[0]:
raise ValueError("gemma4_sdpa query, key, and value shapes do not match")
if query.shape[-1] != key.shape[-1] or query.shape[2] % key.shape[2] != 0:
raise ValueError("gemma4_sdpa requires grouped-query compatible heads")
if attn_mask.dim() != 2 or tuple(attn_mask.shape) != (
query.shape[1],
key.shape[1],
):
raise ValueError("gemma4_sdpa requires a rank-2 [S_q, S_kv] mask")

group_size = query.shape[2] // key.shape[2]
query_bhsd = query.transpose(1, 2)
key_bhsd = key.transpose(1, 2).repeat_interleave(group_size, dim=1)
value_bhsd = value.transpose(1, 2).repeat_interleave(group_size, dim=1)
scores = torch.matmul(query_bhsd, key_bhsd.transpose(-2, -1))
scores = scores + attn_mask
return torch.matmul(torch.softmax(scores, dim=-1), value_bhsd).transpose(1, 2)


def gemma4_sdpa_meta(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
start_pos: int,
attn_mask: torch.Tensor,
dropout_p: float,
is_causal: bool,
scale: float,
) -> torch.Tensor:
return torch.empty_like(query)


name = "gemma4_sdpa"
lib.define(
f"{name}(Tensor query, Tensor key, Tensor value, SymInt start_pos, Tensor attn_mask, float dropout_p, bool is_causal, float scale) -> Tensor"
)
lib.impl(name, gemma4_sdpa_impl, "CompositeExplicitAutograd")
lib.impl(name, gemma4_sdpa_meta, "Meta")
gemma4_sdpa_op = getattr(getattr(torch.ops, namespace), name)

################
## rms_norm ##
################
Expand Down
28 changes: 28 additions & 0 deletions backends/webgpu/runtime/WebGPUDispatchMath.h
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,34 @@ constexpr bool should_record_sdpa_dual_route(
return fd_eligible && (has_dynamic_sequence || has_dynamic_position);
}

constexpr uint32_t kCqpQdqFusedInvocations = 256u;
constexpr uint32_t kCqpQdqFusedStorageBytes =
2u * kCqpQdqFusedInvocations * sizeof(float);

constexpr bool is_cqp_qdq_fusion_eligible(
uint32_t rows,
uint32_t row_width,
uint64_t numel,
int64_t quant_min,
int64_t quant_max,
bool asymmetric,
bool per_row_block,
bool keepdim,
uint32_t max_invocations,
uint32_t max_workgroup_size_x,
uint32_t max_workgroup_storage_bytes) {
return rows > 0u && row_width > 0u &&
numel == static_cast<uint64_t>(rows) * row_width && asymmetric &&
per_row_block && !keepdim && quant_min == -128 && quant_max == 127 &&
max_invocations >= kCqpQdqFusedInvocations &&
max_workgroup_size_x >= kCqpQdqFusedInvocations &&
max_workgroup_storage_bytes >= kCqpQdqFusedStorageBytes;
}

constexpr uint32_t cqp_resize_workgroups(bool producer_elided, uint32_t grid) {
return producer_elided ? 0u : grid;
}

constexpr bool is_q4gsw_bk64_eligible(
uint32_t k,
uint32_t n,
Expand Down
43 changes: 43 additions & 0 deletions backends/webgpu/runtime/WebGPUGraph.cpp
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
/*
* Copyright (c) Meta Platforms, Inc. and affiliates.
* All rights reserved.
Expand Down Expand Up @@ -1006,12 +1006,55 @@
}
}

void WebGPUGraph::offer_cqp_fusion_site(CqpFusionSite site) {
site.valid = true;
cqp_fusion_site_ = std::move(site);
}

WebGPUGraph::CqpFusionSite WebGPUGraph::claim_cqp_fusion_site(
int input_id,
int scales_id,
int zero_points_id,
uint32_t rows,
uint32_t row_width) {
const CqpFusionSite& site = cqp_fusion_site_;
const bool producer_is_choose_qparams =
site.dispatch_index < dispatches_.size() &&
dispatches_[site.dispatch_index].kernel_name == "choose_qparams_affine";
const bool matches = site.valid && site.input_id == input_id &&
site.scales_id == scales_id && site.zero_points_id == zero_points_id &&
site.rows == rows && site.row_width == row_width &&
site.input_buffer == get_tensor(input_id).buffer &&
site.scales_buffer == get_tensor(scales_id).buffer &&
site.zero_points_buffer == get_tensor(zero_points_id).buffer &&
site.dispatch_index + 1u == dispatches_.size() &&
producer_is_choose_qparams && site.producer_elided != nullptr;
if (!matches) {
return CqpFusionSite{};
}
CqpFusionSite claimed = site;
cqp_fusion_site_ = CqpFusionSite{};
return claimed;
}

void WebGPUGraph::build(
const void* flatbuffer_data,
const uint8_t* constant_data,
size_t constant_data_size,
const executorch::runtime::NamedDataMap* named_data_map,
WebGPUGraphConfig config) {
clear_cqp_fusion_site();
clear_rms_fusion_site();
clear_slice_chain();
struct ClearFusionSitesOnExit {
WebGPUGraph* graph;
~ClearFusionSitesOnExit() {
graph->clear_cqp_fusion_site();
graph->clear_rms_fusion_site();
graph->clear_slice_chain();
}
} clear_fusion_sites_on_exit{this};

if (!device_) {
auto* ctx = get_default_webgpu_context();
if (ctx) {
Expand Down
83 changes: 83 additions & 0 deletions backends/webgpu/runtime/WebGPUGraph.h
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@

#include <cstdint>
#include <functional>
#include <memory>
#include <stdexcept>
#include <string>
#include <type_traits>
Expand Down Expand Up @@ -147,6 +148,36 @@ struct WebGPUGraphConfig {

class WebGPUGraph {
public:
struct CqpFusionSite {
bool valid = false;
int input_id = -1;
int scales_id = -1;
int zero_points_id = -1;
uint32_t rows = 0u;
uint32_t row_width = 0u;
int64_t quant_min = 0;
int64_t quant_max = 0;
size_t dispatch_index = 0u;
WGPUBuffer input_buffer = nullptr;
WGPUBuffer scales_buffer = nullptr;
WGPUBuffer zero_points_buffer = nullptr;
std::shared_ptr<bool> producer_elided;
};

struct RmsFusionSite {
bool valid = false;
bool add_fused = false;
int in_id = -1;
int weight_id = -1;
int out_id = -1;
int resid_id = -1;
int addout_id = -1;
uint32_t num_rows = 0u;
uint32_t row_width = 0u;
size_t dispatch_index = 0u;
WGPUBuffer params_buffer = nullptr;
};

WebGPUGraph();
~WebGPUGraph();

Expand Down Expand Up @@ -410,6 +441,55 @@ class WebGPUGraph {
return dispatches_.size();
}

void offer_cqp_fusion_site(CqpFusionSite site);
CqpFusionSite claim_cqp_fusion_site(
int input_id,
int scales_id,
int zero_points_id,
uint32_t rows,
uint32_t row_width);
void clear_cqp_fusion_site() {
cqp_fusion_site_ = CqpFusionSite{};
}

void offer_rms_fusion_site(RmsFusionSite site) {
site.valid = true;
rms_fusion_site_ = std::move(site);
}
const RmsFusionSite& rms_fusion_site() const {
return rms_fusion_site_;
}
void clear_rms_fusion_site() {
rms_fusion_site_ = RmsFusionSite{};
}

// Dual-store slice merge: the preceding slice offers its dispatch so a
// following whole-extent copy can re-bind it to a second destination.
// Graph-instance state, so two graphs can never observe each other's
// dispatch indices or buffer handles.
struct SliceChain {
bool valid = false;
int out_id = -1;
size_t dispatch_idx = 0;
WGPUBuffer in_buffer = nullptr;
size_t in_nbytes = 0;
WGPUBuffer out_buffer = nullptr;
size_t out_nbytes = 0;
WGPUBuffer out_meta_buf = nullptr;
WGPUBuffer in_meta_buf = nullptr;
WGPUBuffer params_buf = nullptr;
};

void offer_slice_chain(SliceChain chain) {
slice_chain_ = chain;
}
const SliceChain& slice_chain() const {
return slice_chain_;
}
void clear_slice_chain() {
slice_chain_ = SliceChain{};
}

size_t register_dispatch_route_group(
const std::vector<utils::DispatchRange>& ranges) {
validate_dynamic_dispatch_route_ranges(ranges);
Expand Down Expand Up @@ -761,6 +841,9 @@ class WebGPUGraph {

std::vector<WebGPUDispatch> dispatches_;
utils::DispatchRouteRegistry dispatch_routes_;
CqpFusionSite cqp_fusion_site_;
RmsFusionSite rms_fusion_site_;
SliceChain slice_chain_;

// Prepack-routed constant sources (offset/named-key + size); the prepack node
// materializes these once. constant_data_/named_data_map_ point at the .pte
Expand Down
Loading
Loading