Skip to main content

runmat_runtime/builtins/deep_learning/
mod.rs

1//! Deep Learning Toolbox compatibility builtins.
2
3use 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;