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 /// The runtime name on the given device.
26 fn name(client: &ComputeClient<Self>) -> &'static str;
27
28 /// Stable loaded driver/runtime identity for persistent autotuning. Unknown backends return
29 /// None and use session-only caches; device ordinal or API major version alone is insufficient.
30 /// Adding a default preserves existing custom Runtime implementations.
31 fn autotune_driver_fingerprint(_client: &ComputeClient<Self>) -> Option<alloc::string::String> {
32 None
33 }
34
35 /// Return true if global input array lengths should be added to kernel info.
36 fn require_array_lengths() -> bool {
37 false
38 }
39
40 /// Returns the maximum ruda count on each dimension that can be launched.
41 fn max_ruda_count() -> (u32, u32, u32);
42
43 /// Whether a tensor with `shape` and `strides` can be read as is. If the result is false, the
44 /// tensor should be made contiguous before reading.
45 fn can_read_tensor(shape: &Shape, strides: &Strides) -> bool;
46
47 /// Returns the properties of the target hardware architecture.
48 fn target_properties() -> TargetProperties;
49
50 /// Returns all devices available under the provided type id.
51 fn enumerate_devices(
52 type_id: u16,
53 info: &<Self::Server as ComputeServer>::Info,
54 ) -> Vec<DeviceId>;
55 /// Returns all devices that can be handled by the runtime.
56 fn enumerate_all_devices(info: &<Self::Server as ComputeServer>::Info) -> Vec<DeviceId> {
57 Self::enumerate_devices(0, info)
58 }
59}