runmat_runtime/builtins/logical/tests/
isgpuarray.rs1use runmat_builtins::{
4 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinIntegerBackendRule,
5 BuiltinIntegerCapabilityDescriptor, BuiltinIntegerComputationDomain,
6 BuiltinIntegerInputAvailability, BuiltinIntegerInputCapability, BuiltinIntegerOutputClassRule,
7 BuiltinIntegerOverflowRule, BuiltinIntegerOverloadKind, BuiltinIntegerScalarDoubleRule,
8 BuiltinOutputMode, BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType,
9 BuiltinSignatureDescriptor, ResolveContext, Type,
10};
11use runmat_macros::runtime_builtin;
12use runmat_value::Value;
13
14use crate::builtins::common::spec::{
15 BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
16 ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
17};
18use crate::BuiltinResult;
19
20#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::logical::tests::isgpuarray")]
21pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
22 name: "isgpuarray",
23 op_kind: GpuOpKind::Custom("metadata"),
24 supported_precisions: &[ScalarType::F32, ScalarType::F64],
25 broadcast: BroadcastSemantics::None,
26 provider_hooks: &[],
27 constant_strategy: ConstantStrategy::InlineLiteral,
28 residency: ResidencyPolicy::GatherImmediately,
29 nan_mode: ReductionNaN::Include,
30 two_pass_threshold: None,
31 workgroup_size: None,
32 accepts_nan_mode: false,
33 notes: "Reports whether the value is a gpuArray without gathering device buffers.",
34};
35
36#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::logical::tests::isgpuarray")]
37pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
38 name: "isgpuarray",
39 shape: ShapeRequirements::Any,
40 constant_strategy: ConstantStrategy::InlineLiteral,
41 elementwise: None,
42 reduction: None,
43 emits_nan: false,
44 notes: "Metadata query that executes outside of fusion pipelines.",
45};
46
47const ISGPUARRAY_OUTPUT: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
48 name: "tf",
49 ty: BuiltinParamType::LogicalArray,
50 arity: BuiltinParamArity::Required,
51 default: None,
52 description: "True when input is a gpuArray handle.",
53}];
54
55const ISGPUARRAY_INPUTS: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
56 name: "A",
57 ty: BuiltinParamType::Any,
58 arity: BuiltinParamArity::Required,
59 default: None,
60 description: "Input value to test.",
61}];
62
63const ISGPUARRAY_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
64 label: "tf = isgpuarray(A)",
65 inputs: &ISGPUARRAY_INPUTS,
66 outputs: &ISGPUARRAY_OUTPUT,
67}];
68
69const ISGPUARRAY_ERRORS: [BuiltinErrorDescriptor; 0] = [];
70
71pub const ISGPUARRAY_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
72 signatures: &ISGPUARRAY_SIGNATURES,
73 output_mode: BuiltinOutputMode::Fixed,
74 completion_policy: BuiltinCompletionPolicy::Public,
75 errors: &ISGPUARRAY_ERRORS,
76};
77
78const ISGPUARRAY_INTEGER_INPUTS: [BuiltinIntegerInputCapability; 1] =
79 [BuiltinIntegerInputCapability {
80 name: "A",
81 classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
82 availability: BuiltinIntegerInputAvailability::Documented,
83 scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
84 notes: "An explicitly constructed gpuArray may contain any of the eight integer classes; the predicate inspects residency intent without downloading its payload.",
85 }];
86
87pub const ISGPUARRAY_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 1] =
88 [BuiltinIntegerCapabilityDescriptor {
89 form: "tf = isgpuarray(integer_gpuArray)",
90 inputs: &ISGPUARRAY_INTEGER_INPUTS,
91 computation_domain: BuiltinIntegerComputationDomain::Predicate,
92 output_class: BuiltinIntegerOutputClassRule::Logical,
93 overflow: BuiltinIntegerOverflowRule::NotApplicable,
94 backend: BuiltinIntegerBackendRule::HostAndGpu,
95 overload: BuiltinIntegerOverloadKind::ScalarOnly,
96 notes: "The metadata query returns true for explicit integer gpuArray values and false for host integers or RunMat-internal automatic residency; it performs no gather or numeric conversion.",
97 }];
98
99#[runtime_builtin(
100 name = "isgpuarray",
101 category = "logical/tests",
102 summary = "Return true when a value is stored as a gpuArray handle.",
103 keywords = "isgpuarray,gpuarray,gpu,type,logical",
104 accel = "metadata",
105 type_resolver(bool_scalar_type),
106 descriptor(crate::builtins::logical::tests::isgpuarray::ISGPUARRAY_DESCRIPTOR),
107 integer_capabilities(
108 crate::builtins::logical::tests::isgpuarray::ISGPUARRAY_INTEGER_CAPABILITIES
109 ),
110 builtin_path = "crate::builtins::logical::tests::isgpuarray"
111)]
112async fn isgpuarray_builtin(value: Value) -> BuiltinResult<Value> {
113 Ok(Value::Bool(match value {
114 Value::GpuTensor(handle) => runmat_accelerate_api::handle_is_explicit(&handle),
115 _ => false,
116 }))
117}
118
119fn bool_scalar_type(_: &[Type], _context: &ResolveContext) -> Type {
120 Type::Bool
121}
122
123#[cfg(test)]
124pub(crate) mod tests {
125 use super::*;
126 use crate::builtins::common::test_support;
127 use futures::executor::block_on;
128 use runmat_accelerate_api::HostTensorView;
129 use runmat_value::{Tensor, Value};
130
131 fn run_isgpuarray(value: Value) -> BuiltinResult<Value> {
132 block_on(super::isgpuarray_builtin(value))
133 }
134
135 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
136 #[test]
137 fn non_gpu_values_report_false() {
138 assert_eq!(run_isgpuarray(Value::Num(1.0)).unwrap(), Value::Bool(false));
139 }
140
141 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
142 #[test]
143 fn only_explicit_gpuarray_handles_report_true() {
144 test_support::with_test_provider(|provider| {
145 let tensor = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
146 let view = HostTensorView {
147 data: &tensor.materialize_f64(),
148 shape: &tensor.shape,
149 };
150 let handle = provider.upload(&view).expect("upload");
151 let automatic =
152 run_isgpuarray(Value::GpuTensor(handle.clone())).expect("isgpuarray automatic");
153 assert_eq!(automatic, Value::Bool(false));
154
155 let handle =
156 handle.with_provenance(runmat_accelerate_api::GpuHandleProvenance::Explicit);
157 let explicit =
158 run_isgpuarray(Value::GpuTensor(handle.clone())).expect("isgpuarray explicit");
159 assert_eq!(explicit, Value::Bool(true));
160 provider.free(&handle).ok();
161 });
162 }
163}