1use std::any::TypeId;
2use tenferro_tensor::backend::{
3 BackendSession, BackendSessionHost, ElementwiseFusionPlan, ElementwiseReadOp,
4 GroupedGemmConfig, SessionCachedDot, TensorAnalytic, TensorBuffer, TensorDeviceTransfer,
5 TensorDot, TensorElementwise, TensorFusion, TensorIndexing, TensorReduction, TensorStructural,
6};
7use tenferro_tensor::config::{
8 CompareDir, DotGeneralConfig, GatherConfig, PadConfig, ScatterConfig, SliceConfig,
9};
10use tenferro_tensor::{
11 with_session_entry_guard, DotGeneralAccumulation, Tensor, TensorRead, TensorValue, TensorWrite,
12};
13
14use super::{WebGpuBackend, WebGpuRuntime, WebGpuRuntimeIdentity};
15
16#[doc(hidden)]
18pub(super) struct WebGpuExecSessionMarker;
19
20#[doc(hidden)]
22#[derive(Debug)]
23pub struct WebGpuExecSession<'a> {
24 backend: &'a mut WebGpuBackend,
25}
26
27impl WebGpuExecSession<'_> {
28 #[doc(hidden)]
30 pub fn runtime(&self) -> &WebGpuRuntime {
31 self.backend.runtime()
32 }
33
34 #[doc(hidden)]
36 pub fn runtime_identity(&self) -> WebGpuRuntimeIdentity {
37 self.backend.runtime_identity()
38 }
39}
40
41#[doc(hidden)]
46pub fn with_webgpu_exec_session<B, R>(
47 session: &mut B,
48 f: impl for<'a> FnOnce(&'a mut WebGpuExecSession<'a>) -> R,
49) -> Option<R>
50where
51 B: BackendSession + ?Sized,
52{
53 if session.session_type_id() != std::any::TypeId::of::<WebGpuExecSessionMarker>() {
54 return None;
55 }
56 let data = unsafe { session.session_data_mut() };
57 Some(unsafe { f(&mut *(data.cast::<WebGpuExecSession<'static>>())) })
60}
61
62macro_rules! delegate {
63 ($trait:path {
64 $(fn $method:ident($($arg:ident: $arg_ty:ty),* $(,)?) -> $ret:ty;)*
65 }) => {
66 impl $trait for WebGpuExecSession<'_> {
67 $(
68 fn $method(&mut self, $($arg: $arg_ty),*) -> $ret {
69 self.backend.$method($($arg),*)
70 }
71 )*
72 }
73 };
74}
75
76delegate!(TensorElementwise {
77 fn elementwise_read_into(op: ElementwiseReadOp, inputs: &[TensorRead<'_>], out: TensorWrite<'_>) -> crate::Result<()>;
78 fn add(lhs: &Tensor, rhs: &Tensor) -> crate::Result<Tensor>;
79 fn sub(lhs: &Tensor, rhs: &Tensor) -> crate::Result<Tensor>;
80 fn mul(lhs: &Tensor, rhs: &Tensor) -> crate::Result<Tensor>;
81 fn neg(input: &Tensor) -> crate::Result<Tensor>;
82 fn conj(input: &Tensor) -> crate::Result<Tensor>;
83 fn div(lhs: &Tensor, rhs: &Tensor) -> crate::Result<Tensor>;
84 fn abs(input: &Tensor) -> crate::Result<Tensor>;
85 fn sign(input: &Tensor) -> crate::Result<Tensor>;
86 fn maximum(lhs: &Tensor, rhs: &Tensor) -> crate::Result<Tensor>;
87 fn minimum(lhs: &Tensor, rhs: &Tensor) -> crate::Result<Tensor>;
88 fn compare(lhs: &Tensor, rhs: &Tensor, dir: &CompareDir) -> crate::Result<Tensor>;
89 fn select(pred: &Tensor, on_true: &Tensor, on_false: &Tensor) -> crate::Result<Tensor>;
90 fn clamp(input: &Tensor, lower: &Tensor, upper: &Tensor) -> crate::Result<Tensor>;
91});
92
93delegate!(TensorAnalytic {
94 fn exp(input: &Tensor) -> crate::Result<Tensor>;
95 fn log(input: &Tensor) -> crate::Result<Tensor>;
96 fn sin(input: &Tensor) -> crate::Result<Tensor>;
97 fn cos(input: &Tensor) -> crate::Result<Tensor>;
98 fn tanh(input: &Tensor) -> crate::Result<Tensor>;
99 fn sqrt(input: &Tensor) -> crate::Result<Tensor>;
100 fn rsqrt(input: &Tensor) -> crate::Result<Tensor>;
101 fn pow(lhs: &Tensor, rhs: &Tensor) -> crate::Result<Tensor>;
102 fn expm1(input: &Tensor) -> crate::Result<Tensor>;
103 fn log1p(input: &Tensor) -> crate::Result<Tensor>;
104});
105
106delegate!(TensorStructural {
107 fn to_contiguous_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
108 fn copy_read_into(src: TensorRead<'_>, dst: TensorWrite<'_>) -> crate::Result<()>;
109 fn transpose(input: &Tensor, perm: &[usize]) -> crate::Result<Tensor>;
110 fn reshape(input: &Tensor, shape: &[usize]) -> crate::Result<Tensor>;
111 fn broadcast_in_dim(input: &Tensor, shape: &[usize], dims: &[usize]) -> crate::Result<Tensor>;
112 fn cast(input: &Tensor, to: tenferro_tensor::DType) -> crate::Result<Tensor>;
113 fn extract_diagonal(input: &Tensor, axis_a: usize, axis_b: usize) -> crate::Result<Tensor>;
114 fn embed_diagonal(input: &Tensor, axis_a: usize, axis_b: usize) -> crate::Result<Tensor>;
115 fn tril(input: &Tensor, k: i64) -> crate::Result<Tensor>;
116 fn triu(input: &Tensor, k: i64) -> crate::Result<Tensor>;
117});
118
119delegate!(TensorReduction {
120 fn reduce_sum(input: &Tensor, axes: &[usize]) -> crate::Result<Tensor>;
121 fn reduce_prod(input: &Tensor, axes: &[usize]) -> crate::Result<Tensor>;
122 fn reduce_max(input: &Tensor, axes: &[usize]) -> crate::Result<Tensor>;
123 fn reduce_min(input: &Tensor, axes: &[usize]) -> crate::Result<Tensor>;
124});
125
126delegate!(TensorDot {
127 fn dot_general(lhs: &Tensor, rhs: &Tensor, config: &DotGeneralConfig) -> crate::Result<Tensor>;
128 fn dot_general_with_conj(
129 lhs: &Tensor,
130 rhs: &Tensor,
131 config: &DotGeneralConfig,
132 lhs_conj: bool,
133 rhs_conj: bool,
134 ) -> crate::Result<Tensor>;
135});
136
137delegate!(TensorIndexing {
138 fn gather(
139 operand: &Tensor,
140 start_indices: &Tensor,
141 config: &GatherConfig,
142 ) -> crate::Result<Tensor>;
143 fn scatter(
144 operand: &Tensor,
145 scatter_indices: &Tensor,
146 updates: &Tensor,
147 config: &ScatterConfig,
148 ) -> crate::Result<Tensor>;
149 fn slice(input: &Tensor, config: &SliceConfig) -> crate::Result<Tensor>;
150 fn dynamic_slice(
151 input: &Tensor,
152 starts: &Tensor,
153 slice_sizes: &[usize],
154 ) -> crate::Result<Tensor>;
155 fn dynamic_update_slice(
156 operand: &Tensor,
157 update: &Tensor,
158 starts: &Tensor,
159 ) -> crate::Result<Tensor>;
160 fn pad(input: &Tensor, config: &PadConfig) -> crate::Result<Tensor>;
161 fn concatenate(inputs: &[&Tensor], axis: usize) -> crate::Result<Tensor>;
162 fn reverse(input: &Tensor, axes: &[usize]) -> crate::Result<Tensor>;
163});
164
165delegate!(TensorFusion {
166 fn execute_elementwise_fusion(
167 inputs: &[&Tensor],
168 plan: &ElementwiseFusionPlan,
169 ) -> crate::Result<Option<Vec<Tensor>>>;
170 fn execute_broadcast_multiply(
171 lhs: TensorRead<'_>,
172 lhs_shape: &[usize],
173 lhs_dims: &[usize],
174 rhs: TensorRead<'_>,
175 rhs_shape: &[usize],
176 rhs_dims: &[usize],
177 ) -> crate::Result<Option<Tensor>>;
178 fn execute_broadcast_multiply_value(
179 lhs: TensorRead<'_>,
180 lhs_shape: &[usize],
181 lhs_dims: &[usize],
182 rhs: TensorRead<'_>,
183 rhs_shape: &[usize],
184 rhs_dims: &[usize],
185 ) -> crate::Result<Option<TensorValue>>;
186});
187
188delegate!(TensorBuffer {
189 fn reclaim_buffer(tensor: Tensor) -> ();
190});
191
192delegate!(TensorDeviceTransfer {
193 fn download_to_host(tensor: TensorRead<'_>) -> crate::Result<Tensor>;
194 fn upload_host_tensor(tensor: TensorRead<'_>) -> crate::Result<Tensor>;
195});
196
197macro_rules! delegate_cached {
198 ($(fn $method:ident($($arg:ident: $arg_ty:ty),* $(,)?) -> $ret:ty;)*) => {
199 impl SessionCachedDot for WebGpuExecSession<'_> {
200 $(
201 fn $method(&mut self, $($arg: $arg_ty),*) -> $ret {
202 <WebGpuBackend as SessionCachedDot>::$method(self.backend, $($arg),*)
203 }
204 )*
205 }
206 };
207}
208
209delegate_cached! {
210 fn dot_general_cached(
211 cache_slot: Option<usize>,
212 lhs: &Tensor,
213 rhs: &Tensor,
214 config: &DotGeneralConfig,
215 ) -> crate::Result<Tensor>;
216 fn dot_general_read_cached(
217 cache_slot: Option<usize>,
218 lhs: TensorRead<'_>,
219 rhs: TensorRead<'_>,
220 config: &DotGeneralConfig,
221 ) -> crate::Result<Tensor>;
222 fn dot_general_with_conj_cached(
223 cache_slot: Option<usize>,
224 lhs: &Tensor,
225 rhs: &Tensor,
226 config: &DotGeneralConfig,
227 lhs_conj: bool,
228 rhs_conj: bool,
229 ) -> crate::Result<Tensor>;
230 fn dot_general_with_conj_read_cached(
231 cache_slot: Option<usize>,
232 lhs: TensorRead<'_>,
233 rhs: TensorRead<'_>,
234 config: &DotGeneralConfig,
235 lhs_conj: bool,
236 rhs_conj: bool,
237 ) -> crate::Result<Tensor>;
238 fn dot_general_read_into_accum_cached(
239 cache_slot: Option<usize>,
240 lhs: TensorRead<'_>,
241 rhs: TensorRead<'_>,
242 config: &DotGeneralConfig,
243 accumulation: DotGeneralAccumulation,
244 out: TensorWrite<'_>,
245 ) -> crate::Result<()>;
246 fn grouped_gemm_cached(
247 cache_slot: Option<usize>,
248 lhs: TensorRead<'_>,
249 rhs: TensorRead<'_>,
250 config: &GroupedGemmConfig<'_>,
251 out: TensorWrite<'_>,
252 ) -> crate::Result<()>;
253}
254
255impl BackendSession for WebGpuExecSession<'_> {
256 fn session_type_id(&self) -> TypeId {
257 TypeId::of::<WebGpuExecSessionMarker>()
258 }
259
260 unsafe fn session_data_mut(&mut self) -> *mut () {
261 self as *mut Self as *mut ()
262 }
263}
264
265impl BackendSessionHost for WebGpuBackend {
266 fn with_backend_session<R: Send>(
267 &mut self,
268 f: impl FnOnce(&mut dyn BackendSession) -> R + Send,
269 ) -> R {
270 let mut session = WebGpuExecSession { backend: self };
271 with_session_entry_guard(|| f(&mut session))
274 }
275}