1use 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
328pub(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;