Skip to main content

runmat_runtime/builtins/math/ode/
ode23.rs

1//! MATLAB-compatible `ode23` builtin.
2
3use runmat_builtins::{
4    BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
5    BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
6};
7use runmat_macros::runtime_builtin;
8use runmat_value::Value;
9
10use crate::builtins::common::spec::{
11    BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
12    ReductionNaN, ResidencyPolicy, ShapeRequirements,
13};
14use crate::builtins::math::ode::common::{
15    build_ode_output, define_ode_integer_contract, ode_options_from_struct, parse_ode_input,
16    parse_options, prepare_ode_options, solve_ode, OdeMethod,
17};
18use crate::builtins::math::ode::type_resolvers::ode_solution_type;
19use crate::{build_runtime_error, BuiltinResult, RuntimeError};
20
21const NAME: &str = "ode23";
22
23define_ode_integer_contract!("ode23", "Ode23");
24
25const ODE23_OUTPUT_Y: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
26    name: "y",
27    ty: BuiltinParamType::NumericArray,
28    arity: BuiltinParamArity::Required,
29    default: None,
30    description: "Solution states evaluated over tspan.",
31}];
32
33const ODE23_OUTPUT_TY: [BuiltinParamDescriptor; 2] = [
34    BuiltinParamDescriptor {
35        name: "t",
36        ty: BuiltinParamType::NumericArray,
37        arity: BuiltinParamArity::Required,
38        default: None,
39        description: "Time points selected by solver.",
40    },
41    BuiltinParamDescriptor {
42        name: "y",
43        ty: BuiltinParamType::NumericArray,
44        arity: BuiltinParamArity::Required,
45        default: None,
46        description: "Solution states at each returned time point.",
47    },
48];
49
50const ODE23_INPUTS_CORE: [BuiltinParamDescriptor; 3] = [
51    BuiltinParamDescriptor {
52        name: "odefun",
53        ty: BuiltinParamType::Any,
54        arity: BuiltinParamArity::Required,
55        default: None,
56        description: "ODE right-hand-side callback f(t,y).",
57    },
58    BuiltinParamDescriptor {
59        name: "tspan",
60        ty: BuiltinParamType::Any,
61        arity: BuiltinParamArity::Required,
62        default: None,
63        description: "Time interval or monotonic time vector.",
64    },
65    BuiltinParamDescriptor {
66        name: "y0",
67        ty: BuiltinParamType::Any,
68        arity: BuiltinParamArity::Required,
69        default: None,
70        description: "Initial state vector/value.",
71    },
72];
73
74const ODE23_INPUTS_WITH_OPTIONS: [BuiltinParamDescriptor; 4] = [
75    BuiltinParamDescriptor {
76        name: "odefun",
77        ty: BuiltinParamType::Any,
78        arity: BuiltinParamArity::Required,
79        default: None,
80        description: "ODE right-hand-side callback f(t,y).",
81    },
82    BuiltinParamDescriptor {
83        name: "tspan",
84        ty: BuiltinParamType::Any,
85        arity: BuiltinParamArity::Required,
86        default: None,
87        description: "Time interval or monotonic time vector.",
88    },
89    BuiltinParamDescriptor {
90        name: "y0",
91        ty: BuiltinParamType::Any,
92        arity: BuiltinParamArity::Required,
93        default: None,
94        description: "Initial state vector/value.",
95    },
96    BuiltinParamDescriptor {
97        name: "options",
98        ty: BuiltinParamType::Any,
99        arity: BuiltinParamArity::Optional,
100        default: None,
101        description: "Optional struct with tolerances and step controls.",
102    },
103];
104
105const ODE23_SIGNATURES: [BuiltinSignatureDescriptor; 4] = [
106    BuiltinSignatureDescriptor {
107        label: "y = ode23(odefun, tspan, y0)",
108        inputs: &ODE23_INPUTS_CORE,
109        outputs: &ODE23_OUTPUT_Y,
110    },
111    BuiltinSignatureDescriptor {
112        label: "y = ode23(odefun, tspan, y0, options)",
113        inputs: &ODE23_INPUTS_WITH_OPTIONS,
114        outputs: &ODE23_OUTPUT_Y,
115    },
116    BuiltinSignatureDescriptor {
117        label: "[t, y] = ode23(odefun, tspan, y0)",
118        inputs: &ODE23_INPUTS_CORE,
119        outputs: &ODE23_OUTPUT_TY,
120    },
121    BuiltinSignatureDescriptor {
122        label: "[t, y] = ode23(odefun, tspan, y0, options)",
123        inputs: &ODE23_INPUTS_WITH_OPTIONS,
124        outputs: &ODE23_OUTPUT_TY,
125    },
126];
127
128const ODE23_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
129    code: "RM.ODE23.INVALID_ARGUMENT",
130    identifier: Some("RunMat:ode23:InvalidArgument"),
131    when: "Input argument count/options struct grammar is invalid.",
132    message: "ode23: invalid argument",
133};
134
135const ODE23_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
136    code: "RM.ODE23.INVALID_INPUT",
137    identifier: Some("RunMat:ode23:InvalidInput"),
138    when: "ODE input/state/callback semantics are invalid for integration.",
139    message: "ode23: invalid input",
140};
141
142const ODE23_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
143    code: "RM.ODE23.INTERNAL",
144    identifier: Some("RunMat:ode23:Internal"),
145    when: "Internal output materialization fails.",
146    message: "ode23: internal runtime failure",
147};
148
149const ODE23_ERRORS: [BuiltinErrorDescriptor; 3] = [
150    ODE23_ERROR_INVALID_ARGUMENT,
151    ODE23_ERROR_INVALID_INPUT,
152    ODE23_ERROR_INTERNAL,
153];
154
155pub const ODE23_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
156    signatures: &ODE23_SIGNATURES,
157    output_mode: BuiltinOutputMode::ByRequestedOutputCount,
158    completion_policy: BuiltinCompletionPolicy::Public,
159    errors: &ODE23_ERRORS,
160};
161
162fn ode23_error_with_detail(
163    error: &'static BuiltinErrorDescriptor,
164    detail: impl AsRef<str>,
165) -> RuntimeError {
166    let detail = detail.as_ref();
167    let message = if detail.starts_with("ode23:") {
168        detail.to_string()
169    } else {
170        format!("{}: {}", error.message, detail)
171    };
172    let mut builder = build_runtime_error(message).with_builtin(NAME);
173    if let Some(identifier) = error.identifier {
174        builder = builder.with_identifier(identifier);
175    }
176    builder.build()
177}
178
179fn ode23_map_error(err: RuntimeError, fallback: &'static BuiltinErrorDescriptor) -> RuntimeError {
180    if err.identifier().is_some() {
181        err
182    } else {
183        ode23_error_with_detail(fallback, err.message())
184    }
185}
186
187#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::math::ode::ode23")]
188pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
189    name: "ode23",
190    op_kind: GpuOpKind::Custom("ode-solve"),
191    supported_precisions: &[],
192    broadcast: BroadcastSemantics::None,
193    provider_hooks: &[],
194    constant_strategy: ConstantStrategy::InlineLiteral,
195    residency: ResidencyPolicy::GatherImmediately,
196    nan_mode: ReductionNaN::Include,
197    two_pass_threshold: None,
198    workgroup_size: None,
199    accepts_nan_mode: false,
200    notes: "Adaptive ODE integration runs on the host. RHS callbacks may call GPU-aware builtins.",
201};
202
203#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::math::ode::ode23")]
204pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
205    name: "ode23",
206    shape: ShapeRequirements::Any,
207    constant_strategy: ConstantStrategy::InlineLiteral,
208    elementwise: None,
209    reduction: None,
210    emits_nan: false,
211    notes: "ODE integration repeatedly invokes user callbacks and terminates fusion planning.",
212};
213
214#[runtime_builtin(
215    name = "ode23",
216    category = "math/ode",
217    summary = "Solve nonstiff ODE systems using adaptive Bogacki-Shampine 3(2) integration.",
218    keywords = "ode23,ode,nonstiff,bogacki-shampine,adaptive step",
219    accel = "sink",
220    type_resolver(ode_solution_type),
221    descriptor(crate::builtins::math::ode::ode23::ODE23_DESCRIPTOR),
222    extensions(crate::builtins::math::ode::ode23::EXTENSIONS),
223    integer_capabilities(crate::builtins::math::ode::ode23::INTEGER_CAPABILITIES),
224    builtin_path = "crate::builtins::math::ode::ode23"
225)]
226async fn ode23_builtin(
227    function: Value,
228    tspan: Value,
229    y0: Value,
230    rest: Vec<Value>,
231) -> BuiltinResult<Value> {
232    if rest.len() > 1 {
233        return Err(ode23_error_with_detail(
234            &ODE23_ERROR_INVALID_ARGUMENT,
235            "too many input arguments",
236        ));
237    }
238    let options = parse_options(NAME, rest.first())
239        .map_err(|err| ode23_map_error(err, &ODE23_ERROR_INVALID_ARGUMENT))?;
240    let options = prepare_ode_options(NAME, options, ODE_COMPATIBILITY_EXTENSIONS)
241        .await
242        .map_err(|err| ode23_map_error(err, &ODE23_ERROR_INVALID_ARGUMENT))?;
243    let opts = ode_options_from_struct(NAME, options.as_ref())
244        .map_err(|err| ode23_map_error(err, &ODE23_ERROR_INVALID_ARGUMENT))?;
245    let input = parse_ode_input(NAME, tspan, y0, ODE_COMPATIBILITY_EXTENSIONS)
246        .await
247        .map_err(|err| ode23_map_error(err, &ODE23_ERROR_INVALID_INPUT))?;
248    let result = solve_ode(NAME, OdeMethod::Ode23, &function, &input, &opts)
249        .await
250        .map_err(|err| ode23_map_error(err, &ODE23_ERROR_INVALID_INPUT))?;
251    build_ode_output(NAME, result).map_err(|err| ode23_map_error(err, &ODE23_ERROR_INTERNAL))
252}
253
254#[cfg(test)]
255mod tests {
256    use super::*;
257    use futures::executor::block_on;
258    use runmat_value::Tensor;
259    use std::sync::Arc;
260
261    #[test]
262    fn ode23_supports_two_output_form() {
263        let _resolver =
264            crate::user_functions::install_semantic_function_resolver(Some(Arc::new(|_name| {
265                Some(0)
266            })));
267        let _invoker = crate::user_functions::install_semantic_function_invoker(Some(Arc::new(
268            move |_function, args, _requested_outputs| {
269                let y = match &args[1] {
270                    Value::Num(n) => *n,
271                    other => panic!("expected scalar state, got {other:?}"),
272                };
273                Box::pin(async move { Ok(Value::Num(-y)) })
274            },
275        )));
276
277        let _out_guard = crate::output_count::push_output_count(Some(2));
278        let out = block_on(ode23_builtin(
279            Value::FunctionHandle("decay".into()),
280            Value::Tensor(Tensor::new(vec![0.0, 0.5, 1.0], vec![1, 3]).unwrap()),
281            Value::Num(1.0),
282            Vec::new(),
283        ))
284        .unwrap();
285
286        match out {
287            Value::OutputList(values) => {
288                assert_eq!(values.len(), 2);
289            }
290            other => panic!("unexpected output {other:?}"),
291        }
292    }
293
294    #[test]
295    fn ode23_accepts_semantic_function_handle_rhs() {
296        let _invoker = crate::user_functions::install_semantic_function_invoker(Some(Arc::new(
297            move |function, args, _requested_outputs| {
298                assert_eq!(function, 55);
299                let y = match &args[1] {
300                    Value::Num(n) => *n,
301                    other => panic!("expected scalar state, got {other:?}"),
302                };
303                Box::pin(async move { Ok(Value::Num(-y)) })
304            },
305        )));
306
307        let out = block_on(ode23_builtin(
308            Value::BoundFunctionHandle {
309                name: "ode_decay".to_string(),
310                function: 55,
311            },
312            Value::Tensor(Tensor::new(vec![0.0, 1.0], vec![1, 2]).unwrap()),
313            Value::Num(1.0),
314            Vec::new(),
315        ))
316        .unwrap();
317
318        match out {
319            Value::Tensor(t) => {
320                assert_eq!(t.cols(), 1);
321                let last = t.materialize_f64()[t.rows() - 1];
322                assert!(last.is_finite());
323                assert!(last > 0.0);
324                assert!(last < 1.0);
325            }
326            other => panic!("unexpected output {other:?}"),
327        }
328    }
329
330    #[test]
331    fn ode23_too_many_inputs_uses_stable_identifier() {
332        let err = block_on(ode23_builtin(
333            Value::FunctionHandle("decay".into()),
334            Value::Tensor(Tensor::new(vec![0.0, 1.0], vec![1, 2]).unwrap()),
335            Value::Num(1.0),
336            vec![Value::Num(1.0), Value::Num(2.0)],
337        ))
338        .expect_err("expected too many inputs error");
339        assert_eq!(err.identifier(), ODE23_ERROR_INVALID_ARGUMENT.identifier);
340    }
341
342    #[test]
343    fn ode23_descriptor_signatures_cover_surface() {
344        let labels: Vec<&str> = ODE23_DESCRIPTOR
345            .signatures
346            .iter()
347            .map(|signature| signature.label)
348            .collect();
349        assert_eq!(
350            labels,
351            vec![
352                "y = ode23(odefun, tspan, y0)",
353                "y = ode23(odefun, tspan, y0, options)",
354                "[t, y] = ode23(odefun, tspan, y0)",
355                "[t, y] = ode23(odefun, tspan, y0, options)",
356            ]
357        );
358    }
359
360    #[test]
361    fn ode23_descriptor_errors_have_stable_codes() {
362        let codes: Vec<&str> = ODE23_DESCRIPTOR
363            .errors
364            .iter()
365            .map(|error| error.code)
366            .collect();
367        assert_eq!(
368            codes,
369            vec![
370                "RM.ODE23.INVALID_ARGUMENT",
371                "RM.ODE23.INVALID_INPUT",
372                "RM.ODE23.INTERNAL",
373            ]
374        );
375    }
376}