1use std::collections::{BTreeMap, HashSet};
4
5use runmat_builtins::{
6 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
7 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
8 ResolveContext, Type,
9};
10use runmat_macros::runtime_builtin;
11use runmat_value::{
12 CellArray, CharArray, LogicalArray, ObjectInstance, StringArray, Tensor, Value,
13};
14
15use crate::builtins::common::spec::{
16 BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
17 ReductionNaN, ResidencyPolicy, ShapeRequirements,
18};
19use crate::builtins::common::tensor as tensor_helpers;
20use crate::{build_runtime_error, gather_if_needed_async, BuiltinResult, RuntimeError};
21
22const NAME: &str = "optimizableVariable";
23const CLASS_NAME: &str = "optimizableVariable";
24
25const OUTPUT_VAR: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
26 name: "variable",
27 ty: BuiltinParamType::Any,
28 arity: BuiltinParamArity::Required,
29 default: None,
30 description: "Bayesian optimization variable metadata object.",
31}];
32
33const INPUTS_BASE: [BuiltinParamDescriptor; 2] = [
34 BuiltinParamDescriptor {
35 name: "Name",
36 ty: BuiltinParamType::StringScalar,
37 arity: BuiltinParamArity::Required,
38 default: None,
39 description: "Variable name.",
40 },
41 BuiltinParamDescriptor {
42 name: "Range",
43 ty: BuiltinParamType::Any,
44 arity: BuiltinParamArity::Required,
45 default: None,
46 description: "Numeric two-element bounds or categorical choices.",
47 },
48];
49
50const INPUTS_OPTIONS: [BuiltinParamDescriptor; 3] = [
51 BuiltinParamDescriptor {
52 name: "Name",
53 ty: BuiltinParamType::StringScalar,
54 arity: BuiltinParamArity::Required,
55 default: None,
56 description: "Variable name.",
57 },
58 BuiltinParamDescriptor {
59 name: "Range",
60 ty: BuiltinParamType::Any,
61 arity: BuiltinParamArity::Required,
62 default: None,
63 description: "Numeric two-element bounds or categorical choices.",
64 },
65 BuiltinParamDescriptor {
66 name: "NameValue",
67 ty: BuiltinParamType::Any,
68 arity: BuiltinParamArity::Variadic,
69 default: None,
70 description: "Name-value options Type, Transform, and Optimize.",
71 },
72];
73
74const SIGNATURES: [BuiltinSignatureDescriptor; 2] = [
75 BuiltinSignatureDescriptor {
76 label: "v = optimizableVariable(Name, Range)",
77 inputs: &INPUTS_BASE,
78 outputs: &OUTPUT_VAR,
79 },
80 BuiltinSignatureDescriptor {
81 label: "v = optimizableVariable(Name, Range, Name, Value)",
82 inputs: &INPUTS_OPTIONS,
83 outputs: &OUTPUT_VAR,
84 },
85];
86
87const ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
88 code: "RM.OPTIMIZABLE_VARIABLE.INVALID_ARGUMENT",
89 identifier: Some("RunMat:optimizableVariable:InvalidArgument"),
90 when: "The variable name, range, or name-value arguments are malformed.",
91 message: "optimizableVariable: invalid argument",
92};
93
94const ERROR_INVALID_RANGE: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
95 code: "RM.OPTIMIZABLE_VARIABLE.INVALID_RANGE",
96 identifier: Some("RunMat:optimizableVariable:InvalidRange"),
97 when: "The supplied range is incompatible with the requested variable type.",
98 message: "optimizableVariable: invalid range",
99};
100
101const ERROR_INVALID_OPTION: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
102 code: "RM.OPTIMIZABLE_VARIABLE.INVALID_OPTION",
103 identifier: Some("RunMat:optimizableVariable:InvalidOption"),
104 when: "A name-value option is unknown or has an unsupported value.",
105 message: "optimizableVariable: invalid option",
106};
107
108const ERROR_FLOW: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
109 code: "RM.OPTIMIZABLE_VARIABLE.FLOW",
110 identifier: Some("RunMat:optimizableVariable:Flow"),
111 when: "Gathering a metadata input fails.",
112 message: "optimizableVariable: flow failure",
113};
114
115const ERRORS: [BuiltinErrorDescriptor; 4] = [
116 ERROR_INVALID_ARGUMENT,
117 ERROR_INVALID_RANGE,
118 ERROR_INVALID_OPTION,
119 ERROR_FLOW,
120];
121
122pub const OPTIMIZABLE_VARIABLE_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
123 signatures: &SIGNATURES,
124 output_mode: BuiltinOutputMode::Fixed,
125 completion_policy: BuiltinCompletionPolicy::Public,
126 errors: &ERRORS,
127};
128
129#[runmat_macros::register_gpu_spec(
130 builtin_path = "crate::builtins::stats::ml::optimizable_variable"
131)]
132pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
133 name: NAME,
134 op_kind: GpuOpKind::Custom("optimizable-variable-metadata"),
135 supported_precisions: &[],
136 broadcast: BroadcastSemantics::None,
137 provider_hooks: &[],
138 constant_strategy: ConstantStrategy::InlineLiteral,
139 residency: ResidencyPolicy::GatherImmediately,
140 nan_mode: ReductionNaN::Include,
141 two_pass_threshold: None,
142 workgroup_size: None,
143 accepts_nan_mode: false,
144 notes: "Host metadata construction for bayesopt; gpuArray range inputs are gathered before validation.",
145};
146
147#[runmat_macros::register_fusion_spec(
148 builtin_path = "crate::builtins::stats::ml::optimizable_variable"
149)]
150pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
151 name: NAME,
152 shape: ShapeRequirements::Any,
153 constant_strategy: ConstantStrategy::InlineLiteral,
154 elementwise: None,
155 reduction: None,
156 emits_nan: false,
157 notes: "optimizableVariable constructs host metadata and terminates fusion planning.",
158};
159
160#[runtime_builtin(
161 name = "optimizableVariable",
162 category = "stats/ml",
163 summary = "Create Bayesian optimization variable metadata.",
164 keywords = "optimizableVariable,bayesopt,optimization,variable,Type,Transform,Optimize",
165 accel = "cpu",
166 type_resolver(optimizable_variable_type),
167 descriptor(crate::builtins::stats::ml::optimizable_variable::OPTIMIZABLE_VARIABLE_DESCRIPTOR),
168 builtin_path = "crate::builtins::stats::ml::optimizable_variable"
169)]
170async fn optimizable_variable_builtin(
171 name: Value,
172 range: Value,
173 rest: Vec<Value>,
174) -> BuiltinResult<Value> {
175 let name = gather_if_needed_async(&name)
176 .await
177 .map_err(|err| remap_flow(err, "name"))?;
178 let range = gather_if_needed_async(&range)
179 .await
180 .map_err(|err| remap_flow(err, "range"))?;
181 let mut gathered = Vec::with_capacity(rest.len());
182 for value in rest {
183 gathered.push(
184 gather_if_needed_async(&value)
185 .await
186 .map_err(|err| remap_flow(err, "name-value argument"))?,
187 );
188 }
189
190 let variable_name = parse_text_scalar(&name, "Name", &ERROR_INVALID_ARGUMENT)?;
191 if variable_name.trim().is_empty() {
192 return Err(error(
193 "optimizableVariable: Name must be a nonempty text scalar",
194 &ERROR_INVALID_ARGUMENT,
195 ));
196 }
197 let options = Options::parse(gathered)?;
198 let range = OptimizableRange::parse(range, options.var_type)?;
199 let var_type = options.var_type.unwrap_or(range.inferred_type());
200 range.validate_for(var_type, options.transform)?;
201
202 let mut object = ObjectInstance::new(CLASS_NAME.to_string());
203 object
204 .properties
205 .insert("Name".to_string(), Value::String(variable_name));
206 object
207 .properties
208 .insert("Range".to_string(), range.value().clone());
209 object.properties.insert(
210 "Type".to_string(),
211 Value::String(var_type.as_str().to_string()),
212 );
213 object.properties.insert(
214 "Transform".to_string(),
215 Value::String(options.transform.as_str().to_string()),
216 );
217 object
218 .properties
219 .insert("Optimize".to_string(), Value::Bool(options.optimize));
220 Ok(Value::Object(object))
221}
222
223fn optimizable_variable_type(_args: &[Type], _ctx: &ResolveContext) -> Type {
224 Type::Unknown
225}
226
227#[derive(Debug, Clone, Copy, PartialEq, Eq)]
228enum VariableType {
229 Real,
230 Integer,
231 Categorical,
232}
233
234impl VariableType {
235 fn parse(value: &Value) -> BuiltinResult<Self> {
236 match parse_text_scalar(value, "Type", &ERROR_INVALID_OPTION)?
237 .trim()
238 .to_ascii_lowercase()
239 .as_str()
240 {
241 "real" => Ok(Self::Real),
242 "integer" => Ok(Self::Integer),
243 "categorical" => Ok(Self::Categorical),
244 other => Err(error(
245 format!("optimizableVariable: unsupported Type '{other}'"),
246 &ERROR_INVALID_OPTION,
247 )),
248 }
249 }
250
251 fn as_str(self) -> &'static str {
252 match self {
253 Self::Real => "real",
254 Self::Integer => "integer",
255 Self::Categorical => "categorical",
256 }
257 }
258}
259
260#[derive(Debug, Clone, Copy, PartialEq, Eq)]
261enum Transform {
262 None,
263 Log,
264}
265
266impl Transform {
267 fn parse(value: &Value) -> BuiltinResult<Self> {
268 match parse_text_scalar(value, "Transform", &ERROR_INVALID_OPTION)?
269 .trim()
270 .to_ascii_lowercase()
271 .as_str()
272 {
273 "none" => Ok(Self::None),
274 "log" => Ok(Self::Log),
275 other => Err(error(
276 format!("optimizableVariable: unsupported Transform '{other}'"),
277 &ERROR_INVALID_OPTION,
278 )),
279 }
280 }
281
282 fn as_str(self) -> &'static str {
283 match self {
284 Self::None => "none",
285 Self::Log => "log",
286 }
287 }
288}
289
290#[derive(Debug)]
291struct Options {
292 var_type: Option<VariableType>,
293 transform: Transform,
294 optimize: bool,
295}
296
297impl Options {
298 fn parse(args: Vec<Value>) -> BuiltinResult<Self> {
299 if !args.len().is_multiple_of(2) {
300 return Err(error(
301 "optimizableVariable: name-value options must be paired",
302 &ERROR_INVALID_ARGUMENT,
303 ));
304 }
305 let mut options = Self {
306 var_type: None,
307 transform: Transform::None,
308 optimize: true,
309 };
310 let mut pairs = BTreeMap::new();
311 for pair in args.chunks_exact(2) {
312 let name = parse_text_scalar(&pair[0], "option name", &ERROR_INVALID_OPTION)?
313 .trim()
314 .to_ascii_lowercase();
315 pairs.insert(name, pair[1].clone());
316 }
317 for (name, value) in pairs {
318 match name.as_str() {
319 "type" => options.var_type = Some(VariableType::parse(&value)?),
320 "transform" => options.transform = Transform::parse(&value)?,
321 "optimize" => options.optimize = parse_bool(&value, "Optimize")?,
322 other => {
323 return Err(error(
324 format!("optimizableVariable: unknown option '{other}'"),
325 &ERROR_INVALID_OPTION,
326 ));
327 }
328 }
329 }
330 Ok(options)
331 }
332}
333
334#[derive(Debug, Clone)]
335enum OptimizableRange {
336 Numeric {
337 value: Value,
338 lower: f64,
339 upper: f64,
340 },
341 Categorical {
342 value: Value,
343 categories: Vec<String>,
344 },
345}
346
347impl OptimizableRange {
348 fn parse(value: Value, explicit_type: Option<VariableType>) -> BuiltinResult<Self> {
349 match value {
350 Value::Tensor(tensor) => {
351 let values = tensor_helpers::tensor_values_f64(&tensor);
352 Self::numeric_from_values(Value::Tensor(tensor), &values)
353 }
354 Value::Num(n) => Self::numeric_from_values(
355 Value::Tensor(Tensor::new(vec![n], vec![1, 1]).unwrap()),
356 &[n],
357 ),
358 Value::Int(i) => {
359 let n = i.to_f64();
360 Self::numeric_from_values(
361 Value::Tensor(Tensor::new(vec![n], vec![1, 1]).unwrap()),
362 &[n],
363 )
364 }
365 Value::StringArray(array) => {
366 let categories = text_categories_from_string_array(&array)?;
367 Ok(Self::Categorical {
368 value: Value::StringArray(array),
369 categories,
370 })
371 }
372 Value::String(text) if explicit_type == Some(VariableType::Categorical) => {
373 check_categories(std::slice::from_ref(&text))?;
374 Ok(Self::Categorical {
375 value: Value::String(text.clone()),
376 categories: vec![text],
377 })
378 }
379 Value::Cell(cell) => {
380 let categories = text_categories_from_cell(&cell)?;
381 Ok(Self::Categorical {
382 value: Value::Cell(cell),
383 categories,
384 })
385 }
386 Value::CharArray(chars) if explicit_type == Some(VariableType::Categorical) => {
387 let category = char_row_to_string(&chars, "Range", &ERROR_INVALID_RANGE)?;
388 check_categories(std::slice::from_ref(&category))?;
389 let value = Value::Cell(
390 CellArray::new(vec![Value::String(category.clone())], 1, 1)
391 .map_err(|err| error(err, &ERROR_INVALID_RANGE))?,
392 );
393 Ok(Self::Categorical {
394 value,
395 categories: vec![category],
396 })
397 }
398 other => Err(error(
399 format!("optimizableVariable: unsupported Range value {other:?}"),
400 &ERROR_INVALID_RANGE,
401 )),
402 }
403 }
404
405 fn numeric_from_values(value: Value, values: &[f64]) -> BuiltinResult<Self> {
406 if values.len() != 2 {
407 return Err(error(
408 "optimizableVariable: numeric Range must contain exactly two elements",
409 &ERROR_INVALID_RANGE,
410 ));
411 }
412 let lower = values[0];
413 let upper = values[1];
414 if !lower.is_finite() || !upper.is_finite() || lower >= upper {
415 return Err(error(
416 "optimizableVariable: numeric Range bounds must be finite and increasing",
417 &ERROR_INVALID_RANGE,
418 ));
419 }
420 Ok(Self::Numeric {
421 value,
422 lower,
423 upper,
424 })
425 }
426
427 fn inferred_type(&self) -> VariableType {
428 match self {
429 Self::Numeric { .. } => VariableType::Real,
430 Self::Categorical { .. } => VariableType::Categorical,
431 }
432 }
433
434 fn value(&self) -> &Value {
435 match self {
436 Self::Numeric { value, .. } | Self::Categorical { value, .. } => value,
437 }
438 }
439
440 fn validate_for(&self, var_type: VariableType, transform: Transform) -> BuiltinResult<()> {
441 match (self, var_type) {
442 (Self::Numeric { lower, .. }, VariableType::Real) => {
443 validate_real_transform(*lower, transform)
444 }
445 (Self::Numeric { lower, upper, .. }, VariableType::Integer) => {
446 if lower.fract().abs() > f64::EPSILON || upper.fract().abs() > f64::EPSILON {
447 return Err(error(
448 "optimizableVariable: integer Range bounds must be whole numbers",
449 &ERROR_INVALID_RANGE,
450 ));
451 }
452 validate_integer_transform(*lower, transform)
453 }
454 (Self::Numeric { .. }, VariableType::Categorical) => Err(error(
455 "optimizableVariable: categorical variables require a categorical text Range",
456 &ERROR_INVALID_RANGE,
457 )),
458 (Self::Categorical { categories, .. }, VariableType::Categorical) => {
459 check_categories(categories)?;
460 if transform != Transform::None {
461 return Err(error(
462 "optimizableVariable: categorical variables only support Transform 'none'",
463 &ERROR_INVALID_OPTION,
464 ));
465 }
466 Ok(())
467 }
468 (Self::Categorical { .. }, _) => Err(error(
469 "optimizableVariable: real and integer variables require a two-element numeric Range",
470 &ERROR_INVALID_RANGE,
471 )),
472 }
473 }
474}
475
476fn validate_real_transform(lower: f64, transform: Transform) -> BuiltinResult<()> {
477 if transform == Transform::Log && lower <= 0.0 {
478 return Err(error(
479 "optimizableVariable: log transform requires positive real Range bounds",
480 &ERROR_INVALID_RANGE,
481 ));
482 }
483 Ok(())
484}
485
486fn validate_integer_transform(lower: f64, transform: Transform) -> BuiltinResult<()> {
487 if transform == Transform::Log && lower < 0.0 {
488 return Err(error(
489 "optimizableVariable: log transform requires nonnegative integer Range bounds",
490 &ERROR_INVALID_RANGE,
491 ));
492 }
493 Ok(())
494}
495
496fn text_categories_from_string_array(array: &StringArray) -> BuiltinResult<Vec<String>> {
497 let categories = array.data.clone();
498 check_categories(&categories)?;
499 Ok(categories)
500}
501
502fn text_categories_from_cell(cell: &CellArray) -> BuiltinResult<Vec<String>> {
503 let mut categories = Vec::with_capacity(cell.data.len());
504 for value in &cell.data {
505 categories.push(parse_text_scalar(value, "Range", &ERROR_INVALID_RANGE)?);
506 }
507 check_categories(&categories)?;
508 Ok(categories)
509}
510
511fn check_categories(categories: &[String]) -> BuiltinResult<()> {
512 if categories.is_empty() {
513 return Err(error(
514 "optimizableVariable: categorical Range must contain at least one category",
515 &ERROR_INVALID_RANGE,
516 ));
517 }
518 let mut seen = HashSet::with_capacity(categories.len());
519 for category in categories {
520 if category.trim().is_empty() {
521 return Err(error(
522 "optimizableVariable: categorical Range values must be nonempty text",
523 &ERROR_INVALID_RANGE,
524 ));
525 }
526 if !seen.insert(category.clone()) {
527 return Err(error(
528 "optimizableVariable: categorical Range values must be unique",
529 &ERROR_INVALID_RANGE,
530 ));
531 }
532 }
533 Ok(())
534}
535
536fn parse_text_scalar(
537 value: &Value,
538 label: &str,
539 descriptor: &'static BuiltinErrorDescriptor,
540) -> BuiltinResult<String> {
541 match value {
542 Value::String(text) => Ok(text.clone()),
543 Value::CharArray(chars) if chars.rows == 1 => char_row_to_string(chars, label, descriptor),
544 Value::StringArray(array) if array.data.len() == 1 => Ok(array.data[0].clone()),
545 other => Err(error(
546 format!("optimizableVariable: {label} must be a text scalar, got {other:?}"),
547 descriptor,
548 )),
549 }
550}
551
552fn char_row_to_string(
553 chars: &CharArray,
554 label: &str,
555 descriptor: &'static BuiltinErrorDescriptor,
556) -> BuiltinResult<String> {
557 if chars.rows != 1 {
558 return Err(error(
559 format!("optimizableVariable: {label} char array must be a row vector"),
560 descriptor,
561 ));
562 }
563 Ok(chars.data.iter().collect())
564}
565
566fn parse_bool(value: &Value, label: &str) -> BuiltinResult<bool> {
567 match value {
568 Value::Bool(flag) => Ok(*flag),
569 Value::LogicalArray(LogicalArray { data, .. }) if data.len() == 1 => Ok(data[0] != 0),
570 Value::Num(n) if *n == 0.0 || *n == 1.0 => Ok(*n != 0.0),
571 Value::Int(i) if i.to_f64() == 0.0 || i.to_f64() == 1.0 => Ok(i.to_f64() != 0.0),
572 other => Err(error(
573 format!("optimizableVariable: {label} must be a scalar logical value, got {other:?}"),
574 &ERROR_INVALID_OPTION,
575 )),
576 }
577}
578
579fn error(message: impl Into<String>, descriptor: &'static BuiltinErrorDescriptor) -> RuntimeError {
580 let mut builder = build_runtime_error(message).with_builtin(NAME);
581 if let Some(identifier) = descriptor.identifier {
582 builder = builder.with_identifier(identifier);
583 }
584 builder.build()
585}
586
587fn remap_flow(err: RuntimeError, label: &str) -> RuntimeError {
588 error(
589 format!(
590 "optimizableVariable: failed to gather {label}: {}",
591 err.message()
592 ),
593 &ERROR_FLOW,
594 )
595}
596
597#[cfg(test)]
598mod tests {
599 use super::*;
600 use futures::executor::block_on;
601 use runmat_value::IntegerStorage;
602
603 fn object(value: Value) -> ObjectInstance {
604 let Value::Object(object) = value else {
605 panic!("expected object");
606 };
607 object
608 }
609
610 fn typed_integer_range(storage: IntegerStorage) -> Value {
611 let len = storage.len();
612 let tensor = Tensor::new_integer(storage, vec![1, len]).expect("integer range");
613 Value::Tensor(tensor)
614 }
615
616 #[test]
617 fn builds_real_variable_with_defaults() {
618 let out = block_on(optimizable_variable_builtin(
619 Value::String("depth".into()),
620 Value::Tensor(Tensor::new(vec![1.0, 10.0], vec![1, 2]).unwrap()),
621 vec![],
622 ))
623 .expect("optimizableVariable");
624 let object = object(out);
625 assert_eq!(object.class_name, CLASS_NAME);
626 assert_eq!(
627 object.properties.get("Name"),
628 Some(&Value::String("depth".into()))
629 );
630 assert_eq!(
631 object.properties.get("Type"),
632 Some(&Value::String("real".into()))
633 );
634 assert_eq!(
635 object.properties.get("Transform"),
636 Some(&Value::String("none".into()))
637 );
638 assert_eq!(object.properties.get("Optimize"), Some(&Value::Bool(true)));
639 }
640
641 #[test]
642 fn builds_integer_log_variable_with_optimize_false() {
643 let out = block_on(optimizable_variable_builtin(
644 Value::String("trees".into()),
645 Value::Tensor(Tensor::new(vec![0.0, 1000.0], vec![1, 2]).unwrap()),
646 vec![
647 Value::String("Type".into()),
648 Value::String("integer".into()),
649 Value::String("Transform".into()),
650 Value::String("log".into()),
651 Value::String("Optimize".into()),
652 Value::Bool(false),
653 ],
654 ))
655 .expect("optimizableVariable integer");
656 let object = object(out);
657 assert_eq!(
658 object.properties.get("Type"),
659 Some(&Value::String("integer".into()))
660 );
661 assert_eq!(
662 object.properties.get("Transform"),
663 Some(&Value::String("log".into()))
664 );
665 assert_eq!(object.properties.get("Optimize"), Some(&Value::Bool(false)));
666 }
667
668 #[test]
669 fn typed_integer_range_bounds_are_read_from_exact_storage() {
670 let range = typed_integer_range(IntegerStorage::I16(vec![0, 1000]));
671 let out = block_on(optimizable_variable_builtin(
672 Value::String("trees".into()),
673 range,
674 vec![
675 Value::String("Type".into()),
676 Value::String("integer".into()),
677 Value::String("Transform".into()),
678 Value::String("log".into()),
679 ],
680 ))
681 .expect("optimizableVariable integer");
682 let object = object(out);
683 assert_eq!(
684 object.properties.get("Type"),
685 Some(&Value::String("integer".into()))
686 );
687 let Some(Value::Tensor(range)) = object.properties.get("Range") else {
688 panic!("expected tensor Range");
689 };
690 assert_eq!(
691 range.integer_storage(),
692 Some(&IntegerStorage::I16(vec![0, 1000]))
693 );
694 assert_eq!(range.materialize_f64(), vec![0.0, 1000.0]);
695 }
696
697 #[test]
698 fn builds_categorical_variable_from_cellstr() {
699 let categories = Value::Cell(
700 CellArray::new(
701 vec![Value::String("linear".into()), Value::String("rbf".into())],
702 1,
703 2,
704 )
705 .unwrap(),
706 );
707 let out = block_on(optimizable_variable_builtin(
708 Value::String("kernel".into()),
709 categories.clone(),
710 vec![
711 Value::String("Type".into()),
712 Value::String("categorical".into()),
713 ],
714 ))
715 .expect("optimizableVariable categorical");
716 let object = object(out);
717 assert_eq!(
718 object.properties.get("Type"),
719 Some(&Value::String("categorical".into()))
720 );
721 assert_eq!(object.properties.get("Range"), Some(&categories));
722 }
723
724 #[test]
725 fn infers_categorical_type_from_text_range() {
726 let out = block_on(optimizable_variable_builtin(
727 Value::String("method".into()),
728 Value::StringArray(StringArray::new(vec!["a".into(), "b".into()], vec![1, 2]).unwrap()),
729 vec![],
730 ))
731 .expect("optimizableVariable categorical inferred");
732 let object = object(out);
733 assert_eq!(
734 object.properties.get("Type"),
735 Some(&Value::String("categorical".into()))
736 );
737 }
738
739 #[test]
740 fn accepts_explicit_categorical_scalar_string_range() {
741 let out = block_on(optimizable_variable_builtin(
742 Value::String("method".into()),
743 Value::String("rbf".into()),
744 vec![
745 Value::String("Type".into()),
746 Value::String("categorical".into()),
747 ],
748 ))
749 .expect("optimizableVariable scalar categorical");
750 let object = object(out);
751 assert_eq!(
752 object.properties.get("Range"),
753 Some(&Value::String("rbf".into()))
754 );
755 assert_eq!(
756 object.properties.get("Type"),
757 Some(&Value::String("categorical".into()))
758 );
759 }
760
761 #[test]
762 fn rejects_invalid_ranges_and_options() {
763 let err = block_on(optimizable_variable_builtin(
764 Value::String("bad".into()),
765 Value::Tensor(Tensor::new(vec![2.0, 1.0], vec![1, 2]).unwrap()),
766 vec![],
767 ))
768 .unwrap_err();
769 assert!(err.to_string().contains("finite and increasing"));
770
771 let err = block_on(optimizable_variable_builtin(
772 Value::String("bad".into()),
773 Value::Tensor(Tensor::new(vec![1.5, 3.0], vec![1, 2]).unwrap()),
774 vec![
775 Value::String("Type".into()),
776 Value::String("integer".into()),
777 ],
778 ))
779 .unwrap_err();
780 assert!(err.to_string().contains("whole numbers"));
781
782 let err = block_on(optimizable_variable_builtin(
783 Value::String("bad".into()),
784 Value::Tensor(Tensor::new(vec![0.0, 3.0], vec![1, 2]).unwrap()),
785 vec![
786 Value::String("Transform".into()),
787 Value::String("log".into()),
788 ],
789 ))
790 .unwrap_err();
791 assert!(err.to_string().contains("positive real"));
792
793 let err = block_on(optimizable_variable_builtin(
794 Value::String("bad".into()),
795 Value::Tensor(Tensor::new(vec![-1.0, 3.0], vec![1, 2]).unwrap()),
796 vec![
797 Value::String("Type".into()),
798 Value::String("integer".into()),
799 Value::String("Transform".into()),
800 Value::String("log".into()),
801 ],
802 ))
803 .unwrap_err();
804 assert!(err.to_string().contains("nonnegative integer"));
805
806 let err = block_on(optimizable_variable_builtin(
807 Value::String("bad".into()),
808 Value::Tensor(Tensor::new(vec![1.0, 3.0], vec![1, 2]).unwrap()),
809 vec![Value::String("Type".into()), Value::String("int".into())],
810 ))
811 .unwrap_err();
812 assert!(err.to_string().contains("unsupported Type"));
813
814 let err = block_on(optimizable_variable_builtin(
815 Value::String("bad".into()),
816 Value::StringArray(StringArray::new(vec!["a".into(), "a".into()], vec![1, 2]).unwrap()),
817 vec![
818 Value::String("Type".into()),
819 Value::String("categorical".into()),
820 ],
821 ))
822 .unwrap_err();
823 assert!(err.to_string().contains("unique"));
824 }
825}