1use runmat_builtins::{
4 Access, BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
5 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
6 ClassDef, MethodDef, ObjectInstance, ResolveContext, StringArray, Tensor, Type, Value,
7};
8use std::cell::Cell;
9use std::collections::HashMap;
10
11use crate::{build_runtime_error, gather_if_needed_async, BuiltinResult, RuntimeError};
12
13pub(super) const MAX_COMBVEC_COLUMNS: usize = 1_000_000;
14pub(super) const MAX_PAD_ELEMENTS: usize = 10_000_000;
15
16const ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
17 code: "RM.DEEP_LEARNING.INVALID_INPUT",
18 identifier: Some("RunMat:deepLearning:InvalidInput"),
19 when:
20 "Inputs or name-value options do not match the supported Deep Learning compatibility forms.",
21 message: "deep learning builtin received invalid input",
22};
23
24const ERROR_UNSUPPORTED: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
25 code: "RM.DEEP_LEARNING.UNSUPPORTED",
26 identifier: Some("RunMat:deepLearning:Unsupported"),
27 when: "The requested operation requires training, autodiff, export, or UI infrastructure outside this compatibility slice.",
28 message: "deep learning operation is not supported in this slice",
29};
30
31const ERRORS: [BuiltinErrorDescriptor; 2] = [ERROR_INVALID_INPUT, ERROR_UNSUPPORTED];
32
33thread_local! {
34 static DLARRAY_CLASS_REGISTERED: Cell<bool> = const { Cell::new(false) };
35}
36
37const OUT_OBJECT: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
38 name: "obj",
39 ty: BuiltinParamType::Any,
40 arity: BuiltinParamArity::Required,
41 default: None,
42 description: "Deep Learning Toolbox compatibility object.",
43}];
44
45const OUT_ARRAY: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
46 name: "Y",
47 ty: BuiltinParamType::Any,
48 arity: BuiltinParamArity::Required,
49 default: None,
50 description: "Numeric or cell array output.",
51}];
52
53const IN_REST: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
54 name: "args",
55 ty: BuiltinParamType::Any,
56 arity: BuiltinParamArity::Variadic,
57 default: None,
58 description: "Builtin-specific positional and name-value arguments.",
59}];
60
61const OBJECT_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
62 label: "obj = deepLearningBuiltin(args...)",
63 inputs: &IN_REST,
64 outputs: &OUT_OBJECT,
65}];
66
67const ARRAY_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
68 label: "Y = deepLearningUtility(args...)",
69 inputs: &IN_REST,
70 outputs: &OUT_ARRAY,
71}];
72
73const OUT_ADAMUPDATE: [BuiltinParamDescriptor; 3] = [
74 BuiltinParamDescriptor {
75 name: "params",
76 ty: BuiltinParamType::Any,
77 arity: BuiltinParamArity::Required,
78 default: None,
79 description: "Updated numeric parameters.",
80 },
81 BuiltinParamDescriptor {
82 name: "averageGrad",
83 ty: BuiltinParamType::Any,
84 arity: BuiltinParamArity::Optional,
85 default: None,
86 description: "Updated first-moment moving average.",
87 },
88 BuiltinParamDescriptor {
89 name: "averageSqGrad",
90 ty: BuiltinParamType::Any,
91 arity: BuiltinParamArity::Optional,
92 default: None,
93 description: "Updated second-moment moving average.",
94 },
95];
96
97const ADAMUPDATE_SIGNATURES: [BuiltinSignatureDescriptor; 2] = [
98 BuiltinSignatureDescriptor {
99 label: "[params, averageGrad, averageSqGrad] = adamupdate(params, grad, averageGrad, averageSqGrad, iteration)",
100 inputs: &IN_REST,
101 outputs: &OUT_ADAMUPDATE,
102 },
103 BuiltinSignatureDescriptor {
104 label: "[params, averageGrad, averageSqGrad] = adamupdate(params, grad, averageGrad, averageSqGrad, iteration, learnRate, gradDecay, sqGradDecay, epsilon)",
105 inputs: &IN_REST,
106 outputs: &OUT_ADAMUPDATE,
107 },
108];
109
110const OUT_VARARG: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
111 name: "varargout",
112 ty: BuiltinParamType::Any,
113 arity: BuiltinParamArity::Variadic,
114 default: None,
115 description: "Outputs returned by the invoked function.",
116}];
117
118const DLFEVAL_INPUTS: [BuiltinParamDescriptor; 2] = [
119 BuiltinParamDescriptor {
120 name: "fun",
121 ty: BuiltinParamType::Any,
122 arity: BuiltinParamArity::Required,
123 default: None,
124 description: "Function handle to evaluate.",
125 },
126 BuiltinParamDescriptor {
127 name: "args",
128 ty: BuiltinParamType::Any,
129 arity: BuiltinParamArity::Variadic,
130 default: None,
131 description: "Arguments forwarded to the function handle.",
132 },
133];
134
135const DLFEVAL_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
136 label: "varargout = dlfeval(fun, args...)",
137 inputs: &DLFEVAL_INPUTS,
138 outputs: &OUT_VARARG,
139}];
140
141const DLUPDATE_INPUTS: [BuiltinParamDescriptor; 2] = [
142 BuiltinParamDescriptor {
143 name: "fun",
144 ty: BuiltinParamType::Any,
145 arity: BuiltinParamArity::Required,
146 default: None,
147 description: "Function handle applied to matching leaves in the parameter trees.",
148 },
149 BuiltinParamDescriptor {
150 name: "args",
151 ty: BuiltinParamType::Any,
152 arity: BuiltinParamArity::Variadic,
153 default: None,
154 description: "One or more compatible parameter trees.",
155 },
156];
157
158const DLUPDATE_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
159 label: "varargout = dlupdate(fun, args...)",
160 inputs: &DLUPDATE_INPUTS,
161 outputs: &OUT_VARARG,
162}];
163
164const DLGRADIENT_INPUTS: [BuiltinParamDescriptor; 2] = [
165 BuiltinParamDescriptor {
166 name: "loss",
167 ty: BuiltinParamType::Any,
168 arity: BuiltinParamArity::Required,
169 default: None,
170 description: "Scalar traced dlarray loss.",
171 },
172 BuiltinParamDescriptor {
173 name: "targets",
174 ty: BuiltinParamType::Any,
175 arity: BuiltinParamArity::Variadic,
176 default: None,
177 description: "Traced dlarray values, learnables, or dlnetworks to differentiate.",
178 },
179];
180
181const DLGRADIENT_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
182 label: "varargout = dlgradient(loss, targets...)",
183 inputs: &DLGRADIENT_INPUTS,
184 outputs: &OUT_VARARG,
185}];
186
187pub const OBJECT_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
188 signatures: &OBJECT_SIGNATURES,
189 output_mode: BuiltinOutputMode::Fixed,
190 completion_policy: BuiltinCompletionPolicy::Public,
191 errors: &ERRORS,
192};
193
194pub const ARRAY_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
195 signatures: &ARRAY_SIGNATURES,
196 output_mode: BuiltinOutputMode::Fixed,
197 completion_policy: BuiltinCompletionPolicy::Public,
198 errors: &ERRORS,
199};
200
201pub const ADAMUPDATE_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
202 signatures: &ADAMUPDATE_SIGNATURES,
203 output_mode: BuiltinOutputMode::ByRequestedOutputCount,
204 completion_policy: BuiltinCompletionPolicy::Public,
205 errors: &ERRORS,
206};
207
208pub const DLFEVAL_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
209 signatures: &DLFEVAL_SIGNATURES,
210 output_mode: BuiltinOutputMode::ByRequestedOutputCount,
211 completion_policy: BuiltinCompletionPolicy::Public,
212 errors: &ERRORS,
213};
214
215pub const DLUPDATE_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
216 signatures: &DLUPDATE_SIGNATURES,
217 output_mode: BuiltinOutputMode::ByRequestedOutputCount,
218 completion_policy: BuiltinCompletionPolicy::Public,
219 errors: &ERRORS,
220};
221
222pub const DLGRADIENT_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
223 signatures: &DLGRADIENT_SIGNATURES,
224 output_mode: BuiltinOutputMode::ByRequestedOutputCount,
225 completion_policy: BuiltinCompletionPolicy::Public,
226 errors: &ERRORS,
227};
228
229pub(super) fn any_type(_args: &[Type], _ctx: &ResolveContext) -> Type {
230 Type::Unknown
231}
232
233pub(super) async fn gather_args(args: Vec<Value>) -> BuiltinResult<Vec<Value>> {
234 let mut gathered = Vec::with_capacity(args.len());
235 for value in args {
236 gathered.push(gather_if_needed_async(&value).await?);
237 }
238 Ok(gathered)
239}
240
241pub(super) fn deep_learning_error(
242 function: &'static str,
243 message: impl Into<String>,
244) -> RuntimeError {
245 descriptor_error(function, message, &ERROR_INVALID_INPUT)
246}
247
248pub(super) fn unsupported_error(
249 function: &'static str,
250 message: impl Into<String>,
251) -> RuntimeError {
252 descriptor_error(function, message, &ERROR_UNSUPPORTED)
253}
254
255fn descriptor_error(
256 function: &'static str,
257 message: impl Into<String>,
258 descriptor: &'static BuiltinErrorDescriptor,
259) -> RuntimeError {
260 let mut builder = build_runtime_error(message).with_builtin(function);
261 if let Some(identifier) = descriptor.identifier {
262 builder = builder.with_identifier(identifier);
263 }
264 builder.build()
265}
266
267pub(super) fn ensure_dlarray_class_registered() {
268 DLARRAY_CLASS_REGISTERED.with(|registered| {
269 if registered.get() {
270 return;
271 }
272 let methods = ["plus", "minus", "times", "rdivide", "mtimes", "sum"]
273 .into_iter()
274 .map(|name| {
275 (
276 name.to_string(),
277 MethodDef {
278 name: name.to_string(),
279 is_static: false,
280 is_abstract: false,
281 is_sealed: false,
282 access: Access::Public,
283 function_name: format!("dlarray.{name}"),
284 implicit_class_argument: None,
285 },
286 )
287 })
288 .collect::<HashMap<_, _>>();
289 runmat_builtins::register_class(ClassDef {
290 name: "dlarray".to_string(),
291 parent: None,
292 properties: HashMap::new(),
293 methods,
294 });
295 registered.set(true);
296 });
297}
298
299pub(super) fn scalar_text(value: &Value, function: &'static str) -> BuiltinResult<String> {
300 match value {
301 Value::String(s) => Ok(s.clone()),
302 Value::CharArray(chars) if chars.rows == 1 => Ok(chars.data.iter().collect()),
303 Value::StringArray(array) if array.data.len() == 1 => Ok(array.data[0].clone()),
304 other => Err(deep_learning_error(
305 function,
306 format!("{function}: expected text scalar, got {other:?}"),
307 )),
308 }
309}
310
311pub(super) fn numeric_scalar(
312 value: &Value,
313 function: &'static str,
314 label: &str,
315) -> BuiltinResult<f64> {
316 match value {
317 Value::Num(n) if n.is_finite() => Ok(*n),
318 Value::Int(i) => Ok(i.to_f64()),
319 Value::Tensor(t) if t.data.len() == 1 && t.data[0].is_finite() => Ok(t.data[0]),
320 other => Err(deep_learning_error(
321 function,
322 format!("{function}: {label} must be a finite numeric scalar, got {other:?}"),
323 )),
324 }
325}
326
327pub(super) fn positive_usize(
328 value: &Value,
329 function: &'static str,
330 label: &str,
331) -> BuiltinResult<usize> {
332 let n = numeric_scalar(value, function, label)?;
333 if n.fract().abs() > f64::EPSILON || n < 1.0 || n > usize::MAX as f64 {
334 return Err(deep_learning_error(
335 function,
336 format!("{function}: {label} must be a positive integer"),
337 ));
338 }
339 Ok(n as usize)
340}
341
342pub(super) fn numeric_vector(
343 value: &Value,
344 function: &'static str,
345 label: &str,
346) -> BuiltinResult<Vec<usize>> {
347 match value {
348 Value::Num(_) | Value::Int(_) | Value::Tensor(_) => {
349 let values = numeric_values(value, function, label)?;
350 let mut out = Vec::with_capacity(values.len());
351 for item in values {
352 if !item.is_finite()
353 || item.fract().abs() > f64::EPSILON
354 || item < 1.0
355 || item > usize::MAX as f64
356 {
357 return Err(deep_learning_error(
358 function,
359 format!("{function}: {label} must contain positive integers"),
360 ));
361 }
362 out.push(item as usize);
363 }
364 Ok(out)
365 }
366 other => Err(deep_learning_error(
367 function,
368 format!("{function}: {label} must be numeric, got {other:?}"),
369 )),
370 }
371}
372
373pub(super) fn numeric_values(
374 value: &Value,
375 function: &'static str,
376 label: &str,
377) -> BuiltinResult<Vec<f64>> {
378 match value {
379 Value::Num(n) => Ok(vec![*n]),
380 Value::Int(i) => Ok(vec![i.to_f64()]),
381 Value::Tensor(t) => Ok(t.data.clone()),
382 other => Err(deep_learning_error(
383 function,
384 format!("{function}: {label} must be numeric, got {other:?}"),
385 )),
386 }
387}
388
389pub(super) fn text_or_missing(
390 value: Option<&Value>,
391 default: &str,
392 function: &'static str,
393) -> BuiltinResult<String> {
394 match value {
395 Some(v) => scalar_text(v, function),
396 None => Ok(default.to_string()),
397 }
398}
399
400pub(super) fn string_array(
401 values: Vec<String>,
402 shape: Vec<usize>,
403 function: &'static str,
404) -> BuiltinResult<Value> {
405 StringArray::new(values, shape)
406 .map(Value::StringArray)
407 .map_err(|err| deep_learning_error(function, err))
408}
409
410pub(super) fn tensor_value(
411 data: Vec<f64>,
412 shape: Vec<usize>,
413 function: &'static str,
414) -> BuiltinResult<Value> {
415 Tensor::new(data, shape)
416 .map(Value::Tensor)
417 .map_err(|err| deep_learning_error(function, err))
418}
419
420pub(super) fn object<K, I>(class_name: &str, properties: I) -> Value
421where
422 K: Into<String>,
423 I: IntoIterator<Item = (K, Value)>,
424{
425 let mut object = ObjectInstance::new(class_name.to_string());
426 for (name, value) in properties {
427 object.properties.insert(name.into(), value);
428 }
429 Value::Object(object)
430}
431
432pub(super) fn layer_object(
433 class_name: &str,
434 type_name: &str,
435 mut properties: Vec<(&str, Value)>,
436 rest: Vec<Value>,
437 function: &'static str,
438) -> BuiltinResult<Value> {
439 let mut owned_properties = properties
440 .drain(..)
441 .map(|(name, value)| (name.to_string(), value))
442 .collect::<Vec<_>>();
443 let mut name = String::new();
444 let mut description = String::new();
445 let mut extra = parse_name_values(rest, function)?;
446 if let Some(value) = extra.remove("name") {
447 name = scalar_text(&value, function)?;
448 }
449 if let Some(value) = extra.remove("description") {
450 description = scalar_text(&value, function)?;
451 }
452 owned_properties.push(("Type".to_string(), Value::String(type_name.to_string())));
453 owned_properties.push(("Name".to_string(), Value::String(name)));
454 owned_properties.push(("Description".to_string(), Value::String(description)));
455 for (key, value) in extra {
456 owned_properties.push((canonical_property_name(&key), value));
457 }
458 Ok(object(class_name, owned_properties))
459}
460
461fn canonical_property_name(name: &str) -> String {
462 match name.to_ascii_lowercase().as_str() {
463 "biaslearnratefactor" => "BiasLearnRateFactor",
464 "biasl2factor" => "BiasL2Factor",
465 "biasinitializer" => "BiasInitializer",
466 "weightslearnratefactor" => "WeightsLearnRateFactor",
467 "weightsl2factor" => "WeightsL2Factor",
468 "weightsinitializer" => "WeightsInitializer",
469 "inputnames" => "InputNames",
470 "outputnames" => "OutputNames",
471 "padding" => "Padding",
472 "stride" => "Stride",
473 "dilationfactor" => "DilationFactor",
474 "numchannels" => "NumChannels",
475 "hasstateinputs" => "HasStateInputs",
476 "hasstateoutputs" => "HasStateOutputs",
477 "outputmode" => "OutputMode",
478 "stateactivationfunction" => "StateActivationFunction",
479 "gateactivationfunction" => "GateActivationFunction",
480 "normalization" => "Normalization",
481 "splitcomplexinputs" => "SplitComplexInputs",
482 "weights" => "Weights",
483 "bias" => "Bias",
484 "classes" => "Classes",
485 "epsilon" => "Epsilon",
486 "alphalearnratefactor" => "AlphaLearnRateFactor",
487 "betalearnratefactor" => "BetaLearnRateFactor",
488 "offset" => "Offset",
489 "scale" => "Scale",
490 other => other,
491 }
492 .to_string()
493}
494
495pub(super) fn parse_name_values(
496 args: Vec<Value>,
497 function: &'static str,
498) -> BuiltinResult<std::collections::BTreeMap<String, Value>> {
499 if !args.len().is_multiple_of(2) {
500 return Err(deep_learning_error(
501 function,
502 format!("{function}: name-value options must be paired"),
503 ));
504 }
505 let mut map = std::collections::BTreeMap::new();
506 let mut idx = 0;
507 while idx < args.len() {
508 let name = scalar_text(&args[idx], function)?.to_ascii_lowercase();
509 map.insert(name, args[idx + 1].clone());
510 idx += 2;
511 }
512 Ok(map)
513}
514
515pub(super) fn layers_from_value(value: Value, function: &'static str) -> BuiltinResult<Vec<Value>> {
516 match value {
517 Value::Object(_) => Ok(vec![value]),
518 Value::Cell(cell) => Ok(cell.data),
519 Value::OutputList(values) => Ok(values),
520 other => Err(deep_learning_error(
521 function,
522 format!("{function}: layers must be a layer object, cell array, or object array, got {other:?}"),
523 )),
524 }
525}
526
527pub(super) fn layer_names(layers: &[Value], function: &'static str) -> BuiltinResult<Vec<String>> {
528 let mut names = Vec::with_capacity(layers.len());
529 for (idx, layer) in layers.iter().enumerate() {
530 match layer {
531 Value::Object(object) => {
532 let name = object
533 .properties
534 .get("Name")
535 .and_then(|value| match value {
536 Value::String(s) if !s.is_empty() => Some(s.clone()),
537 _ => None,
538 })
539 .unwrap_or_else(|| format!("layer_{}", idx + 1));
540 names.push(name);
541 }
542 other => {
543 return Err(deep_learning_error(
544 function,
545 format!("{function}: layer list contains non-object value {other:?}"),
546 ));
547 }
548 }
549 }
550 Ok(names)
551}
552
553pub(crate) mod autodiff;
554pub(crate) mod graph;
555pub(crate) mod layers;
556pub(crate) mod losses;
557pub(crate) mod model;
558pub(crate) mod onnx;
559pub(crate) mod sequences;
560pub(crate) mod supervised;
561pub(crate) mod training;
562
563#[cfg(test)]
564mod tests;