Skip to main content

runmat_runtime/builtins/math/linalg/factor/
lu.rs

1//! MATLAB-compatible `lu` builtin with CPU-backed semantics.
2
3use 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/// Output form for `lu`, reused by both the builtin wrapper and the VM multi-output path.
304#[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    /// Combined LU factor (single-output form).
316    pub fn combined(&self) -> Value {
317        self.combined.clone()
318    }
319
320    /// Lower-triangular factor.
321    pub fn lower(&self) -> Value {
322        self.lower.clone()
323    }
324
325    /// Upper-triangular factor.
326    pub fn upper(&self) -> Value {
327        self.upper.clone()
328    }
329
330    /// Permutation value respecting the selected pivot mode.
331    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    /// Permutation matrix (always available, useful for tests).
339    pub fn permutation_matrix(&self) -> Value {
340        self.perm_matrix.clone()
341    }
342
343    /// Pivot vector (always available, useful for tests).
344    pub fn pivot_vector(&self) -> Value {
345        self.perm_vector.clone()
346    }
347
348    /// The pivot mode that was requested.
349    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/// Permutation output mode.
396#[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
408/// Evaluate `lu` while preserving all output forms for later extraction.
409pub 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        // Select pivot row with maximal absolute value in column k.
675        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            // Entire column is effectively zero; set multipliers to zero and continue.
693            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}