1use 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}