1use crate::builtins::common::spec::{
4 BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
5 ProviderHook, ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
6};
7use crate::builtins::common::{gpu_helpers, tensor};
8use crate::builtins::math::linalg::type_resolvers::matrix_unary_type;
9use crate::{build_runtime_error, BuiltinResult, RuntimeError};
10
11use num_complex::Complex64;
12use runmat_accelerate_api::{GpuTensorHandle, ProviderLuResult};
13use runmat_builtins::{
14 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinExtensionDescriptor,
15 BuiltinExtensionMode, BuiltinIntegerBackendRule, BuiltinIntegerCapabilityDescriptor,
16 BuiltinIntegerComputationDomain, BuiltinIntegerInputAvailability,
17 BuiltinIntegerInputCapability, BuiltinIntegerOutputClassRule, BuiltinIntegerOverflowRule,
18 BuiltinIntegerOverloadKind, BuiltinIntegerScalarDoubleRule, BuiltinOutputMode,
19 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
20};
21use runmat_macros::runtime_builtin;
22use runmat_value::{ComplexTensor, Tensor, Value};
23
24const BUILTIN_NAME: &str = "lu";
25
26const LU_INTEGER_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
27 id: "lu-integer-input",
28 mode: BuiltinExtensionMode::RunMatOnly,
29 description: "lu with integer input is a RunMat extension",
30 error_identifier: Some("RunMat:compatibility:LuIntegerInputExtension"),
31};
32const LU_LOGICAL_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
33 id: "lu-logical-input",
34 mode: BuiltinExtensionMode::RunMatOnly,
35 description: "lu with logical input is a RunMat extension",
36 error_identifier: Some("RunMat:compatibility:LuLogicalInputExtension"),
37};
38pub const LU_EXTENSIONS: [BuiltinExtensionDescriptor; 2] =
39 [LU_INTEGER_EXTENSION, LU_LOGICAL_EXTENSION];
40const LU_INTEGER_INPUTS: [BuiltinIntegerInputCapability; 1] = [BuiltinIntegerInputCapability {
41 name: "A",
42 classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
43 availability: BuiltinIntegerInputAvailability::RunMatOnly,
44 scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
45 notes: "RunMat mode admits integer matrices at an explicit binary64 factorization boundary.",
46}];
47pub const LU_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 1] =
48 [BuiltinIntegerCapabilityDescriptor {
49 form: "[L,U,P] = lu(integer_A)",
50 inputs: &LU_INTEGER_INPUTS,
51 computation_domain: BuiltinIntegerComputationDomain::FloatingPoint,
52 output_class: BuiltinIntegerOutputClassRule::Double,
53 overflow: BuiltinIntegerOverflowRule::Error,
54 backend: BuiltinIntegerBackendRule::GatherFallback,
55 overload: BuiltinIntegerOverloadKind::FunctionSpecific,
56 notes: "Gated RunMat extension; documented MATLAB input classes remain single and double.",
57 }];
58
59const LU_OUTPUT_COMBINED: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
60 name: "LU",
61 ty: BuiltinParamType::NumericArray,
62 arity: BuiltinParamArity::Required,
63 default: None,
64 description: "Combined LU factors.",
65}];
66
67const LU_OUTPUT_LU: [BuiltinParamDescriptor; 2] = [
68 BuiltinParamDescriptor {
69 name: "L",
70 ty: BuiltinParamType::NumericArray,
71 arity: BuiltinParamArity::Required,
72 default: None,
73 description: "Lower-triangular factor.",
74 },
75 BuiltinParamDescriptor {
76 name: "U",
77 ty: BuiltinParamType::NumericArray,
78 arity: BuiltinParamArity::Required,
79 default: None,
80 description: "Upper-triangular factor.",
81 },
82];
83
84const LU_OUTPUT_LUP: [BuiltinParamDescriptor; 3] = [
85 BuiltinParamDescriptor {
86 name: "L",
87 ty: BuiltinParamType::NumericArray,
88 arity: BuiltinParamArity::Required,
89 default: None,
90 description: "Lower-triangular factor.",
91 },
92 BuiltinParamDescriptor {
93 name: "U",
94 ty: BuiltinParamType::NumericArray,
95 arity: BuiltinParamArity::Required,
96 default: None,
97 description: "Upper-triangular factor.",
98 },
99 BuiltinParamDescriptor {
100 name: "P",
101 ty: BuiltinParamType::NumericArray,
102 arity: BuiltinParamArity::Required,
103 default: None,
104 description: "Permutation matrix or vector based on pivot mode.",
105 },
106];
107
108const LU_INPUTS_A: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
109 name: "A",
110 ty: BuiltinParamType::NumericArray,
111 arity: BuiltinParamArity::Required,
112 default: None,
113 description: "Input matrix to factorize.",
114}];
115
116const LU_INPUTS_A_MODE: [BuiltinParamDescriptor; 2] = [
117 BuiltinParamDescriptor {
118 name: "A",
119 ty: BuiltinParamType::NumericArray,
120 arity: BuiltinParamArity::Required,
121 default: None,
122 description: "Input matrix to factorize.",
123 },
124 BuiltinParamDescriptor {
125 name: "pivotMode",
126 ty: BuiltinParamType::StringScalar,
127 arity: BuiltinParamArity::Required,
128 default: Some("\"matrix\""),
129 description: "Permutation mode (`\"matrix\"` or `\"vector\"`).",
130 },
131];
132
133const LU_SIGNATURES: [BuiltinSignatureDescriptor; 6] = [
134 BuiltinSignatureDescriptor {
135 label: "LU = lu(A)",
136 inputs: &LU_INPUTS_A,
137 outputs: &LU_OUTPUT_COMBINED,
138 },
139 BuiltinSignatureDescriptor {
140 label: "LU = lu(A, pivotMode)",
141 inputs: &LU_INPUTS_A_MODE,
142 outputs: &LU_OUTPUT_COMBINED,
143 },
144 BuiltinSignatureDescriptor {
145 label: "[L, U] = lu(A)",
146 inputs: &LU_INPUTS_A,
147 outputs: &LU_OUTPUT_LU,
148 },
149 BuiltinSignatureDescriptor {
150 label: "[L, U] = lu(A, pivotMode)",
151 inputs: &LU_INPUTS_A_MODE,
152 outputs: &LU_OUTPUT_LU,
153 },
154 BuiltinSignatureDescriptor {
155 label: "[L, U, P] = lu(A)",
156 inputs: &LU_INPUTS_A,
157 outputs: &LU_OUTPUT_LUP,
158 },
159 BuiltinSignatureDescriptor {
160 label: "[L, U, P] = lu(A, pivotMode)",
161 inputs: &LU_INPUTS_A_MODE,
162 outputs: &LU_OUTPUT_LUP,
163 },
164];
165
166const LU_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
167 code: "RM.LU.INVALID_ARGUMENT",
168 identifier: Some("RunMat:lu:InvalidArgument"),
169 when: "Option arguments or requested output count are invalid.",
170 message: "lu currently supports at most three outputs",
171};
172
173const LU_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
174 code: "RM.LU.INVALID_INPUT",
175 identifier: Some("RunMat:lu:InvalidInput"),
176 when: "Input is unsupported for LU factorization.",
177 message: "lu: expected numeric or logical input values",
178};
179
180const LU_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
181 code: "RM.LU.INTERNAL",
182 identifier: Some("RunMat:lu:Internal"),
183 when: "Runtime cannot materialize LU outputs.",
184 message: "lu: internal runtime failure",
185};
186
187const LU_ERRORS: [BuiltinErrorDescriptor; 3] = [
188 LU_ERROR_INVALID_ARGUMENT,
189 LU_ERROR_INVALID_INPUT,
190 LU_ERROR_INTERNAL,
191];
192
193pub const LU_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
194 signatures: &LU_SIGNATURES,
195 output_mode: BuiltinOutputMode::ByRequestedOutputCount,
196 completion_policy: BuiltinCompletionPolicy::Public,
197 errors: &LU_ERRORS,
198};
199
200#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::math::linalg::factor::lu")]
201pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
202 name: "lu",
203 op_kind: GpuOpKind::Custom("lu-factor"),
204 supported_precisions: &[ScalarType::F32, ScalarType::F64],
205 broadcast: BroadcastSemantics::None,
206 provider_hooks: &[ProviderHook::Custom("lu")],
207 constant_strategy: ConstantStrategy::InlineLiteral,
208 residency: ResidencyPolicy::NewHandle,
209 nan_mode: ReductionNaN::Include,
210 two_pass_threshold: None,
211 workgroup_size: None,
212 accepts_nan_mode: false,
213 notes: "Prefers the provider `lu` hook; automatically gathers and falls back to the CPU implementation when no provider support is registered.",
214};
215
216fn lu_error(error: &'static BuiltinErrorDescriptor) -> RuntimeError {
217 lu_error_with_message(error.message, error)
218}
219
220fn lu_error_with_message(
221 message: impl Into<String>,
222 error: &'static BuiltinErrorDescriptor,
223) -> RuntimeError {
224 let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
225 if let Some(identifier) = error.identifier {
226 builder = builder.with_identifier(identifier);
227 }
228 builder.build()
229}
230
231fn lu_invalid_argument(message: impl Into<String>) -> RuntimeError {
232 lu_error_with_message(message, &LU_ERROR_INVALID_ARGUMENT)
233}
234
235fn lu_invalid_input(message: impl Into<String>) -> RuntimeError {
236 lu_error_with_message(message, &LU_ERROR_INVALID_INPUT)
237}
238
239fn lu_internal_error(message: impl Into<String>) -> RuntimeError {
240 lu_error_with_message(message, &LU_ERROR_INTERNAL)
241}
242
243fn with_lu_context(mut error: RuntimeError) -> RuntimeError {
244 if error.message() == "interaction pending..." {
245 return build_runtime_error("interaction pending...")
246 .with_builtin(BUILTIN_NAME)
247 .build();
248 }
249 if error.context.builtin.is_none() {
250 error.context = error.context.with_builtin(BUILTIN_NAME);
251 }
252 error
253}
254
255#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::math::linalg::factor::lu")]
256pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
257 name: "lu",
258 shape: ShapeRequirements::Any,
259 constant_strategy: ConstantStrategy::InlineLiteral,
260 elementwise: None,
261 reduction: None,
262 emits_nan: false,
263 notes: "LU decomposition is not part of expression fusion; calls execute eagerly on the CPU.",
264};
265
266#[runtime_builtin(
267 name = "lu",
268 category = "math/linalg/factor",
269 summary = "Compute LU decompositions with partial pivoting.",
270 keywords = "lu,factorization,decomposition,permutation",
271 accel = "sink",
272 sink = true,
273 type_resolver(matrix_unary_type),
274 descriptor(crate::builtins::math::linalg::factor::lu::LU_DESCRIPTOR),
275 extensions(crate::builtins::math::linalg::factor::lu::LU_EXTENSIONS),
276 integer_capabilities(crate::builtins::math::linalg::factor::lu::LU_INTEGER_CAPABILITIES),
277 builtin_path = "crate::builtins::math::linalg::factor::lu"
278)]
279async fn lu_builtin(value: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
280 let eval = evaluate(value, &rest).await?;
281 if let Some(out_count) = crate::output_count::current_output_count() {
282 if out_count == 0 {
283 return Ok(Value::OutputList(Vec::new()));
284 }
285 if out_count == 1 {
286 return Ok(Value::OutputList(vec![eval.combined()]));
287 }
288 if out_count == 2 {
289 return Ok(Value::OutputList(vec![eval.lower(), eval.upper()]));
290 }
291 if out_count == 3 {
292 return Ok(Value::OutputList(vec![
293 eval.lower(),
294 eval.upper(),
295 eval.permutation(),
296 ]));
297 }
298 return Err(lu_error(&LU_ERROR_INVALID_ARGUMENT));
299 }
300 Ok(eval.combined())
301}
302
303#[derive(Clone)]
305pub struct LuEval {
306 combined: Value,
307 lower: Value,
308 upper: Value,
309 perm_matrix: Value,
310 perm_vector: Value,
311 pivot_mode: PivotMode,
312}
313
314impl LuEval {
315 pub fn combined(&self) -> Value {
317 self.combined.clone()
318 }
319
320 pub fn lower(&self) -> Value {
322 self.lower.clone()
323 }
324
325 pub fn upper(&self) -> Value {
327 self.upper.clone()
328 }
329
330 pub fn permutation(&self) -> Value {
332 match self.pivot_mode {
333 PivotMode::Matrix => self.perm_matrix.clone(),
334 PivotMode::Vector => self.perm_vector.clone(),
335 }
336 }
337
338 pub fn permutation_matrix(&self) -> Value {
340 self.perm_matrix.clone()
341 }
342
343 pub fn pivot_vector(&self) -> Value {
345 self.perm_vector.clone()
346 }
347
348 pub fn pivot_mode(&self) -> PivotMode {
350 self.pivot_mode
351 }
352
353 fn from_components(components: LuComponents, pivot_mode: PivotMode) -> BuiltinResult<Self> {
354 let combined = matrix_to_value(&components.combined)?;
355 let lower = matrix_to_value(&components.lower)?;
356 let upper = matrix_to_value(&components.upper)?;
357 let perm_matrix = matrix_to_value(&components.permutation)?;
358 let perm_vector = pivot_vector_to_value(&components.pivot_vector)?;
359 Ok(Self {
360 combined,
361 lower,
362 upper,
363 perm_matrix,
364 perm_vector,
365 pivot_mode,
366 })
367 }
368
369 fn from_provider(
370 mut result: ProviderLuResult,
371 pivot_mode: PivotMode,
372 provenance: runmat_accelerate_api::GpuHandleProvenance,
373 ) -> Self {
374 for handle in [
375 &mut result.combined,
376 &mut result.lower,
377 &mut result.upper,
378 &mut result.perm_matrix,
379 &mut result.perm_vector,
380 ] {
381 runmat_accelerate_api::set_handle_provenance(handle, provenance);
382 runmat_accelerate_api::mark_residency(handle);
383 }
384 Self {
385 combined: Value::GpuTensor(result.combined),
386 lower: Value::GpuTensor(result.lower),
387 upper: Value::GpuTensor(result.upper),
388 perm_matrix: Value::GpuTensor(result.perm_matrix),
389 perm_vector: Value::GpuTensor(result.perm_vector),
390 pivot_mode,
391 }
392 }
393}
394
395#[derive(Clone, Copy, Debug, PartialEq, Eq)]
397pub enum PivotMode {
398 Matrix,
399 Vector,
400}
401
402impl Default for PivotMode {
403 fn default() -> Self {
404 Self::Matrix
405 }
406}
407
408pub async fn evaluate(value: Value, args: &[Value]) -> BuiltinResult<LuEval> {
410 let pivot_mode = parse_pivot_mode(args)?;
411 ensure_lu_extensions(&value).await?;
412 crate::builtins::common::validation::reject_typed_complex_integer(&value, BUILTIN_NAME)?;
413 match value {
414 Value::GpuTensor(handle) => {
415 if let Some(eval) = evaluate_gpu(&handle, pivot_mode).await? {
416 return Ok(eval);
417 }
418 let owner = gpu_helpers::exact_provider_for_handle(&handle);
419 let explicit = runmat_accelerate_api::handle_is_explicit(&handle);
420 let tensor = gpu_helpers::gather_tensor_async(&handle)
421 .await
422 .map_err(with_lu_context)?;
423 let eval = evaluate_host_value(Value::Tensor(tensor), pivot_mode).await?;
424 if explicit {
425 let owner = owner.ok_or_else(|| {
426 lu_invalid_input("lu: no exact owner for explicit gpuArray input")
427 })?;
428 restore_lu_eval_to_provider(eval, owner)
429 } else {
430 Ok(eval)
431 }
432 }
433 other => evaluate_host_value(other, pivot_mode).await,
434 }
435}
436
437fn restore_lu_eval_to_provider(
438 eval: LuEval,
439 owner: &'static dyn runmat_accelerate_api::AccelProvider,
440) -> BuiltinResult<LuEval> {
441 fn upload(
442 owner: &'static dyn runmat_accelerate_api::AccelProvider,
443 value: &Value,
444 ) -> BuiltinResult<GpuTensorHandle> {
445 let handle = match value {
446 Value::Tensor(tensor) => gpu_helpers::upload_tensor(owner, tensor)
447 .map_err(|error| lu_internal_error(format!("lu: GPU upload failed: {error}")))?,
448 Value::ComplexTensor(tensor) => gpu_helpers::upload_complex_tensor(owner, tensor)?,
449 _ => return Err(lu_internal_error("lu: unexpected host factor value")),
450 };
451 Ok(handle)
452 }
453 let mut uploaded = Vec::with_capacity(5);
454 for value in [
455 &eval.combined,
456 &eval.lower,
457 &eval.upper,
458 &eval.perm_matrix,
459 &eval.perm_vector,
460 ] {
461 match upload(owner, value) {
462 Ok(handle) => uploaded.push(handle),
463 Err(error) => {
464 for handle in &uploaded {
465 gpu_helpers::free_unprotected_exact_owner(handle, &[]);
466 }
467 return Err(error);
468 }
469 }
470 }
471 let mut uploaded = uploaded.into_iter();
472 let mut combined = uploaded.next().expect("combined upload");
473 let mut lower = uploaded.next().expect("lower upload");
474 let mut upper = uploaded.next().expect("upper upload");
475 let mut perm_matrix = uploaded.next().expect("permutation upload");
476 let mut perm_vector = uploaded.next().expect("pivot upload");
477 for handle in [
478 &mut combined,
479 &mut lower,
480 &mut upper,
481 &mut perm_matrix,
482 &mut perm_vector,
483 ] {
484 runmat_accelerate_api::mark_handle_explicit(handle);
485 runmat_accelerate_api::mark_residency(handle);
486 }
487 Ok(LuEval {
488 combined: Value::GpuTensor(combined),
489 lower: Value::GpuTensor(lower),
490 upper: Value::GpuTensor(upper),
491 perm_matrix: Value::GpuTensor(perm_matrix),
492 perm_vector: Value::GpuTensor(perm_vector),
493 pivot_mode: eval.pivot_mode,
494 })
495}
496
497async fn ensure_lu_extensions(value: &Value) -> BuiltinResult<()> {
498 let extension = match value {
499 Value::Int(_) => Some(&LU_INTEGER_EXTENSION),
500 Value::Tensor(tensor) if tensor.integer_storage().is_some() => Some(&LU_INTEGER_EXTENSION),
501 Value::GpuTensor(handle)
502 if runmat_accelerate_api::handle_integer_type(handle).is_some() =>
503 {
504 Some(&LU_INTEGER_EXTENSION)
505 }
506 Value::Bool(_) | Value::LogicalArray(_) => Some(&LU_LOGICAL_EXTENSION),
507 Value::GpuTensor(handle) if runmat_accelerate_api::handle_is_logical(handle) => {
508 Some(&LU_LOGICAL_EXTENSION)
509 }
510 _ => None,
511 };
512 if let Some(extension) = extension {
513 crate::compatibility::ensure_builtin_extension_enabled(extension, BUILTIN_NAME)?;
514 }
515 if crate::builtins::common::validation::value_has_native_integer_class(value)
516 && !crate::builtins::common::validation::native_integer_value_is_exact_f64_async(value)
517 .await?
518 {
519 return Err(lu_invalid_input(
520 "lu: integer input lies outside the exact binary64 interval",
521 ));
522 }
523 Ok(())
524}
525
526async fn evaluate_host_value(value: Value, pivot_mode: PivotMode) -> BuiltinResult<LuEval> {
527 let matrix = extract_matrix(value).await?;
528 let components = lu_factor(matrix)?;
529 LuEval::from_components(components, pivot_mode)
530}
531
532async fn evaluate_gpu(
533 handle: &GpuTensorHandle,
534 pivot_mode: PivotMode,
535) -> BuiltinResult<Option<LuEval>> {
536 if let Some(provider) = gpu_helpers::exact_provider_for_handle(handle) {
537 if let Ok(result) = provider.lu(handle).await {
538 if valid_provider_lu_result(&result, handle, provider) {
539 let provenance = runmat_accelerate_api::handle_provenance(handle)
540 .unwrap_or(runmat_accelerate_api::GpuHandleProvenance::Automatic);
541 return Ok(Some(LuEval::from_provider(result, pivot_mode, provenance)));
542 }
543 free_invalid_provider_lu_result(&result, handle);
544 }
545 }
546 Ok(None)
547}
548
549fn valid_provider_lu_result(
550 result: &ProviderLuResult,
551 input: &GpuTensorHandle,
552 owner: &'static dyn runmat_accelerate_api::AccelProvider,
553) -> bool {
554 let rows = input.shape.first().copied().unwrap_or(1);
555 let outputs = [
556 &result.combined,
557 &result.lower,
558 &result.upper,
559 &result.perm_matrix,
560 &result.perm_vector,
561 ];
562 let expected = [
563 input.shape.clone(),
564 vec![rows, rows],
565 input.shape.clone(),
566 vec![rows, rows],
567 vec![rows, 1],
568 ];
569 outputs.iter().zip(expected.iter()).all(|(output, shape)| {
570 output.shape == *shape
571 && output.device_id == input.device_id
572 && !gpu_helpers::same_gpu_handle(output, input)
573 && runmat_accelerate_api::handle_storage(output)
574 == runmat_accelerate_api::GpuTensorStorage::Real
575 && runmat_accelerate_api::handle_integer_type(output).is_none()
576 && !runmat_accelerate_api::handle_is_logical(output)
577 && runmat_accelerate_api::handle_precision(output)
578 == runmat_accelerate_api::handle_precision(input)
579 && gpu_helpers::exact_provider_for_handle(output)
580 .is_some_and(|candidate| std::ptr::eq(candidate, owner))
581 }) && outputs.iter().enumerate().all(|(index, output)| {
582 outputs
583 .iter()
584 .skip(index + 1)
585 .all(|other| !gpu_helpers::same_gpu_handle(output, other))
586 })
587}
588
589fn free_invalid_provider_lu_result(result: &ProviderLuResult, input: &GpuTensorHandle) {
590 let outputs = [
591 &result.combined,
592 &result.lower,
593 &result.upper,
594 &result.perm_matrix,
595 &result.perm_vector,
596 ];
597 for (index, output) in outputs.iter().enumerate() {
598 if outputs[..index]
599 .iter()
600 .any(|prior| gpu_helpers::same_gpu_handle(output, prior))
601 {
602 continue;
603 }
604 gpu_helpers::free_unprotected_exact_owner(output, &[input]);
605 }
606}
607
608fn parse_pivot_mode(args: &[Value]) -> BuiltinResult<PivotMode> {
609 if args.is_empty() {
610 return Ok(PivotMode::Matrix);
611 }
612 if args.len() > 1 {
613 return Err(lu_invalid_argument("lu: too many option arguments"));
614 }
615 let Some(option) = tensor::value_to_string(&args[0]) else {
616 return Err(lu_invalid_argument(
617 "lu: option must be a string or character vector",
618 ));
619 };
620 match option.trim().to_ascii_lowercase().as_str() {
621 "matrix" => Ok(PivotMode::Matrix),
622 "vector" => Ok(PivotMode::Vector),
623 other => Err(lu_invalid_argument(format!("lu: unknown option '{other}'"))),
624 }
625}
626
627async fn extract_matrix(value: Value) -> BuiltinResult<RowMajorMatrix> {
628 match value {
629 Value::Tensor(t) => RowMajorMatrix::from_tensor(&t),
630 Value::ComplexTensor(ct) => RowMajorMatrix::from_complex_tensor(&ct),
631 Value::GpuTensor(handle) => {
632 let tensor = gpu_helpers::gather_tensor_async(&handle)
633 .await
634 .map_err(with_lu_context)?;
635 RowMajorMatrix::from_tensor(&tensor)
636 }
637 Value::LogicalArray(logical) => {
638 let tensor = tensor::logical_to_tensor(&logical)
639 .map_err(|err| lu_invalid_input(format!("lu: {err}")))?;
640 RowMajorMatrix::from_tensor(&tensor)
641 }
642 Value::Num(n) => Ok(RowMajorMatrix::from_scalar(Complex64::new(n, 0.0))),
643 Value::Int(i) => Ok(RowMajorMatrix::from_scalar(Complex64::new(i.to_f64(), 0.0))),
644 Value::Bool(b) => Ok(RowMajorMatrix::from_scalar(Complex64::new(
645 if b { 1.0 } else { 0.0 },
646 0.0,
647 ))),
648 Value::Complex(re, im) => Ok(RowMajorMatrix::from_scalar(Complex64::new(re, im))),
649 Value::CharArray(_) | Value::String(_) | Value::StringArray(_) => Err(lu_invalid_input(
650 "lu: character data is not supported; convert to numeric values first",
651 )),
652 other => Err(lu_invalid_input(format!(
653 "lu: unsupported input type {:?}",
654 other
655 ))),
656 }
657}
658
659struct LuComponents {
660 combined: RowMajorMatrix,
661 lower: RowMajorMatrix,
662 upper: RowMajorMatrix,
663 permutation: RowMajorMatrix,
664 pivot_vector: Vec<f64>,
665}
666
667fn lu_factor(mut matrix: RowMajorMatrix) -> BuiltinResult<LuComponents> {
668 let rows = matrix.rows;
669 let cols = matrix.cols;
670 let min_dim = rows.min(cols);
671 let mut perm: Vec<usize> = (0..rows).collect();
672
673 for k in 0..min_dim {
674 let mut pivot_row = k;
676 let mut pivot_abs = 0.0;
677 for r in k..rows {
678 let val = matrix.get(r, k);
679 let abs = val.norm();
680 if abs > pivot_abs {
681 pivot_abs = abs;
682 pivot_row = r;
683 }
684 }
685
686 if pivot_row != k {
687 matrix.swap_rows(pivot_row, k);
688 perm.swap(pivot_row, k);
689 }
690
691 if pivot_abs == 0.0 {
692 for r in (k + 1)..rows {
694 matrix.set(r, k, Complex64::new(0.0, 0.0));
695 }
696 continue;
697 }
698
699 let pivot_value = matrix.get(k, k);
700 for r in (k + 1)..rows {
701 let factor = matrix.get(r, k) / pivot_value;
702 matrix.set(r, k, factor);
703 for c in (k + 1)..cols {
704 let updated = matrix.get(r, c) - factor * matrix.get(k, c);
705 matrix.set(r, c, updated);
706 }
707 }
708 }
709
710 let combined = matrix.clone();
711 let lower = build_lower(&matrix);
712 let upper = build_upper(&matrix);
713 let mut permutation = build_permutation(rows, &perm);
714 permutation.single = matrix.single;
715 let pivot_vector: Vec<f64> = perm.iter().map(|idx| (*idx + 1) as f64).collect();
716
717 Ok(LuComponents {
718 combined,
719 lower,
720 upper,
721 permutation,
722 pivot_vector,
723 })
724}
725
726fn build_lower(matrix: &RowMajorMatrix) -> RowMajorMatrix {
727 let rows = matrix.rows;
728 let cols = matrix.cols;
729 let min_dim = rows.min(cols);
730 let mut lower = RowMajorMatrix::identity(rows);
731 lower.single = matrix.single;
732 for i in 0..rows {
733 for j in 0..min_dim {
734 if i > j {
735 lower.set(i, j, matrix.get(i, j));
736 }
737 }
738 }
739 lower
740}
741
742fn build_upper(matrix: &RowMajorMatrix) -> RowMajorMatrix {
743 let rows = matrix.rows;
744 let cols = matrix.cols;
745 let mut upper = RowMajorMatrix::zeros(rows, cols);
746 upper.single = matrix.single;
747 for i in 0..rows {
748 for j in 0..cols {
749 if i <= j {
750 upper.set(i, j, matrix.get(i, j));
751 }
752 }
753 }
754 upper
755}
756
757fn build_permutation(rows: usize, perm: &[usize]) -> RowMajorMatrix {
758 let mut matrix = RowMajorMatrix::zeros(rows, rows);
759 for (i, &col) in perm.iter().enumerate() {
760 if col < rows {
761 matrix.set(i, col, Complex64::new(1.0, 0.0));
762 }
763 }
764 matrix
765}
766
767const EPS: f64 = 1.0e-12;
768
769fn matrix_to_value(matrix: &RowMajorMatrix) -> BuiltinResult<Value> {
770 let mut has_imag = false;
771 for val in &matrix.data {
772 if val.im.abs() > EPS {
773 has_imag = true;
774 break;
775 }
776 }
777 if has_imag {
778 let mut data = Vec::with_capacity(matrix.rows * matrix.cols);
779 for col in 0..matrix.cols {
780 for row in 0..matrix.rows {
781 let idx = row * matrix.cols + col;
782 let v = matrix.data[idx];
783 data.push((v.re, v.im));
784 }
785 }
786 let tensor = if matrix.single {
787 ComplexTensor::from_f32(
788 data.into_iter()
789 .map(|(re, im)| (re as f32, im as f32))
790 .collect(),
791 vec![matrix.rows, matrix.cols],
792 )
793 } else {
794 ComplexTensor::new(data, vec![matrix.rows, matrix.cols])
795 }
796 .map_err(|e| lu_internal_error(format!("lu: {e}")))?;
797 Ok(Value::ComplexTensor(tensor))
798 } else {
799 let mut data = Vec::with_capacity(matrix.rows * matrix.cols);
800 for col in 0..matrix.cols {
801 for row in 0..matrix.rows {
802 let idx = row * matrix.cols + col;
803 data.push(matrix.data[idx].re);
804 }
805 }
806 let tensor = if matrix.single {
807 Tensor::from_f32(
808 data.into_iter().map(|value| value as f32).collect(),
809 vec![matrix.rows, matrix.cols],
810 )
811 } else {
812 Tensor::new(data, vec![matrix.rows, matrix.cols])
813 }
814 .map_err(|e| lu_internal_error(format!("lu: {e}")))?;
815 Ok(Value::Tensor(tensor))
816 }
817}
818
819fn pivot_vector_to_value(pivot: &[f64]) -> BuiltinResult<Value> {
820 let rows = pivot.len();
821 let tensor = Tensor::new(pivot.to_vec(), vec![rows, 1])
822 .map_err(|e| lu_internal_error(format!("lu: {e}")))?;
823 Ok(Value::Tensor(tensor))
824}
825
826#[derive(Clone)]
827struct RowMajorMatrix {
828 rows: usize,
829 cols: usize,
830 data: Vec<Complex64>,
831 single: bool,
832}
833
834impl RowMajorMatrix {
835 fn zeros(rows: usize, cols: usize) -> Self {
836 Self {
837 rows,
838 cols,
839 data: vec![Complex64::new(0.0, 0.0); rows.saturating_mul(cols)],
840 single: false,
841 }
842 }
843
844 fn identity(size: usize) -> Self {
845 let mut matrix = Self::zeros(size, size);
846 for i in 0..size {
847 matrix.set(i, i, Complex64::new(1.0, 0.0));
848 }
849 matrix
850 }
851
852 fn from_scalar(value: Complex64) -> Self {
853 Self {
854 rows: 1,
855 cols: 1,
856 data: vec![value],
857 single: false,
858 }
859 }
860
861 fn from_tensor(tensor: &Tensor) -> BuiltinResult<Self> {
862 if tensor.shape.len() > 2 {
863 return Err(lu_invalid_input("lu: input must be 2-D"));
864 }
865 let rows = tensor.rows();
866 let cols = tensor.cols();
867 let values = tensor::tensor_values_f64_cow(tensor);
868 let mut data = vec![Complex64::new(0.0, 0.0); rows.saturating_mul(cols)];
869 for col in 0..cols {
870 for row in 0..rows {
871 let idx_col_major = row + col * rows;
872 let idx_row_major = row * cols + col;
873 data[idx_row_major] = Complex64::new(values[idx_col_major], 0.0);
874 }
875 }
876 Ok(Self {
877 rows,
878 cols,
879 data,
880 single: tensor.numeric_dtype() == runmat_value::NumericDType::F32,
881 })
882 }
883
884 fn from_complex_tensor(tensor: &ComplexTensor) -> BuiltinResult<Self> {
885 if tensor.shape.len() > 2 {
886 return Err(lu_invalid_input("lu: input must be 2-D"));
887 }
888 let rows = tensor.rows;
889 let cols = tensor.cols;
890 let mut data = vec![Complex64::new(0.0, 0.0); rows.saturating_mul(cols)];
891 for col in 0..cols {
892 for row in 0..rows {
893 let idx_col_major = row + col * rows;
894 let idx_row_major = row * cols + col;
895 let (re, im) = tensor.materialize_f64()[idx_col_major];
896 data[idx_row_major] = Complex64::new(re, im);
897 }
898 }
899 Ok(Self {
900 rows,
901 cols,
902 data,
903 single: tensor.numeric_dtype() == runmat_value::NumericDType::F32,
904 })
905 }
906
907 fn get(&self, row: usize, col: usize) -> Complex64 {
908 self.data[row * self.cols + col]
909 }
910
911 fn set(&mut self, row: usize, col: usize, value: Complex64) {
912 self.data[row * self.cols + col] = value;
913 }
914
915 fn swap_rows(&mut self, r1: usize, r2: usize) {
916 if r1 == r2 {
917 return;
918 }
919 for col in 0..self.cols {
920 self.data.swap(r1 * self.cols + col, r2 * self.cols + col);
921 }
922 }
923}
924
925#[cfg(test)]
926pub(crate) mod tests {
927 use super::*;
928 use crate::builtins::common::test_support;
929 use futures::executor::block_on;
930 #[cfg(feature = "wgpu")]
931 use runmat_accelerate_api::AccelProvider;
932 use runmat_builtins::{ResolveContext, Type};
933 use runmat_value::{ComplexTensor as CMatrix, IntegerStorage, Tensor as Matrix};
934
935 fn error_message(err: RuntimeError) -> String {
936 err.message().to_string()
937 }
938
939 fn tensor_from_value(value: Value) -> Matrix {
940 match value {
941 Value::Tensor(t) => t,
942 other => panic!("expected dense tensor, got {other:?}"),
943 }
944 }
945
946 fn row_major_from_value(value: Value) -> RowMajorMatrix {
947 match value {
948 Value::Tensor(t) => RowMajorMatrix::from_tensor(&t).expect("row-major tensor"),
949 Value::ComplexTensor(ct) => {
950 RowMajorMatrix::from_complex_tensor(&ct).expect("row-major complex tensor")
951 }
952 other => panic!("expected tensor value, got {other:?}"),
953 }
954 }
955
956 #[test]
957 fn lu_type_preserves_matrix_shape() {
958 let out = matrix_unary_type(
959 &[Type::Tensor {
960 shape: Some(vec![Some(2), Some(3)]),
961 }],
962 &ResolveContext::new(Vec::new()),
963 );
964 assert_eq!(
965 out,
966 Type::Tensor {
967 shape: Some(vec![Some(2), Some(3)])
968 }
969 );
970 }
971
972 #[test]
973 fn lu_descriptor_signatures_cover_core_forms() {
974 let labels: Vec<&str> = LU_DESCRIPTOR
975 .signatures
976 .iter()
977 .map(|signature| signature.label)
978 .collect();
979 assert!(labels.contains(&"LU = lu(A)"));
980 assert!(labels.contains(&"LU = lu(A, pivotMode)"));
981 assert!(labels.contains(&"[L, U] = lu(A)"));
982 assert!(labels.contains(&"[L, U] = lu(A, pivotMode)"));
983 assert!(labels.contains(&"[L, U, P] = lu(A)"));
984 assert!(labels.contains(&"[L, U, P] = lu(A, pivotMode)"));
985 }
986
987 #[test]
988 fn lu_descriptor_errors_have_stable_codes() {
989 let codes: Vec<&str> = LU_DESCRIPTOR.errors.iter().map(|err| err.code).collect();
990 assert!(codes.contains(&"RM.LU.INVALID_ARGUMENT"));
991 assert!(codes.contains(&"RM.LU.INVALID_INPUT"));
992 assert!(codes.contains(&"RM.LU.INTERNAL"));
993 }
994
995 #[test]
996 fn lu_matrix_conversion_reads_typed_integer_storage_exactly() {
997 let tensor = Matrix::new_integer(IntegerStorage::I16(vec![4, 6, 3, 3]), vec![2, 2])
998 .expect("typed integer tensor");
999
1000 let matrix = RowMajorMatrix::from_tensor(&tensor).expect("matrix");
1001 assert_eq!(matrix.rows, 2);
1002 assert_eq!(matrix.cols, 2);
1003 assert_eq!(
1004 matrix.data,
1005 vec![
1006 Complex64::new(4.0, 0.0),
1007 Complex64::new(3.0, 0.0),
1008 Complex64::new(6.0, 0.0),
1009 Complex64::new(3.0, 0.0),
1010 ]
1011 );
1012 }
1013
1014 #[test]
1015 fn lu_compatibility_mode_gates_integer_and_logical_extensions() {
1016 let _compat = crate::compatibility::push_runmat_extensions_enabled(false);
1017 let integer = Matrix::new_integer(IntegerStorage::I16(vec![1]), vec![1, 1]).unwrap();
1018 let error = match evaluate(Value::Tensor(integer), &[]) {
1019 Err(error) => error,
1020 Ok(_) => panic!("integer extension must be gated"),
1021 };
1022 assert_eq!(
1023 error.identifier(),
1024 Some("RunMat:compatibility:LuIntegerInputExtension")
1025 );
1026 let error = match evaluate(Value::Bool(true), &[]) {
1027 Err(error) => error,
1028 Ok(_) => panic!("logical extension must be gated"),
1029 };
1030 assert_eq!(
1031 error.identifier(),
1032 Some("RunMat:compatibility:LuLogicalInputExtension")
1033 );
1034 }
1035
1036 #[test]
1037 fn lu_rejects_integer_values_outside_exact_binary64_interval() {
1038 let _extensions = crate::compatibility::push_runmat_extensions_enabled(true);
1039 let integer =
1040 Matrix::new_integer(IntegerStorage::U64(vec![9_007_199_254_740_993]), vec![1, 1])
1041 .unwrap();
1042 let error = match evaluate(Value::Tensor(integer), &[]) {
1043 Err(error) => error,
1044 Ok(_) => panic!("wide integer must reject"),
1045 };
1046 assert_eq!(error.identifier(), LU_ERROR_INVALID_INPUT.identifier);
1047 }
1048
1049 #[test]
1050 fn lu_preserves_single_precision_outputs() {
1051 let input = Matrix::from_f32(vec![2.0, 1.0, 1.0, 2.0], vec![2, 2]).unwrap();
1052 let eval = evaluate(Value::Tensor(input), &[]).expect("single LU");
1053 for output in [
1054 eval.combined(),
1055 eval.lower(),
1056 eval.upper(),
1057 eval.permutation_matrix(),
1058 ] {
1059 let Value::Tensor(tensor) = output else {
1060 panic!("expected real tensor");
1061 };
1062 assert_eq!(tensor.numeric_dtype(), runmat_value::NumericDType::F32);
1063 }
1064 }
1065
1066 #[test]
1067 fn lu_does_not_treat_small_nonzero_pivot_as_zero() {
1068 let input = Matrix::new(vec![1.0e-14, 0.0, 0.0, 2.0e-14], vec![2, 2]).unwrap();
1069 let eval = evaluate(Value::Tensor(input.clone()), &[]).expect("small-scale LU");
1070 let l = tensor_from_value(eval.lower());
1071 let u = tensor_from_value(eval.upper());
1072 let p = tensor_from_value(eval.permutation_matrix());
1073 let pa = crate::builtins::common::matrix::matrix_mul(&p, &input).unwrap();
1074 let product = crate::builtins::common::matrix::matrix_mul(&l, &u).unwrap();
1075 assert_tensor_close(&pa, &product, 1e-28);
1076 }
1077
1078 fn row_major_matmul(a: &RowMajorMatrix, b: &RowMajorMatrix) -> RowMajorMatrix {
1079 assert_eq!(a.cols, b.rows, "incompatible shapes for matmul");
1080 let mut out = RowMajorMatrix::zeros(a.rows, b.cols);
1081 for i in 0..a.rows {
1082 for k in 0..a.cols {
1083 let aik = a.get(i, k);
1084 for j in 0..b.cols {
1085 let acc = out.get(i, j) + aik * b.get(k, j);
1086 out.set(i, j, acc);
1087 }
1088 }
1089 }
1090 out
1091 }
1092
1093 fn assert_tensor_close(a: &Matrix, b: &Matrix, tol: f64) {
1094 assert_eq!(a.shape, b.shape);
1095 for (lhs, rhs) in a.materialize_f64().iter().zip(&b.materialize_f64()) {
1096 assert!(
1097 (lhs - rhs).abs() <= tol,
1098 "mismatch: lhs={lhs}, rhs={rhs}, tol={tol}"
1099 );
1100 }
1101 }
1102
1103 fn assert_row_major_close(a: &RowMajorMatrix, b: &RowMajorMatrix, tol: f64) {
1104 assert_eq!(a.rows, b.rows, "row mismatch");
1105 assert_eq!(a.cols, b.cols, "col mismatch");
1106 for row in 0..a.rows {
1107 for col in 0..a.cols {
1108 let lhs = a.get(row, col);
1109 let rhs = b.get(row, col);
1110 let diff = (lhs - rhs).norm();
1111 assert!(
1112 diff <= tol,
1113 "mismatch at ({row}, {col}): lhs={lhs:?}, rhs={rhs:?}, diff={diff}, tol={tol}"
1114 );
1115 }
1116 }
1117 }
1118
1119 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1120 #[test]
1121 fn lu_single_output_produces_combined_matrix() {
1122 let a = Matrix::new(
1123 vec![2.0, 4.0, -2.0, 1.0, -6.0, 7.0, 1.0, 0.0, 2.0],
1124 vec![3, 3],
1125 )
1126 .unwrap();
1127 let result = lu_builtin(Value::Tensor(a.clone()), Vec::new()).expect("lu");
1128 let lu = tensor_from_value(result);
1129 let eval = evaluate(Value::Tensor(a), &[]).expect("evaluate");
1130 let expected = tensor_from_value(eval.combined());
1131 assert_tensor_close(&lu, &expected, 1e-12);
1132 }
1133
1134 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1135 #[test]
1136 fn lu_three_outputs_matches_factorization() {
1137 let data = vec![2.0, 4.0, -2.0, 1.0, -6.0, 7.0, 1.0, 0.0, 2.0];
1138 let a = Matrix::new(data.clone(), vec![3, 3]).unwrap();
1139 let eval = evaluate(Value::Tensor(a.clone()), &[]).expect("evaluate");
1140 let l = tensor_from_value(eval.lower());
1141 let u = tensor_from_value(eval.upper());
1142 let p = tensor_from_value(eval.permutation_matrix());
1143
1144 let pa = crate::builtins::common::matrix::matrix_mul(&p, &a).expect("P*A");
1145 let lu_product = crate::builtins::common::matrix::matrix_mul(&l, &u).expect("L*U");
1146 assert_tensor_close(&pa, &lu_product, 1e-9);
1147 }
1148
1149 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1150 #[test]
1151 fn lu_complex_matrix_factorization() {
1152 let data = vec![(1.0, 2.0), (3.0, -1.0), (2.0, -1.0), (4.0, 2.0)];
1153 let a = CMatrix::new(data.clone(), vec![2, 2]).expect("complex tensor");
1154 let eval = evaluate(Value::ComplexTensor(a.clone()), &[]).expect("evaluate complex");
1155
1156 let l = row_major_from_value(eval.lower());
1157 let u = row_major_from_value(eval.upper());
1158 let p = row_major_from_value(eval.permutation_matrix());
1159 let input = RowMajorMatrix::from_complex_tensor(&a).expect("row-major input");
1160
1161 let pa = row_major_matmul(&p, &input);
1162 let lu = row_major_matmul(&l, &u);
1163 assert_row_major_close(&pa, &lu, 1e-9);
1164 }
1165
1166 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1167 #[test]
1168 fn lu_handles_singular_matrix() {
1169 let a = Matrix::new(vec![0.0, 0.0, 0.0, 0.0], vec![2, 2]).unwrap();
1170 let eval = evaluate(Value::Tensor(a.clone()), &[]).expect("evaluate singular");
1171 let l = tensor_from_value(eval.lower());
1172 let u = tensor_from_value(eval.upper());
1173 let p = tensor_from_value(eval.permutation_matrix());
1174
1175 assert!(u.materialize_f64().iter().any(|&v| v.abs() <= 1e-12));
1176
1177 let pa = crate::builtins::common::matrix::matrix_mul(&p, &a).expect("P*A");
1178 let lu_product = crate::builtins::common::matrix::matrix_mul(&l, &u).expect("L*U");
1179 assert_tensor_close(&pa, &lu_product, 1e-9);
1180 }
1181
1182 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1183 #[test]
1184 fn lu_vector_option_returns_pivot_vector() {
1185 let a = Matrix::new(vec![4.0, 6.0, 3.0, 3.0], vec![2, 2]).unwrap();
1186 let eval =
1187 evaluate(Value::Tensor(a), &[Value::from("vector")]).expect("evaluate vector mode");
1188 assert_eq!(eval.pivot_mode(), PivotMode::Vector);
1189 let pivot = tensor_from_value(eval.pivot_vector());
1190 assert_eq!(pivot.shape, vec![2, 1]);
1191 assert_eq!(pivot.materialize_f64(), vec![2.0, 1.0]);
1192 }
1193
1194 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1195 #[test]
1196 fn lu_vector_option_case_insensitive() {
1197 let a = Matrix::new(vec![4.0, 6.0, 3.0, 3.0], vec![2, 2]).unwrap();
1198 let eval =
1199 evaluate(Value::Tensor(a), &[Value::from("VECTOR")]).expect("evaluate vector option");
1200 assert_eq!(eval.pivot_mode(), PivotMode::Vector);
1201 }
1202
1203 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1204 #[test]
1205 fn lu_matrix_option_returns_permutation_matrix() {
1206 let a = Matrix::new(vec![2.0, 1.0, 3.0, 4.0], vec![2, 2]).unwrap();
1207 let eval =
1208 evaluate(Value::Tensor(a), &[Value::from("matrix")]).expect("evaluate matrix option");
1209 assert_eq!(eval.pivot_mode(), PivotMode::Matrix);
1210 let perm_selected = tensor_from_value(eval.permutation());
1211 let perm_matrix = tensor_from_value(eval.permutation_matrix());
1212 assert_eq!(perm_selected.shape, perm_matrix.shape);
1213 assert_tensor_close(&perm_selected, &perm_matrix, 1e-12);
1214 }
1215
1216 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1217 #[test]
1218 fn lu_handles_rectangular_matrices() {
1219 let a = Matrix::new(vec![3.0, 6.0, 1.0, 3.0, 2.0, 4.0], vec![2, 3]).unwrap();
1220 let eval = evaluate(Value::Tensor(a.clone()), &[]).expect("evaluate rectangular");
1221 let l = tensor_from_value(eval.lower());
1222 let u = tensor_from_value(eval.upper());
1223 let p = tensor_from_value(eval.permutation_matrix());
1224 assert_eq!(l.shape, vec![2, 2]);
1225 assert_eq!(u.shape, vec![2, 3]);
1226 assert_eq!(p.shape, vec![2, 2]);
1227
1228 let pa = crate::builtins::common::matrix::matrix_mul(&p, &a).expect("P*A");
1229 let lu_product = crate::builtins::common::matrix::matrix_mul(&l, &u).expect("L*U");
1230 assert_tensor_close(&pa, &lu_product, 1e-9);
1231 }
1232
1233 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1234 #[test]
1235 fn lu_rejects_unknown_option() {
1236 let a = Matrix::new(vec![1.0], vec![1, 1]).unwrap();
1237 let err = match evaluate(Value::Tensor(a), &[Value::from("invalid")]) {
1238 Ok(_) => panic!("expected option parse failure"),
1239 Err(err) => {
1240 assert_eq!(err.identifier(), LU_ERROR_INVALID_ARGUMENT.identifier);
1241 error_message(err)
1242 }
1243 };
1244 assert!(err.contains("unknown option"));
1245 }
1246
1247 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1248 #[test]
1249 fn lu_rejects_non_string_option() {
1250 let a = Matrix::new(vec![1.0], vec![1, 1]).unwrap();
1251 let err = match evaluate(Value::Tensor(a), &[Value::Num(2.0)]) {
1252 Ok(_) => panic!("expected option parse failure"),
1253 Err(err) => {
1254 assert_eq!(err.identifier(), LU_ERROR_INVALID_ARGUMENT.identifier);
1255 error_message(err)
1256 }
1257 };
1258 assert!(err.contains("unknown option"));
1259 }
1260
1261 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1262 #[test]
1263 fn lu_rejects_multiple_options() {
1264 let a = Matrix::new(vec![1.0], vec![1, 1]).unwrap();
1265 let err = match evaluate(
1266 Value::Tensor(a),
1267 &[Value::from("matrix"), Value::from("vector")],
1268 ) {
1269 Ok(_) => panic!("expected option arity failure"),
1270 Err(err) => {
1271 assert_eq!(err.identifier(), LU_ERROR_INVALID_ARGUMENT.identifier);
1272 error_message(err)
1273 }
1274 };
1275 assert!(err.contains("too many option arguments"));
1276 }
1277
1278 #[test]
1279 fn lu_invalid_input_identifier_is_stable() {
1280 let tensor = Matrix::new(vec![1.0, 2.0, 3.0, 4.0], vec![1, 2, 2]).expect("tensor");
1281 let err = match evaluate(Value::Tensor(tensor), &[]) {
1282 Ok(_) => panic!("expected 2-D input failure"),
1283 Err(err) => err,
1284 };
1285 assert_eq!(err.identifier(), LU_ERROR_INVALID_INPUT.identifier);
1286 }
1287
1288 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1289 #[test]
1290 fn lu_gpu_provider_roundtrip() {
1291 test_support::with_test_provider(|provider| {
1292 let host = Matrix::new(vec![10.0, 3.0, 7.0, 2.0], vec![2, 2]).unwrap();
1293 let view = runmat_accelerate_api::HostTensorView {
1294 data: &host.materialize_f64(),
1295 shape: &host.shape,
1296 };
1297 let handle = provider.upload(&view).expect("upload");
1298 let eval = evaluate(Value::GpuTensor(handle.clone()), &[]).expect("evaluate gpu input");
1299 let lower_val = eval.lower();
1300 let upper_val = eval.upper();
1301 let perm_val = eval.permutation_matrix();
1302 assert!(matches!(lower_val, Value::GpuTensor(_)));
1303 assert!(matches!(upper_val, Value::GpuTensor(_)));
1304 assert!(matches!(perm_val, Value::GpuTensor(_)));
1305 let l = test_support::gather(lower_val).expect("gather lower");
1306 let u = test_support::gather(upper_val).expect("gather upper");
1307 let p = test_support::gather(perm_val).expect("gather permutation");
1308 let pa = crate::builtins::common::matrix::matrix_mul(&p, &host).expect("P*A");
1309 let lu_product = crate::builtins::common::matrix::matrix_mul(&l, &u).expect("L*U");
1310 assert_tensor_close(&pa, &lu_product, 1e-9);
1311 });
1312 }
1313
1314 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1315 #[test]
1316 fn lu_gpu_vector_option_roundtrip() {
1317 test_support::with_test_provider(|provider| {
1318 let host = Matrix::new(vec![4.0, 6.0, 3.0, 3.0], vec![2, 2]).unwrap();
1319 let view = runmat_accelerate_api::HostTensorView {
1320 data: &host.materialize_f64(),
1321 shape: &host.shape,
1322 };
1323 let handle = provider.upload(&view).expect("upload");
1324 let eval =
1325 evaluate(Value::GpuTensor(handle), &[Value::from("vector")]).expect("gpu vector");
1326 let pivot_val = eval.permutation();
1327 assert!(matches!(pivot_val, Value::GpuTensor(_)));
1328 let pivot = test_support::gather(pivot_val).expect("gather pivot");
1329 assert_eq!(pivot.shape, vec![2, 1]);
1330 let expected = Matrix::new(vec![2.0, 1.0], vec![2, 1]).unwrap();
1331 assert_tensor_close(&pivot, &expected, 1e-12);
1332 });
1333 }
1334
1335 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1336 #[test]
1337 fn lu_accepts_scalar_inputs() {
1338 let eval = evaluate(Value::Num(5.0), &[]).expect("evaluate scalar");
1339 let l = tensor_from_value(eval.lower());
1340 let u = tensor_from_value(eval.upper());
1341 let p = tensor_from_value(eval.permutation_matrix());
1342 assert_eq!(l.materialize_f64(), vec![1.0]);
1343 assert_eq!(u.materialize_f64(), vec![5.0]);
1344 assert_eq!(p.materialize_f64(), vec![1.0]);
1345 }
1346
1347 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1348 #[test]
1349 #[cfg(feature = "wgpu")]
1350 fn lu_wgpu_matches_cpu() {
1351 let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
1352 runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
1353 ) else {
1354 return;
1355 };
1356 let host = Matrix::new(
1357 vec![2.0, 4.0, -2.0, 1.0, -6.0, 7.0, 1.0, 0.0, 2.0],
1358 vec![3, 3],
1359 )
1360 .unwrap();
1361 let cpu_eval = evaluate(Value::Tensor(host.clone()), &[]).expect("cpu evaluate");
1362 let view = runmat_accelerate_api::HostTensorView {
1363 data: &host.materialize_f64(),
1364 shape: &host.shape,
1365 };
1366 let handle = provider.upload(&view).expect("upload");
1367 let gpu_eval = evaluate(Value::GpuTensor(handle), &[]).expect("gpu evaluate");
1368
1369 let l_cpu = tensor_from_value(cpu_eval.lower());
1370 let u_cpu = tensor_from_value(cpu_eval.upper());
1371 let p_cpu = tensor_from_value(cpu_eval.permutation_matrix());
1372 let lu_cpu = tensor_from_value(cpu_eval.combined());
1373
1374 let l_gpu = test_support::gather(gpu_eval.lower()).expect("gather L");
1375 let u_gpu = test_support::gather(gpu_eval.upper()).expect("gather U");
1376 let p_gpu = test_support::gather(gpu_eval.permutation_matrix()).expect("gather P");
1377 let lu_gpu = test_support::gather(gpu_eval.combined()).expect("gather LU");
1378
1379 assert_tensor_close(&l_cpu, &l_gpu, 1e-12);
1380 assert_tensor_close(&u_cpu, &u_gpu, 1e-12);
1381 assert_tensor_close(&p_cpu, &p_gpu, 1e-12);
1382 assert_tensor_close(&lu_cpu, &lu_gpu, 1e-12);
1383
1384 let pivot_cpu = tensor_from_value(cpu_eval.pivot_vector());
1385 let pivot_gpu = test_support::gather(gpu_eval.pivot_vector()).expect("gather pivot vector");
1386 assert_tensor_close(&pivot_cpu, &pivot_gpu, 1e-12);
1387
1388 let handle_vector = provider.upload(&view).expect("upload vector option");
1389 let gpu_vector_eval = evaluate(Value::GpuTensor(handle_vector), &[Value::from("vector")])
1390 .expect("gpu vector evaluate");
1391 let pivot_vector =
1392 test_support::gather(gpu_vector_eval.permutation()).expect("gather vector pivot");
1393 assert_tensor_close(&pivot_cpu, &pivot_vector, 1e-12);
1394 }
1395
1396 fn lu_builtin(value: Value, rest: Vec<Value>) -> BuiltinResult<Value> {
1397 block_on(super::lu_builtin(value, rest))
1398 }
1399
1400 fn evaluate(value: Value, args: &[Value]) -> BuiltinResult<LuEval> {
1401 block_on(super::evaluate(value, args))
1402 }
1403}