Skip to main content

runmat_runtime/builtins/deep_learning/
mod.rs

1//! Deep Learning Toolbox compatibility builtins.
2use runmat_types::MemberAccess;
3
4use runmat_builtins::{
5    BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
6    BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
7    ResolveContext, Type,
8};
9use runmat_value::{ObjectInstance, StringArray, Tensor, Value};
10use std::collections::HashMap;
11
12use crate::{build_runtime_error, gather_if_needed_async, BuiltinResult, RuntimeError};
13
14pub(super) const MAX_COMBVEC_COLUMNS: usize = 1_000_000;
15pub(super) const MAX_PAD_ELEMENTS: usize = 10_000_000;
16
17const ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
18    code: "RM.DEEP_LEARNING.INVALID_INPUT",
19    identifier: Some("RunMat:deepLearning:InvalidInput"),
20    when:
21        "Inputs or name-value options do not match the supported Deep Learning compatibility forms.",
22    message: "deep learning builtin received invalid input",
23};
24
25const ERROR_UNSUPPORTED: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
26    code: "RM.DEEP_LEARNING.UNSUPPORTED",
27    identifier: Some("RunMat:deepLearning:Unsupported"),
28    when: "The requested operation requires training, autodiff, export, or UI infrastructure outside this compatibility slice.",
29    message: "deep learning operation is not supported in this slice",
30};
31
32const ERRORS: [BuiltinErrorDescriptor; 2] = [ERROR_INVALID_INPUT, ERROR_UNSUPPORTED];
33
34static DLARRAY_CLASS_REGISTERED: crate::class_registry::ClassRegistration =
35    crate::class_registry::ClassRegistration::new("dlarray");
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.ensure(|| {
269        let methods = ["plus", "minus", "times", "rdivide", "mtimes", "sum"]
270            .into_iter()
271            .map(|name| {
272                (
273                    name.to_string(),
274                    crate::class_registry::RuntimeMethod {
275                        name: name.to_string(),
276                        is_static: false,
277                        is_abstract: false,
278                        is_sealed: false,
279                        access: MemberAccess::Public,
280                        function_name: format!("dlarray.{name}"),
281                        implicit_class_argument: None,
282                    },
283                )
284            })
285            .collect::<HashMap<_, _>>();
286        crate::class_registry::register_class(crate::class_registry::RuntimeClass {
287            name: "dlarray".to_string(),
288            parent: None,
289            properties: HashMap::new(),
290            methods,
291        });
292    });
293}
294
295pub(super) fn scalar_text(value: &Value, function: &'static str) -> BuiltinResult<String> {
296    match value {
297        Value::String(s) => Ok(s.clone()),
298        Value::CharArray(chars) if chars.rows == 1 => Ok(chars.data.iter().collect()),
299        Value::StringArray(array) if array.data.len() == 1 => Ok(array.data[0].clone()),
300        other => Err(deep_learning_error(
301            function,
302            format!("{function}: expected text scalar, got {other:?}"),
303        )),
304    }
305}
306
307pub(super) fn numeric_scalar(
308    value: &Value,
309    function: &'static str,
310    label: &str,
311) -> BuiltinResult<f64> {
312    match value {
313        Value::Num(n) if n.is_finite() => Ok(*n),
314        Value::Int(i) => Ok(i.to_f64()),
315        Value::Tensor(t)
316            if crate::builtins::common::tensor::is_scalar_tensor(t)
317                && crate::builtins::common::tensor::tensor_value_f64(t, 0).is_finite() =>
318        {
319            Ok(crate::builtins::common::tensor::tensor_value_f64(t, 0))
320        }
321        other => Err(deep_learning_error(
322            function,
323            format!("{function}: {label} must be a finite numeric scalar, got {other:?}"),
324        )),
325    }
326}
327
328/// Parse a scalar flag without consulting an integer tensor's compatibility
329/// `f64` mirror.  Structural options use this rather than treating an integer
330/// value as ordinary numeric data.
331pub(super) fn logical_scalar(
332    value: &Value,
333    function: &'static str,
334    label: &str,
335) -> BuiltinResult<bool> {
336    if let Value::Bool(flag) = value {
337        return Ok(*flag);
338    }
339    if let Some(integer) = crate::builtins::common::tensor::scalar_integer_value(value) {
340        return match integer.try_to_i64() {
341            Some(0) => Ok(false),
342            Some(1) => Ok(true),
343            _ => Err(deep_learning_error(
344                function,
345                format!("{function}: {label} must be logical scalar true or false"),
346            )),
347        };
348    }
349    let number = match value {
350        Value::Num(number) => *number,
351        Value::Tensor(tensor) if crate::builtins::common::tensor::is_scalar_tensor(tensor) => {
352            crate::builtins::common::tensor::tensor_value_f64(tensor, 0)
353        }
354        other => {
355            return Err(deep_learning_error(
356                function,
357                format!("{function}: {label} must be logical scalar true or false, got {other:?}"),
358            ));
359        }
360    };
361    match number {
362        0.0 => Ok(false),
363        1.0 => Ok(true),
364        _ => Err(deep_learning_error(
365            function,
366            format!("{function}: {label} must be logical scalar true or false"),
367        )),
368    }
369}
370
371pub(super) fn positive_i64(
372    value: &Value,
373    function: &'static str,
374    label: &str,
375) -> BuiltinResult<i64> {
376    if let Some(integer) = crate::builtins::common::tensor::scalar_integer_value(value) {
377        return integer
378            .try_to_i64()
379            .filter(|value| *value >= 1)
380            .ok_or_else(|| {
381                deep_learning_error(
382                    function,
383                    format!("{function}: {label} must be a positive integer scalar"),
384                )
385            });
386    }
387    let number = numeric_scalar(value, function, label)?;
388    if number.fract().abs() > f64::EPSILON || number < 1.0 || number >= i64::MAX as f64 {
389        return Err(deep_learning_error(
390            function,
391            format!("{function}: {label} must be a positive integer scalar"),
392        ));
393    }
394    Ok(number as i64)
395}
396
397pub(super) fn positive_usize(
398    value: &Value,
399    function: &'static str,
400    label: &str,
401) -> BuiltinResult<usize> {
402    match value {
403        Value::Int(value) => value
404            .try_to_usize()
405            .filter(|value| *value >= 1)
406            .ok_or_else(|| {
407                deep_learning_error(
408                    function,
409                    format!("{function}: {label} must be a positive integer"),
410                )
411            }),
412        Value::Tensor(tensor) if crate::builtins::common::tensor::is_scalar_tensor(tensor) => {
413            if let Some(value) = tensor
414                .integer_storage()
415                .and_then(|storage| storage.value_at(0))
416            {
417                return value
418                    .try_to_usize()
419                    .filter(|value| *value >= 1)
420                    .ok_or_else(|| {
421                        deep_learning_error(
422                            function,
423                            format!("{function}: {label} must be a positive integer"),
424                        )
425                    });
426            }
427            let n = crate::builtins::common::tensor::tensor_value_f64(tensor, 0);
428            positive_usize_from_f64(n, function, label)
429        }
430        _ => {
431            let n = numeric_scalar(value, function, label)?;
432            positive_usize_from_f64(n, function, label)
433        }
434    }
435}
436
437pub(super) fn nonnegative_usize(
438    value: &Value,
439    function: &'static str,
440    label: &str,
441) -> Option<usize> {
442    match value {
443        Value::Int(value) => value.try_to_usize(),
444        Value::Tensor(tensor) if crate::builtins::common::tensor::is_scalar_tensor(tensor) => {
445            if let Some(value) = tensor
446                .integer_storage()
447                .and_then(|storage| storage.value_at(0))
448            {
449                return value.try_to_usize();
450            }
451            nonnegative_usize_from_f64(crate::builtins::common::tensor::tensor_value_f64(tensor, 0))
452        }
453        Value::Num(n) => nonnegative_usize_from_f64(*n),
454        _ => {
455            let _ = (function, label);
456            None
457        }
458    }
459}
460
461fn positive_usize_from_f64(n: f64, function: &'static str, label: &str) -> BuiltinResult<usize> {
462    if !n.is_finite()
463        || n.fract().abs() > f64::EPSILON
464        || n < 1.0
465        || n > usize::MAX as f64
466        || (usize::BITS == 64 && n == usize::MAX as f64)
467    {
468        return Err(deep_learning_error(
469            function,
470            format!("{function}: {label} must be a positive integer"),
471        ));
472    }
473    Ok(n as usize)
474}
475
476fn nonnegative_usize_from_f64(n: f64) -> Option<usize> {
477    if n.is_finite()
478        && n >= 0.0
479        && n.fract() == 0.0
480        && (n < usize::MAX as f64 || (usize::BITS < 64 && n == usize::MAX as f64))
481    {
482        Some(n as usize)
483    } else {
484        None
485    }
486}
487
488pub(super) fn numeric_vector(
489    value: &Value,
490    function: &'static str,
491    label: &str,
492) -> BuiltinResult<Vec<usize>> {
493    match value {
494        Value::Int(value) => value
495            .try_to_usize()
496            .filter(|value| *value >= 1)
497            .map(|value| vec![value])
498            .ok_or_else(|| {
499                deep_learning_error(
500                    function,
501                    format!("{function}: {label} must contain positive integers"),
502                )
503            }),
504        Value::Tensor(tensor) if tensor.integer_storage().is_some() => {
505            let storage = tensor.integer_storage().expect("checked integer storage");
506            let mut out = Vec::with_capacity(storage.len());
507            for index in 0..storage.len() {
508                let Some(value) = storage
509                    .value_at(index)
510                    .and_then(|value| value.try_to_usize())
511                    .filter(|value| *value >= 1)
512                else {
513                    return Err(deep_learning_error(
514                        function,
515                        format!("{function}: {label} must contain positive integers"),
516                    ));
517                };
518                out.push(value);
519            }
520            Ok(out)
521        }
522        Value::Num(_) | Value::Tensor(_) => {
523            let values = numeric_values(value, function, label)?;
524            let mut out = Vec::with_capacity(values.len());
525            for item in values {
526                if !item.is_finite()
527                    || item.fract().abs() > f64::EPSILON
528                    || item < 1.0
529                    || item > usize::MAX as f64
530                    || (usize::BITS == 64 && item == usize::MAX as f64)
531                {
532                    return Err(deep_learning_error(
533                        function,
534                        format!("{function}: {label} must contain positive integers"),
535                    ));
536                }
537                out.push(item as usize);
538            }
539            Ok(out)
540        }
541        other => Err(deep_learning_error(
542            function,
543            format!("{function}: {label} must be numeric, got {other:?}"),
544        )),
545    }
546}
547
548pub(super) fn numeric_values(
549    value: &Value,
550    function: &'static str,
551    label: &str,
552) -> BuiltinResult<Vec<f64>> {
553    match value {
554        Value::Num(n) => Ok(vec![*n]),
555        Value::Int(i) => Ok(vec![i.to_f64()]),
556        Value::Tensor(t) => Ok(crate::builtins::common::tensor::tensor_values_f64(t)),
557        other => Err(deep_learning_error(
558            function,
559            format!("{function}: {label} must be numeric, got {other:?}"),
560        )),
561    }
562}
563
564pub(super) fn text_or_missing(
565    value: Option<&Value>,
566    default: &str,
567    function: &'static str,
568) -> BuiltinResult<String> {
569    match value {
570        Some(v) => scalar_text(v, function),
571        None => Ok(default.to_string()),
572    }
573}
574
575pub(super) fn string_array(
576    values: Vec<String>,
577    shape: Vec<usize>,
578    function: &'static str,
579) -> BuiltinResult<Value> {
580    StringArray::new(values, shape)
581        .map(Value::StringArray)
582        .map_err(|err| deep_learning_error(function, err))
583}
584
585pub(super) fn tensor_value(
586    data: Vec<f64>,
587    shape: Vec<usize>,
588    function: &'static str,
589) -> BuiltinResult<Value> {
590    Tensor::new(data, shape)
591        .map(Value::Tensor)
592        .map_err(|err| deep_learning_error(function, err))
593}
594
595pub(super) fn object<K, I>(class_name: &str, properties: I) -> Value
596where
597    K: Into<String>,
598    I: IntoIterator<Item = (K, Value)>,
599{
600    let mut object = ObjectInstance::new(class_name.to_string());
601    for (name, value) in properties {
602        object.properties.insert(name.into(), value);
603    }
604    Value::Object(object)
605}
606
607pub(super) fn layer_object(
608    class_name: &str,
609    type_name: &str,
610    mut properties: Vec<(&str, Value)>,
611    rest: Vec<Value>,
612    function: &'static str,
613) -> BuiltinResult<Value> {
614    let mut owned_properties = properties
615        .drain(..)
616        .map(|(name, value)| (name.to_string(), value))
617        .collect::<Vec<_>>();
618    let mut name = String::new();
619    let mut description = String::new();
620    let mut extra = parse_name_values(rest, function)?;
621    if let Some(value) = extra.remove("name") {
622        name = scalar_text(&value, function)?;
623    }
624    if let Some(value) = extra.remove("description") {
625        description = scalar_text(&value, function)?;
626    }
627    owned_properties.push(("Type".to_string(), Value::String(type_name.to_string())));
628    owned_properties.push(("Name".to_string(), Value::String(name)));
629    owned_properties.push(("Description".to_string(), Value::String(description)));
630    for (key, value) in extra {
631        owned_properties.push((canonical_property_name(&key), value));
632    }
633    Ok(object(class_name, owned_properties))
634}
635
636fn canonical_property_name(name: &str) -> String {
637    match name.to_ascii_lowercase().as_str() {
638        "biaslearnratefactor" => "BiasLearnRateFactor",
639        "biasl2factor" => "BiasL2Factor",
640        "biasinitializer" => "BiasInitializer",
641        "weightslearnratefactor" => "WeightsLearnRateFactor",
642        "weightsl2factor" => "WeightsL2Factor",
643        "weightsinitializer" => "WeightsInitializer",
644        "inputnames" => "InputNames",
645        "outputnames" => "OutputNames",
646        "padding" => "Padding",
647        "stride" => "Stride",
648        "dilationfactor" => "DilationFactor",
649        "numchannels" => "NumChannels",
650        "hasstateinputs" => "HasStateInputs",
651        "hasstateoutputs" => "HasStateOutputs",
652        "outputmode" => "OutputMode",
653        "stateactivationfunction" => "StateActivationFunction",
654        "gateactivationfunction" => "GateActivationFunction",
655        "normalization" => "Normalization",
656        "splitcomplexinputs" => "SplitComplexInputs",
657        "weights" => "Weights",
658        "bias" => "Bias",
659        "classes" => "Classes",
660        "epsilon" => "Epsilon",
661        "alphalearnratefactor" => "AlphaLearnRateFactor",
662        "betalearnratefactor" => "BetaLearnRateFactor",
663        "offset" => "Offset",
664        "scale" => "Scale",
665        other => other,
666    }
667    .to_string()
668}
669
670pub(super) fn parse_name_values(
671    args: Vec<Value>,
672    function: &'static str,
673) -> BuiltinResult<std::collections::BTreeMap<String, Value>> {
674    if !args.len().is_multiple_of(2) {
675        return Err(deep_learning_error(
676            function,
677            format!("{function}: name-value options must be paired"),
678        ));
679    }
680    let mut map = std::collections::BTreeMap::new();
681    let mut idx = 0;
682    while idx < args.len() {
683        let name = scalar_text(&args[idx], function)?.to_ascii_lowercase();
684        map.insert(name, args[idx + 1].clone());
685        idx += 2;
686    }
687    Ok(map)
688}
689
690pub(super) fn layers_from_value(value: Value, function: &'static str) -> BuiltinResult<Vec<Value>> {
691    match value {
692        Value::Object(_) => Ok(vec![value]),
693        Value::Cell(cell) => Ok(cell.data),
694        Value::OutputList(values) => Ok(values),
695        other => Err(deep_learning_error(
696            function,
697            format!("{function}: layers must be a layer object, cell array, or object array, got {other:?}"),
698        )),
699    }
700}
701
702pub(super) fn layer_names(layers: &[Value], function: &'static str) -> BuiltinResult<Vec<String>> {
703    let mut names = Vec::with_capacity(layers.len());
704    for (idx, layer) in layers.iter().enumerate() {
705        match layer {
706            Value::Object(object) => {
707                let name = object
708                    .properties
709                    .get("Name")
710                    .and_then(|value| match value {
711                        Value::String(s) if !s.is_empty() => Some(s.clone()),
712                        _ => None,
713                    })
714                    .unwrap_or_else(|| format!("layer_{}", idx + 1));
715                names.push(name);
716            }
717            other => {
718                return Err(deep_learning_error(
719                    function,
720                    format!("{function}: layer list contains non-object value {other:?}"),
721                ));
722            }
723        }
724    }
725    Ok(names)
726}
727
728pub(crate) mod autodiff;
729pub(crate) mod graph;
730pub(crate) mod layers;
731pub(crate) mod losses;
732pub(crate) mod model;
733pub(crate) mod onnx;
734pub(crate) mod sequences;
735pub(crate) mod supervised;
736pub(crate) mod training;
737
738#[cfg(test)]
739mod tests;