Skip to main content

rlx_runtime/
session.rs

1// RLX — versatile ML compiler + runtime.
2// Copyright (C) 2026 Eugene Hauptmann, Nataliya Kosmyna.
3//
4// This program is free software: you can redistribute it and/or modify
5// it under the terms of the GNU General Public License as published by
6// the Free Software Foundation, version 3.
7//
8// This program is distributed in the hope that it will be useful,
9// but WITHOUT ANY WARRANTY; without even the implied warranty of
10// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
11// GNU General Public License for more details.
12//
13// You should have received a copy of the GNU General Public License
14// along with this program. If not, see <https://www.gnu.org/licenses/>.
15
16//! Session — the main entry point for compiling and executing graphs.
17
18use crate::backend::Backend;
19use crate::compiled::CompiledGraph;
20use crate::precision::Precision;
21use rlx_driver::Device;
22use rlx_ir::Graph;
23use rlx_ir::GraphModule;
24use rlx_ir::hir::HirModule;
25use rlx_opt::PrecisionPolicy;
26
27/// A session manages graph compilation and execution on a device.
28pub struct Session {
29    device: Device,
30    precision: Precision,
31    /// Optional per-op precision policy. If set, runs AutoMixedPrecision
32    /// rewrite before backend compile. Works identically across all modes
33    /// (AOT compile, trace/JIT, proc-macro AOT) — it's just a graph pass.
34    policy: Option<PrecisionPolicy>,
35}
36
37impl Session {
38    /// Create a session for the given device at default (F32) precision.
39    ///
40    /// # Panics
41    /// Panics if the device is not available (missing feature flag).
42    pub fn new(device: Device) -> Self {
43        Self::new_with_precision(device, Precision::F32)
44    }
45
46    /// Create a session targeting a specific numeric precision.
47    /// Backends fall back to F32 if the requested precision isn't supported.
48    pub fn new_with_precision(device: Device, precision: Precision) -> Self {
49        assert!(
50            crate::device_ext::is_available(device),
51            "device {} is not available — enable the `{}` Cargo feature",
52            device,
53            feature_name(device)
54        );
55        Self {
56            device,
57            precision,
58            policy: None,
59        }
60    }
61
62    /// Builder: set a per-op precision policy. Applied as a graph rewrite
63    /// before backend compile. Same mechanism works for AOT compile, JIT
64    /// tracing, and proc-macro AOT — it's a graph pass, not a runtime mode.
65    pub fn with_policy(mut self, policy: PrecisionPolicy) -> Self {
66        self.policy = Some(policy);
67        self
68    }
69
70    pub fn device(&self) -> Device {
71        self.device
72    }
73    pub fn precision(&self) -> Precision {
74        self.precision
75    }
76    pub fn policy(&self) -> Option<&PrecisionPolicy> {
77        self.policy.as_ref()
78    }
79
80    /// Compile a MIR graph through the fusion-first pipeline (`GraphModule` → LIR).
81    ///
82    /// Prefer [`Self::compile_hir`] or [`Self::compile_module`] for new code.
83    /// This entry wraps the graph as a MIR-stage [`GraphModule`].
84    ///
85    /// If the graph carries `Dim::Dynamic` dims it is compiled lazily: the first
86    /// `run` infers the concrete shape and specializes, so the same graph runs at
87    /// multiple input shapes on every backend. See [`crate::deferred`].
88    ///
89    /// On ANE with `RLX_COREML_NATIVE_FLEX=1`, dynamic graphs compile once with
90    /// CoreML `ShapeRange` and infer shapes at predict time instead.
91    pub fn compile(&self, graph: Graph) -> CompiledGraph {
92        if rlx_ir::dynamic::has_dynamic_dims(&graph) && !self.coreml_native_flex() {
93            return self.compile_deferred(graph, self.default_options());
94        }
95        let opts = self.default_options();
96        if opts.cache_param_invariant
97            && let Some(staged) = self.try_hoist_param_invariant(&graph, &opts)
98        {
99            return staged;
100        }
101        self.compile_module(GraphModule::from_graph(graph))
102            .expect("compile MIR graph through fusion pipeline")
103    }
104
105    /// If `graph` has a hoistable param-invariant closure, compile it as a
106    /// separate run-once `prepare` graph and attach it to the compiled main
107    /// graph. Returns `None` when there is nothing to hoist. Sub-compiles clear
108    /// the flag so this does not recurse.
109    fn try_hoist_param_invariant(
110        &self,
111        graph: &Graph,
112        options: &crate::CompileOptions,
113    ) -> Option<CompiledGraph> {
114        let split = rlx_compile::split_param_invariant(graph)?;
115        let mut sub = options.clone();
116        sub.cache_param_invariant = false;
117        let prepare = self.compile_with(split.prepare, &sub);
118        let main = self.compile_with(split.main, &sub);
119        Some(main.with_staging(
120            prepare,
121            split.boundary,
122            split.prepare_params,
123            split.main_params,
124        ))
125    }
126
127    /// Wrap a dynamic-shape graph so it specializes + recompiles per input shape.
128    fn compile_deferred(&self, graph: Graph, options: crate::CompileOptions) -> CompiledGraph {
129        let backend = self.create_backend();
130        let inner = Box::new(crate::deferred::DeferredExecutable::new(
131            graph,
132            backend,
133            options,
134            self.device,
135        ));
136        CompiledGraph::new(inner, self.device)
137    }
138
139    /// Explicit legacy alias — same as [`Self::compile`].
140    pub fn compile_graph(&self, graph: Graph) -> CompiledGraph {
141        self.compile(graph)
142    }
143
144    /// Compile with explicit options (full control over the pipeline).
145    /// Most callers use `compile()` and configure the session via
146    /// `new_with_precision` / `with_policy`. This escape hatch is for
147    /// callers that need finer control (e.g., disable DCE for debugging).
148    pub fn compile_with(&self, graph: Graph, options: &crate::CompileOptions) -> CompiledGraph {
149        // Opt-in native low-precision GEMM: rewrite every 2-D matmul into a
150        // dynamically-quantized `ScaledMatMul` in the requested format before
151        // the rest of the pipeline (must precede constant folding). Requires
152        // static matmul shapes.
153        let graph = match options.scaled_quant {
154            Some(cfg) => {
155                rlx_opt::rlx_compile::scaled_quant_insert::insert_scaled_matmul(graph, cfg)
156            }
157            None => graph,
158        };
159        if rlx_ir::dynamic::has_dynamic_dims(&graph) && !self.coreml_native_flex() {
160            return self.compile_deferred(graph, options.clone());
161        }
162        if options.cache_param_invariant
163            && let Some(staged) = self.try_hoist_param_invariant(&graph, options)
164        {
165            return staged;
166        }
167        self.compile_module_with(GraphModule::from_graph(graph), options)
168            .expect("compile MIR graph through fusion pipeline")
169    }
170
171    /// Compile a fusion-first HIR module through HIR → MIR → LIR.
172    pub fn compile_hir(&self, hir: HirModule) -> Result<CompiledGraph, rlx_ir::hir::LowerError> {
173        self.compile_hir_with(hir, &self.default_options())
174    }
175
176    /// Compile HIR with explicit compile options.
177    pub fn compile_hir_with(
178        &self,
179        hir: HirModule,
180        options: &crate::CompileOptions,
181    ) -> Result<CompiledGraph, rlx_ir::hir::LowerError> {
182        let backend = self.create_backend();
183        let executable = backend.compile_hir(hir, self.device, options)?;
184        Ok(CompiledGraph::new(executable, self.device))
185    }
186
187    /// Compile a [`GraphModule`] (HIR/MIR/LIR stage) through the pipeline.
188    pub fn compile_module(
189        &self,
190        module: GraphModule,
191    ) -> Result<CompiledGraph, rlx_ir::hir::LowerError> {
192        self.compile_module_with(module, &self.default_options())
193    }
194
195    /// Compile a [`GraphModule`] with explicit compile options.
196    pub fn compile_module_with(
197        &self,
198        module: GraphModule,
199        options: &crate::CompileOptions,
200    ) -> Result<CompiledGraph, rlx_ir::hir::LowerError> {
201        let backend = self.create_backend();
202        let executable = backend.compile_module(module, self.device, options)?;
203        Ok(CompiledGraph::new(executable, self.device))
204    }
205
206    fn default_options(&self) -> crate::CompileOptions {
207        let opts = crate::CompileOptions::new().precision(self.precision);
208        match &self.policy {
209            Some(p) => opts.policy(p.clone()),
210            None => opts,
211        }
212    }
213
214    /// ANE + `RLX_COREML_NATIVE_FLEX=1`: one CoreML model with `ShapeRange`.
215    fn coreml_native_flex(&self) -> bool {
216        #[cfg(all(
217            feature = "coreml",
218            target_vendor = "apple",
219            not(target_os = "watchos")
220        ))]
221        {
222            self.device == Device::Ane
223                && std::env::var("RLX_COREML_NATIVE_FLEX").ok().as_deref() == Some("1")
224        }
225        #[cfg(not(all(
226            feature = "coreml",
227            target_vendor = "apple",
228            not(target_os = "watchos")
229        )))]
230        {
231            false
232        }
233    }
234
235    fn create_backend(&self) -> Box<dyn Backend> {
236        // Single dispatch point: consult the registry. Backends register
237        // themselves (builtins via cfg-gated `register_builtin`; external
238        // crates via `register_backend`). No hardcoded match here.
239        crate::registry::backend_for(self.device).unwrap_or_else(|| {
240            panic!(
241                "no backend registered for device {} — enable feature `{}` \
242                 (or call `rlx_runtime::register_backend` for an external backend)",
243                self.device,
244                feature_name(self.device)
245            )
246        })
247    }
248}
249
250fn feature_name(device: Device) -> &'static str {
251    match device {
252        Device::Cpu => "cpu",
253        Device::Metal => "metal",
254        Device::Mlx => "mlx",
255        Device::Ane => "ane",
256        Device::Cuda => "cuda",
257        Device::Rocm => "rocm",
258        Device::OneApi => "oneapi",
259        Device::Tpu => "tpu",
260        Device::Hexagon => "qnn",
261        Device::Gpu => "gpu",
262        Device::Vulkan => "vulkan",
263        Device::OpenGl => "opengl",
264        Device::DirectX => "directx",
265        Device::WebGpu => "webgpu",
266    }
267}