Skip to main content

tenferro_gpu/webgpu/
exec_session.rs

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/// Marker for the concrete erased WebGPU execution-session target.
17#[doc(hidden)]
18pub(super) struct WebGpuExecSessionMarker;
19
20/// Borrowed WebGPU execution capability.
21#[doc(hidden)]
22#[derive(Debug)]
23pub struct WebGpuExecSession<'a> {
24    backend: &'a mut WebGpuBackend,
25}
26
27impl WebGpuExecSession<'_> {
28    /// Borrow the provider runtime without exposing the owning backend.
29    #[doc(hidden)]
30    pub fn runtime(&self) -> &WebGpuRuntime {
31        self.backend.runtime()
32    }
33
34    /// Return the identity of the borrowed provider runtime.
35    #[doc(hidden)]
36    pub fn runtime_identity(&self) -> WebGpuRuntimeIdentity {
37        self.backend.runtime_identity()
38    }
39}
40
41/// Visit a WebGPU execution session through the erased backend-session surface.
42///
43/// The callback receives only the lifetime-bound session capability. The
44/// owning [`WebGpuBackend`] never crosses this boundary.
45#[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    // SAFETY: the exact marker check and BackendSession erased-pointer contract
58    // identify the value as WebGpuExecSession for this scoped visit.
59    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        // Nested entry is caught by the portable in-session guard in debug
272        // builds; the WebGPU runtime must never re-enter a session closure.
273        with_session_entry_guard(|| f(&mut session))
274    }
275}