Skip to main content

ruda_runtime/runtime/
backend.rs

1use alloc::boxed::Box;
2use alloc::vec::Vec;
3use ruda_core::device::{Device, DeviceId};
4use ruda_core::ir::TargetProperties;
5use ruda_core::tensor::{Shape, Strides};
6
7use crate::runtime::{
8    client::ComputeClient,
9    compiler::{Compiler, RudaTask},
10    server::ComputeServer,
11};
12
13/// Runtime for the `Ruda`.
14pub trait Runtime: Sized + Send + Sync + 'static + core::fmt::Debug + Clone {
15    /// The compiler used to compile the inner representation into tokens.
16    type Compiler: Compiler;
17    /// The compute server used to run kernels and perform autotuning.
18    type Server: ComputeServer<Kernel = Box<dyn RudaTask<Self::Compiler>>>;
19    /// The device used to retrieve the compute client.
20    type Device: Device;
21
22    /// Retrieve the compute client from the runtime device.
23    fn client(device: &Self::Device) -> ComputeClient<Self>;
24
25    /// Whether this runtime supplies native LayerNorm forward and all first-order gradients.
26    fn has_native_layer_norm() -> bool { false }
27
28    /// Native LayerNorm, returning output, mean and reciprocal standard deviation.
29    fn layer_norm(
30        _client: &ComputeClient<Self>,
31        _input: super::normalization::TensorBuffer,
32        _weight: super::normalization::TensorBuffer,
33        _bias: Option<super::normalization::TensorBuffer>,
34        _epsilon: f64,
35    ) -> [super::normalization::TensorBuffer; 3] {
36        unimplemented!("runtime does not supply native LayerNorm")
37    }
38
39    /// Native LayerNorm gradients, returning input, weight and bias gradients.
40    fn layer_norm_backward(
41        _client: &ComputeClient<Self>,
42        _input: super::normalization::TensorBuffer,
43        _weight: super::normalization::TensorBuffer,
44        _grad: super::normalization::TensorBuffer,
45        _mean: super::normalization::TensorBuffer,
46        _rstd: super::normalization::TensorBuffer,
47    ) -> [super::normalization::TensorBuffer; 3] {
48        unimplemented!("runtime does not supply native LayerNorm backward")
49    }
50
51    /// The runtime name on the given device.
52    fn name(client: &ComputeClient<Self>) -> &'static str;
53
54    /// Stable loaded driver/runtime identity for persistent autotuning. Unknown backends return
55    /// None and use session-only caches; device ordinal or API major version alone is insufficient.
56    /// Adding a default preserves existing custom Runtime implementations.
57    fn autotune_driver_fingerprint(_client: &ComputeClient<Self>) -> Option<alloc::string::String> {
58        None
59    }
60
61    /// Return true if global input array lengths should be added to kernel info.
62    fn require_array_lengths() -> bool {
63        false
64    }
65
66    /// Returns the maximum ruda count on each dimension that can be launched.
67    fn max_ruda_count() -> (u32, u32, u32);
68
69    /// Whether a tensor with `shape` and `strides` can be read as is. If the result is false, the
70    /// tensor should be made contiguous before reading.
71    fn can_read_tensor(shape: &Shape, strides: &Strides) -> bool;
72
73    /// Returns the properties of the target hardware architecture.
74    fn target_properties() -> TargetProperties;
75
76    /// Returns all devices available under the provided type id.
77    fn enumerate_devices(
78        type_id: u16,
79        info: &<Self::Server as ComputeServer>::Info,
80    ) -> Vec<DeviceId>;
81    /// Returns all devices that can be handled by the runtime.
82    fn enumerate_all_devices(info: &<Self::Server as ComputeServer>::Info) -> Vec<DeviceId> {
83        Self::enumerate_devices(0, info)
84    }
85}