Skip to main content

runmat_runtime/builtins/math/linalg/solve/
linsolve.rs

1//! MATLAB-compatible `linsolve` builtin with structural hints and GPU-aware fallbacks.
2
3use nalgebra::{linalg::SVD, DMatrix};
4use num_complex::Complex64;
5use runmat_accelerate_api::{
6    AccelProvider, GpuTensorHandle, HostTensorView, ProviderLinsolveOptions, ProviderLinsolveResult,
7};
8use runmat_builtins::{
9    BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinExtensionDescriptor,
10    BuiltinExtensionMode, BuiltinIntegerBackendRule, BuiltinIntegerCapabilityDescriptor,
11    BuiltinIntegerComputationDomain, BuiltinIntegerInputAvailability,
12    BuiltinIntegerInputCapability, BuiltinIntegerOutputClassRule, BuiltinIntegerOverflowRule,
13    BuiltinIntegerOverloadKind, BuiltinIntegerScalarDoubleRule, BuiltinOutputMode,
14    BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
15};
16use runmat_macros::runtime_builtin;
17use runmat_value::{
18    ComplexStorage, ComplexTensor, IntValue, IntegerComplexStorage, NumericDType, Tensor, Value,
19};
20
21use crate::builtins::common::spec::{
22    BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
23    ProviderHook, ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
24};
25use crate::builtins::common::{
26    gpu_helpers,
27    linalg::{diagonal_rcond, singular_value_rcond},
28    tensor,
29};
30use crate::builtins::math::elementwise::conj::conjugate_integer_imaginary_storage;
31use crate::builtins::math::linalg::type_resolvers::left_divide_type;
32use crate::{build_runtime_error, BuiltinResult, RuntimeError};
33
34const NAME: &str = "linsolve";
35
36const LINSOLVE_OUTPUT_X: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
37    name: "X",
38    ty: BuiltinParamType::NumericArray,
39    arity: BuiltinParamArity::Required,
40    default: None,
41    description: "Solution to A * X = B.",
42}];
43
44const LINSOLVE_OUTPUT_XR: [BuiltinParamDescriptor; 2] = [
45    BuiltinParamDescriptor {
46        name: "X",
47        ty: BuiltinParamType::NumericArray,
48        arity: BuiltinParamArity::Required,
49        default: None,
50        description: "Solution to A * X = B.",
51    },
52    BuiltinParamDescriptor {
53        name: "R",
54        ty: BuiltinParamType::NumericScalar,
55        arity: BuiltinParamArity::Required,
56        default: None,
57        description: "Reciprocal condition estimate.",
58    },
59];
60
61const LINSOLVE_INPUTS_AB: [BuiltinParamDescriptor; 2] = [
62    BuiltinParamDescriptor {
63        name: "A",
64        ty: BuiltinParamType::Any,
65        arity: BuiltinParamArity::Required,
66        default: None,
67        description: "Coefficient matrix.",
68    },
69    BuiltinParamDescriptor {
70        name: "B",
71        ty: BuiltinParamType::Any,
72        arity: BuiltinParamArity::Required,
73        default: None,
74        description: "Right-hand side matrix or vector.",
75    },
76];
77
78const LINSOLVE_INPUTS_AB_OPTS: [BuiltinParamDescriptor; 3] = [
79    BuiltinParamDescriptor {
80        name: "A",
81        ty: BuiltinParamType::Any,
82        arity: BuiltinParamArity::Required,
83        default: None,
84        description: "Coefficient matrix.",
85    },
86    BuiltinParamDescriptor {
87        name: "B",
88        ty: BuiltinParamType::Any,
89        arity: BuiltinParamArity::Required,
90        default: None,
91        description: "Right-hand side matrix or vector.",
92    },
93    BuiltinParamDescriptor {
94        name: "opts",
95        ty: BuiltinParamType::Any,
96        arity: BuiltinParamArity::Optional,
97        default: None,
98        description: "Structural options (LT, UT, RECT, SYM, POSDEF, TRANSA, RCOND).",
99    },
100];
101
102const LINSOLVE_SIGNATURES: [BuiltinSignatureDescriptor; 4] = [
103    BuiltinSignatureDescriptor {
104        label: "X = linsolve(A, B)",
105        inputs: &LINSOLVE_INPUTS_AB,
106        outputs: &LINSOLVE_OUTPUT_X,
107    },
108    BuiltinSignatureDescriptor {
109        label: "X = linsolve(A, B, opts)",
110        inputs: &LINSOLVE_INPUTS_AB_OPTS,
111        outputs: &LINSOLVE_OUTPUT_X,
112    },
113    BuiltinSignatureDescriptor {
114        label: "[X, R] = linsolve(A, B)",
115        inputs: &LINSOLVE_INPUTS_AB,
116        outputs: &LINSOLVE_OUTPUT_XR,
117    },
118    BuiltinSignatureDescriptor {
119        label: "[X, R] = linsolve(A, B, opts)",
120        inputs: &LINSOLVE_INPUTS_AB_OPTS,
121        outputs: &LINSOLVE_OUTPUT_XR,
122    },
123];
124
125const LINSOLVE_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
126    code: "RM.LINSOLVE.INVALID_ARGUMENT",
127    identifier: Some("RunMat:linsolve:InvalidArgument"),
128    when: "Options/output count/auxiliary arguments are malformed or unsupported.",
129    message: "linsolve: invalid argument",
130};
131
132const LINSOLVE_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
133    code: "RM.LINSOLVE.INVALID_INPUT",
134    identifier: Some("RunMat:linsolve:InvalidInput"),
135    when: "Input shape/type cannot be solved under linsolve semantics.",
136    message: "linsolve: invalid input",
137};
138
139const LINSOLVE_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
140    code: "RM.LINSOLVE.INTERNAL",
141    identifier: Some("RunMat:linsolve:Internal"),
142    when: "Runtime fails while solving or executing provider fallback paths.",
143    message: "linsolve: internal runtime failure",
144};
145
146const LINSOLVE_ERRORS: [BuiltinErrorDescriptor; 3] = [
147    LINSOLVE_ERROR_INVALID_ARGUMENT,
148    LINSOLVE_ERROR_INVALID_INPUT,
149    LINSOLVE_ERROR_INTERNAL,
150];
151const LINSOLVE_INTEGER_INPUT_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
152    id: "linsolve-integer-input",
153    mode: BuiltinExtensionMode::RunMatOnly,
154    description: "linsolve with integer A or B is a RunMat extension",
155    error_identifier: Some("RunMat:compatibility:LinsolveIntegerInputExtension"),
156};
157const LINSOLVE_LOGICAL_INPUT_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
158    id: "linsolve-logical-input",
159    mode: BuiltinExtensionMode::RunMatOnly,
160    description: "linsolve with logical A or B is a RunMat extension",
161    error_identifier: Some("RunMat:compatibility:LinsolveLogicalInputExtension"),
162};
163const LINSOLVE_EXPLICIT_GPU_TWO_OUTPUT_EXTENSION: BuiltinExtensionDescriptor =
164    BuiltinExtensionDescriptor {
165        id: "linsolve-explicit-gpu-two-output",
166        mode: BuiltinExtensionMode::RunMatOnly,
167        description: "two-output linsolve with explicit gpuArray input is a RunMat extension",
168        error_identifier: Some("RunMat:compatibility:LinsolveExplicitGpuTwoOutputExtension"),
169    };
170const LINSOLVE_INTEGER_OPTION_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
171    id: "linsolve-integer-option-control",
172    mode: BuiltinExtensionMode::RunMatOnly,
173    description: "linsolve with a typed-integer structural option is a RunMat extension",
174    error_identifier: Some("RunMat:compatibility:LinsolveIntegerOptionExtension"),
175};
176const LINSOLVE_TEXT_TRANSA_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
177    id: "linsolve-text-transa-option",
178    mode: BuiltinExtensionMode::RunMatOnly,
179    description: "linsolve with a text-valued TRANSA option is a RunMat extension",
180    error_identifier: Some("RunMat:compatibility:LinsolveTextTransaExtension"),
181};
182const LINSOLVE_RCOND_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
183    id: "linsolve-rcond-option",
184    mode: BuiltinExtensionMode::RunMatOnly,
185    description: "linsolve with the RCOND option is a RunMat extension",
186    error_identifier: Some("RunMat:compatibility:LinsolveRcondExtension"),
187};
188pub const LINSOLVE_EXTENSIONS: [BuiltinExtensionDescriptor; 6] = [
189    LINSOLVE_INTEGER_INPUT_EXTENSION,
190    LINSOLVE_LOGICAL_INPUT_EXTENSION,
191    LINSOLVE_EXPLICIT_GPU_TWO_OUTPUT_EXTENSION,
192    LINSOLVE_INTEGER_OPTION_EXTENSION,
193    LINSOLVE_TEXT_TRANSA_EXTENSION,
194    LINSOLVE_RCOND_EXTENSION,
195];
196const LINSOLVE_INTEGER_INPUTS: [BuiltinIntegerInputCapability; 2] = [
197    BuiltinIntegerInputCapability {
198        name: "A",
199        classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
200        availability: BuiltinIntegerInputAvailability::RunMatOnly,
201        scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
202        notes: "RunMat-only exact-owner promotion into an exact binary64 boundary.",
203    },
204    BuiltinIntegerInputCapability {
205        name: "B",
206        classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
207        availability: BuiltinIntegerInputAvailability::RunMatOnly,
208        scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
209        notes: "RunMat-only exact-owner promotion into an exact binary64 boundary.",
210    },
211];
212const LINSOLVE_INTEGER_OPTION_INPUTS: [BuiltinIntegerInputCapability; 1] =
213    [BuiltinIntegerInputCapability {
214        name: "opts structural fields",
215        classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
216        availability: BuiltinIntegerInputAvailability::RunMatOnly,
217        scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
218        notes: "Typed-integer LT, UT, RECT, SYM, and POSDEF values are independently gated structural controls.",
219    }];
220pub const INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 2] = [
221    BuiltinIntegerCapabilityDescriptor { form: "X = linsolve(integer_A, integer_B)", inputs: &LINSOLVE_INTEGER_INPUTS, computation_domain: BuiltinIntegerComputationDomain::FloatingPoint, output_class: BuiltinIntegerOutputClassRule::Double, overflow: BuiltinIntegerOverflowRule::NotApplicable, backend: BuiltinIntegerBackendRule::GatherFallback, overload: BuiltinIntegerOverloadKind::Multiple, notes: "Integer operands are a gated RunMat extension and are promoted only when exactly representable in binary64." },
222    BuiltinIntegerCapabilityDescriptor { form: "X = linsolve(A, B, opts_with_integer_field)", inputs: &LINSOLVE_INTEGER_OPTION_INPUTS, computation_domain: BuiltinIntegerComputationDomain::Structural, output_class: BuiltinIntegerOutputClassRule::NotApplicable, overflow: BuiltinIntegerOverflowRule::NotApplicable, backend: BuiltinIntegerBackendRule::HostOnly, overload: BuiltinIntegerOverloadKind::FunctionSpecific, notes: "RunMat-only typed-integer truth controls are classified before option coercion." },
223];
224
225pub const LINSOLVE_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
226    signatures: &LINSOLVE_SIGNATURES,
227    output_mode: BuiltinOutputMode::ByRequestedOutputCount,
228    completion_policy: BuiltinCompletionPolicy::Public,
229    errors: &LINSOLVE_ERRORS,
230};
231
232#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::math::linalg::solve::linsolve")]
233pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
234    name: "linsolve",
235    op_kind: GpuOpKind::Custom("solve"),
236    supported_precisions: &[ScalarType::F32, ScalarType::F64],
237    broadcast: BroadcastSemantics::None,
238    provider_hooks: &[ProviderHook::Custom("linsolve")],
239    constant_strategy: ConstantStrategy::UniformBuffer,
240    residency: ResidencyPolicy::NewHandle,
241    nan_mode: ReductionNaN::Include,
242    two_pass_threshold: None,
243    workgroup_size: None,
244    accepts_nan_mode: false,
245    notes: "Prefers the provider linsolve hook; WGPU currently supports triangular solves, real F32 TRANSA='T'/'C' variants, a dedicated real F32 POSDEF/Cholesky path, and selected real F32 QR-backed square and rectangular solves, otherwise it gathers to the host solver and re-uploads the result.",
246};
247
248fn linsolve_error_with_message(
249    message: impl Into<String>,
250    error: &'static BuiltinErrorDescriptor,
251) -> RuntimeError {
252    let mut builder = build_runtime_error(message).with_builtin(NAME);
253    if let Some(identifier) = error.identifier {
254        builder = builder.with_identifier(identifier);
255    }
256    builder.build()
257}
258
259fn builtin_error(message: impl Into<String>) -> RuntimeError {
260    linsolve_error_with_message(message, &LINSOLVE_ERROR_INVALID_INPUT)
261}
262
263fn argument_error(message: impl Into<String>) -> RuntimeError {
264    linsolve_error_with_message(message, &LINSOLVE_ERROR_INVALID_ARGUMENT)
265}
266
267fn map_control_flow(err: RuntimeError) -> RuntimeError {
268    let mut builder = build_runtime_error(err.message()).with_builtin(NAME);
269    if let Some(identifier) = err.identifier() {
270        builder = builder.with_identifier(identifier.to_string());
271    }
272    if let Some(task_id) = err.context.task_id.clone() {
273        builder = builder.with_task_id(task_id);
274    }
275    if !err.context.call_stack.is_empty() {
276        builder = builder.with_call_stack(err.context.call_stack.clone());
277    }
278    if let Some(phase) = err.context.phase.clone() {
279        builder = builder.with_phase(phase);
280    }
281    builder.with_source(err).build()
282}
283
284#[runmat_macros::register_fusion_spec(
285    builtin_path = "crate::builtins::math::linalg::solve::linsolve"
286)]
287pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
288    name: "linsolve",
289    shape: ShapeRequirements::Any,
290    constant_strategy: ConstantStrategy::UniformBuffer,
291    elementwise: None,
292    reduction: None,
293    emits_nan: false,
294    notes: "Linear solves are terminal operations and do not fuse with surrounding kernels.",
295};
296
297#[runtime_builtin(
298    name = "linsolve",
299    category = "math/linalg/solve",
300    summary = "Solve A * X = B with structural hints such as LT, UT, POSDEF, or TRANSA.",
301    keywords = "linsolve,linear system,triangular,gpu",
302    accel = "linsolve",
303    type_resolver(left_divide_type),
304    descriptor(crate::builtins::math::linalg::solve::linsolve::LINSOLVE_DESCRIPTOR),
305    extensions(LINSOLVE_EXTENSIONS),
306    integer_capabilities(crate::builtins::math::linalg::solve::linsolve::INTEGER_CAPABILITIES),
307    builtin_path = "crate::builtins::math::linalg::solve::linsolve"
308)]
309async fn linsolve_builtin(lhs: Value, rhs: Value, rest: Vec<Value>) -> BuiltinResult<Value> {
310    let eval = evaluate_args(lhs, rhs, &rest).await?;
311    if let Some(out_count) = crate::output_count::current_output_count() {
312        if out_count == 0 {
313            return Ok(Value::OutputList(Vec::new()));
314        }
315        if out_count == 1 {
316            return Ok(Value::OutputList(vec![eval.solution()]));
317        }
318        if out_count == 2 {
319            return Ok(Value::OutputList(vec![
320                eval.solution(),
321                eval.reciprocal_condition(),
322            ]));
323        }
324        return Err(argument_error(
325            "linsolve currently supports at most two outputs",
326        ));
327    }
328    Ok(eval.solution())
329}
330
331/// Evaluate `linsolve`, returning both the solution and the estimated reciprocal condition number.
332pub async fn evaluate(
333    lhs: Value,
334    rhs: Value,
335    options: SolveOptions,
336) -> BuiltinResult<LinsolveEval> {
337    if let Some(eval) = try_gpu_linsolve(&lhs, &rhs, &options).await? {
338        return Ok(eval);
339    }
340
341    let lhs_host = crate::dispatcher::gather_if_needed_async(&lhs)
342        .await
343        .map_err(map_control_flow)?;
344    let rhs_host = crate::dispatcher::gather_if_needed_async(&rhs)
345        .await
346        .map_err(map_control_flow)?;
347    let pair = coerce_numeric_pair(lhs_host, rhs_host).await?;
348    match pair {
349        NumericPair::Real(lhs_r, rhs_r) => {
350            let (solution, rcond) = solve_real(lhs_r, rhs_r, &options)?;
351            Ok(LinsolveEval::new(
352                tensor::tensor_into_value(solution),
353                Some(rcond),
354            ))
355        }
356        NumericPair::Complex(lhs_c, rhs_c) => {
357            let (solution, rcond) = solve_complex(lhs_c, rhs_c, &options)?;
358            Ok(LinsolveEval::new(
359                Value::ComplexTensor(solution),
360                Some(rcond),
361            ))
362        }
363    }
364}
365
366/// Host implementation shared with acceleration providers that fall back to CPU execution.
367pub fn linsolve_host_real_for_provider(
368    lhs: &Tensor,
369    rhs: &Tensor,
370    options: &ProviderLinsolveOptions,
371) -> BuiltinResult<(Tensor, f64)> {
372    let opts = SolveOptions::from(options);
373    let lhs = tensor::integer_tensor_to_f64(lhs.clone()).map_err(builtin_error)?;
374    let rhs = tensor::integer_tensor_to_f64(rhs.clone()).map_err(builtin_error)?;
375    solve_real(lhs, rhs, &opts)
376}
377
378/// Result wrapper that exposes both primary and secondary outputs.
379#[derive(Clone)]
380pub struct LinsolveEval {
381    solution: Value,
382    rcond: Option<f64>,
383}
384
385impl LinsolveEval {
386    fn new(solution: Value, rcond: Option<f64>) -> Self {
387        Self { solution, rcond }
388    }
389
390    /// Primary solution output.
391    pub fn solution(&self) -> Value {
392        self.solution.clone()
393    }
394
395    /// Estimated reciprocal condition number (second output).
396    pub fn reciprocal_condition(&self) -> Value {
397        match self.rcond {
398            Some(r) => Value::Num(r),
399            None => Value::Num(f64::NAN),
400        }
401    }
402}
403
404#[derive(Clone, Default)]
405pub struct SolveOptions {
406    lower: bool,
407    upper: bool,
408    rectangular: bool,
409    transposed: bool,
410    conjugate: bool,
411    symmetric: bool,
412    posdef: bool,
413    rcond: Option<f64>,
414}
415
416impl From<&SolveOptions> for ProviderLinsolveOptions {
417    fn from(opts: &SolveOptions) -> Self {
418        Self {
419            lower: opts.lower,
420            upper: opts.upper,
421            rectangular: opts.rectangular,
422            transposed: opts.transposed,
423            conjugate: opts.conjugate,
424            symmetric: opts.symmetric,
425            posdef: opts.posdef,
426            need_rcond: false,
427            rcond: opts.rcond,
428        }
429    }
430}
431
432impl From<&ProviderLinsolveOptions> for SolveOptions {
433    fn from(opts: &ProviderLinsolveOptions) -> Self {
434        Self {
435            lower: opts.lower,
436            upper: opts.upper,
437            rectangular: opts.rectangular,
438            transposed: opts.transposed,
439            conjugate: opts.conjugate,
440            symmetric: opts.symmetric,
441            posdef: opts.posdef,
442            rcond: opts.rcond,
443        }
444    }
445}
446
447fn options_from_rest(rest: &[Value]) -> BuiltinResult<SolveOptions> {
448    match rest.len() {
449        0 => Ok(SolveOptions::default()),
450        1 => parse_options(&rest[0]),
451        _ => Err(argument_error("linsolve: too many input arguments")),
452    }
453}
454
455/// Public helper for the VM multi-output surface.
456pub async fn evaluate_args(lhs: Value, rhs: Value, rest: &[Value]) -> BuiltinResult<LinsolveEval> {
457    let options = options_from_rest(rest)?;
458    ensure_linsolve_extensions(&lhs, &rhs).await?;
459    crate::builtins::common::validation::reject_typed_complex_integer(&lhs, NAME)?;
460    crate::builtins::common::validation::reject_typed_complex_integer(&rhs, NAME)?;
461    evaluate(lhs, rhs, options).await
462}
463
464async fn ensure_linsolve_extensions(lhs: &Value, rhs: &Value) -> BuiltinResult<()> {
465    let integer = |value: &Value| {
466        matches!(value, Value::Int(_))
467            || matches!(value, Value::Tensor(t) if t.integer_storage().is_some())
468            || matches!(value, Value::GpuTensor(h) if runmat_accelerate_api::handle_integer_type(h).is_some())
469    };
470    if integer(lhs) || integer(rhs) {
471        crate::compatibility::ensure_builtin_extension_enabled(
472            &LINSOLVE_INTEGER_INPUT_EXTENSION,
473            NAME,
474        )?;
475        for value in [lhs, rhs] {
476            if integer(value)
477                && !crate::builtins::common::validation::native_integer_value_is_exact_f64_async(
478                    value,
479                )
480                .await?
481            {
482                return Err(builtin_error(
483                    "linsolve: integer input lies outside the exact binary64 interval",
484                ));
485            }
486        }
487    }
488    if crate::builtins::common::validation::value_has_logical_class(lhs)
489        || crate::builtins::common::validation::value_has_logical_class(rhs)
490    {
491        crate::compatibility::ensure_builtin_extension_enabled(
492            &LINSOLVE_LOGICAL_INPUT_EXTENSION,
493            NAME,
494        )?;
495    }
496    if matches!(crate::output_count::current_output_count(), Some(2)) && [lhs, rhs].iter().any(|value| matches!(value, Value::GpuTensor(h) if runmat_accelerate_api::handle_is_explicit(h))) {
497        crate::compatibility::ensure_builtin_extension_enabled(&LINSOLVE_EXPLICIT_GPU_TWO_OUTPUT_EXTENSION, NAME)?;
498    }
499    Ok(())
500}
501
502async fn try_gpu_linsolve(
503    lhs: &Value,
504    rhs: &Value,
505    options: &SolveOptions,
506) -> BuiltinResult<Option<LinsolveEval>> {
507    if matches!(crate::output_count::current_output_count(), Some(n) if n > 2) {
508        return Ok(None);
509    }
510    let gpu_handles: Vec<&GpuTensorHandle> = [lhs, rhs]
511        .into_iter()
512        .filter_map(|value| {
513            if let Value::GpuTensor(handle) = value {
514                Some(handle)
515            } else {
516                None
517            }
518        })
519        .collect();
520    let provider = match gpu_handles
521        .first()
522        .map(|handle| gpu_helpers::exact_provider_for_handle(handle))
523        .unwrap_or_else(runmat_accelerate_api::provider)
524    {
525        Some(p) => p,
526        None => return Ok(None),
527    };
528    if gpu_handles.iter().any(|handle| {
529        gpu_helpers::exact_provider_for_handle(handle)
530            .is_none_or(|owner| !std::ptr::eq(owner, provider))
531    }) {
532        return Ok(None);
533    }
534
535    if contains_complex(lhs) || contains_complex(rhs) {
536        return Ok(None);
537    }
538    let host_extension_input = gpu_handles.is_empty()
539        && (value_has_integer_class(lhs)
540            || value_has_integer_class(rhs)
541            || crate::builtins::common::validation::value_has_logical_class(lhs)
542            || crate::builtins::common::validation::value_has_logical_class(rhs));
543    if host_extension_input
544        || (provider.precision() != runmat_accelerate_api::ProviderPrecision::F64
545            && [lhs, rhs].iter().any(|value| {
546                value_has_integer_class(value)
547                    || crate::builtins::common::validation::value_has_logical_class(value)
548            }))
549    {
550        return Ok(None);
551    }
552
553    let mut lhs_operand = match prepare_gpu_operand(lhs, provider)? {
554        Some(op) => op,
555        None => return Ok(None),
556    };
557    let mut rhs_operand = match prepare_gpu_operand(rhs, provider)? {
558        Some(op) => op,
559        None => {
560            release_operand(provider, &mut lhs_operand);
561            return Ok(None);
562        }
563    };
564
565    if is_scalar_handle(lhs_operand.handle()) || is_scalar_handle(rhs_operand.handle()) {
566        release_operand(provider, &mut lhs_operand);
567        release_operand(provider, &mut rhs_operand);
568        return Ok(None);
569    }
570
571    let mut provider_opts: ProviderLinsolveOptions = options.into();
572    let lhs_rows = lhs_operand.handle().shape.first().copied().unwrap_or(1);
573    let lhs_cols = lhs_operand.handle().shape.get(1).copied().unwrap_or(1);
574    let effective_rows = if options.transposed {
575        lhs_cols
576    } else {
577        lhs_rows
578    };
579    let effective_cols = if options.transposed {
580        lhs_rows
581    } else {
582        lhs_cols
583    };
584    let rectangular = effective_rows != effective_cols;
585    let wants_second_output = matches!(crate::output_count::current_output_count(), Some(2));
586    if rectangular && wants_second_output {
587        release_operand(provider, &mut lhs_operand);
588        release_operand(provider, &mut rhs_operand);
589        return Ok(None);
590    }
591    provider_opts.need_rcond = wants_second_output || options.rcond.is_some();
592    let result = provider
593        .linsolve(lhs_operand.handle(), rhs_operand.handle(), &provider_opts)
594        .await
595        .ok();
596
597    if let Some(ProviderLinsolveResult {
598        mut solution,
599        reciprocal_condition,
600    }) = result
601    {
602        let aliases_lhs = gpu_helpers::same_gpu_handle(&solution, lhs_operand.handle());
603        let aliases_rhs = gpu_helpers::same_gpu_handle(&solution, rhs_operand.handle());
604        let expected_rows = effective_cols;
605        let expected_cols = rhs_operand.handle().shape.get(1).copied().unwrap_or(1);
606        let valid = !aliases_lhs
607            && !aliases_rhs
608            && solution.shape == vec![expected_rows, expected_cols]
609            && solution.device_id == provider.device_id()
610            && gpu_helpers::exact_provider_for_handle(&solution)
611                .is_some_and(|owner| std::ptr::eq(owner, provider))
612            && runmat_accelerate_api::handle_storage(&solution)
613                == runmat_accelerate_api::GpuTensorStorage::Real
614            && runmat_accelerate_api::handle_precision(&solution) == Some(provider.precision())
615            && runmat_accelerate_api::handle_integer_type(&solution).is_none()
616            && !runmat_accelerate_api::handle_is_logical(&solution);
617        if !valid {
618            if !aliases_lhs && !aliases_rhs {
619                gpu_helpers::free_unprotected_exact_owner(
620                    &solution,
621                    &[lhs_operand.handle(), rhs_operand.handle()],
622                );
623            }
624            release_operand(provider, &mut lhs_operand);
625            release_operand(provider, &mut rhs_operand);
626            return Err(builtin_error(
627                "linsolve: provider returned malformed or aliased output",
628            ));
629        }
630        let provenance = gpu_handles
631            .iter()
632            .filter_map(|handle| runmat_accelerate_api::handle_provenance(handle))
633            .find(|provenance| *provenance == runmat_accelerate_api::GpuHandleProvenance::Explicit)
634            .unwrap_or(runmat_accelerate_api::GpuHandleProvenance::Automatic);
635        runmat_accelerate_api::set_handle_provenance(&mut solution, provenance);
636        runmat_accelerate_api::mark_residency(&solution);
637        release_operand(provider, &mut lhs_operand);
638        release_operand(provider, &mut rhs_operand);
639        let eval = LinsolveEval::new(Value::GpuTensor(solution), Some(reciprocal_condition));
640        return Ok(Some(eval));
641    }
642
643    release_operand(provider, &mut lhs_operand);
644    release_operand(provider, &mut rhs_operand);
645
646    Ok(None)
647}
648
649fn value_has_integer_class(value: &Value) -> bool {
650    matches!(value, Value::Int(_))
651        || matches!(value, Value::Tensor(t) if t.integer_storage().is_some())
652        || matches!(value, Value::GpuTensor(h) if runmat_accelerate_api::handle_integer_type(h).is_some())
653}
654
655fn parse_options(value: &Value) -> BuiltinResult<SolveOptions> {
656    let struct_val = match value {
657        Value::Struct(s) => s,
658        other => {
659            return Err(argument_error(format!(
660                "linsolve: opts must be a struct, got {other:?}"
661            )))
662        }
663    };
664    let mut opts = SolveOptions::default();
665    for (key, raw_value) in &struct_val.fields {
666        let name = key.to_ascii_uppercase();
667        match name.as_str() {
668            "LT" => opts.lower = parse_bool_field("LT", raw_value)?,
669            "UT" => opts.upper = parse_bool_field("UT", raw_value)?,
670            "RECT" => opts.rectangular = parse_bool_field("RECT", raw_value)?,
671            "SYM" => opts.symmetric = parse_bool_field("SYM", raw_value)?,
672            "POSDEF" => opts.posdef = parse_bool_field("POSDEF", raw_value)?,
673            "TRANSA" => {
674                if matches!(
675                    raw_value,
676                    Value::CharArray(_) | Value::String(_) | Value::StringArray(_)
677                ) {
678                    crate::compatibility::ensure_builtin_extension_enabled(
679                        &LINSOLVE_TEXT_TRANSA_EXTENSION,
680                        NAME,
681                    )?;
682                }
683                let transa = parse_transa(raw_value)?;
684                opts.transposed = transa != TransposeMode::None;
685                opts.conjugate = transa == TransposeMode::Conjugate;
686            }
687            "RCOND" => {
688                crate::compatibility::ensure_builtin_extension_enabled(
689                    &LINSOLVE_RCOND_EXTENSION,
690                    NAME,
691                )?;
692                let threshold = parse_scalar_f64("RCOND", raw_value)?;
693                if threshold < 0.0 {
694                    return Err(argument_error("linsolve: RCOND must be non-negative"));
695                }
696                opts.rcond = Some(threshold);
697            }
698            other => {
699                return Err(argument_error(format!(
700                    "linsolve: unknown option '{other}'"
701                )))
702            }
703        }
704    }
705    if opts.lower && opts.upper {
706        return Err(argument_error(
707            "linsolve: LT and UT are mutually exclusive.",
708        ));
709    }
710    Ok(opts)
711}
712
713fn parse_bool_field(name: &str, value: &Value) -> BuiltinResult<bool> {
714    if matches!(value, Value::Int(_))
715        || matches!(value, Value::Tensor(t) if t.integer_storage().is_some())
716    {
717        crate::compatibility::ensure_builtin_extension_enabled(
718            &LINSOLVE_INTEGER_OPTION_EXTENSION,
719            NAME,
720        )?;
721    }
722    match value {
723        Value::Bool(b) => Ok(*b),
724        Value::Int(i) => Ok(!i.is_zero()),
725        Value::Num(n) => Ok(*n != 0.0),
726        Value::Tensor(t) if tensor::is_scalar_tensor(t) => Ok(match scalar_tensor_integer(t) {
727            Some(value) => !value.is_zero(),
728            None => tensor::tensor_value_f64(t, 0) != 0.0,
729        }),
730        Value::LogicalArray(arr) if arr.len() == 1 => Ok(arr.data[0] != 0),
731        other => Err(argument_error(format!(
732            "linsolve: option '{name}' must be logical or numeric, got {other:?}"
733        ))),
734    }
735}
736
737fn parse_scalar_f64(name: &str, value: &Value) -> BuiltinResult<f64> {
738    match value {
739        Value::Num(n) => Ok(*n),
740        Value::Int(i) => Ok(i.to_f64()),
741        Value::Tensor(t) if tensor::is_scalar_tensor(t) => Ok(match scalar_tensor_integer(t) {
742            Some(value) => value.to_f64(),
743            None => tensor::tensor_value_f64(t, 0),
744        }),
745        other => Err(argument_error(format!(
746            "linsolve: option '{name}' must be a scalar numeric value, got {other:?}"
747        ))),
748    }
749}
750
751fn scalar_tensor_integer(tensor: &Tensor) -> Option<IntValue> {
752    tensor
753        .integer_storage()
754        .and_then(|storage| storage.value_at(0))
755}
756
757#[derive(Copy, Clone, PartialEq, Eq)]
758enum TransposeMode {
759    None,
760    Transpose,
761    Conjugate,
762}
763
764fn parse_transa(value: &Value) -> BuiltinResult<TransposeMode> {
765    match value {
766        Value::Bool(false) => return Ok(TransposeMode::None),
767        Value::Bool(true) => return Ok(TransposeMode::Conjugate),
768        Value::LogicalArray(array) if array.len() == 1 && array.data[0] == 0 => {
769            return Ok(TransposeMode::None)
770        }
771        Value::LogicalArray(array) if array.len() == 1 => return Ok(TransposeMode::Conjugate),
772        _ => {}
773    }
774    let text = tensor::value_to_string(value)
775        .ok_or_else(|| argument_error("linsolve: TRANSA must be a logical scalar"))?;
776    if text.is_empty() {
777        return Err(argument_error("linsolve: TRANSA cannot be empty"));
778    }
779    match text.trim().to_ascii_uppercase().as_str() {
780        "N" => Ok(TransposeMode::None),
781        "T" => Ok(TransposeMode::Transpose),
782        "C" => Ok(TransposeMode::Conjugate),
783        other => Err(argument_error(format!(
784            "linsolve: extended text TRANSA must be 'N', 'T', or 'C', got '{other}'"
785        ))),
786    }
787}
788
789enum NumericInput {
790    Real(Tensor),
791    Complex(ComplexTensor),
792}
793
794enum NumericPair {
795    Real(Tensor, Tensor),
796    Complex(ComplexTensor, ComplexTensor),
797}
798
799async fn coerce_numeric_pair(lhs: Value, rhs: Value) -> BuiltinResult<NumericPair> {
800    let lhs_num = coerce_numeric(lhs).await?;
801    let rhs_num = coerce_numeric(rhs).await?;
802    match (lhs_num, rhs_num) {
803        (NumericInput::Real(lhs_r), NumericInput::Real(rhs_r)) => {
804            Ok(NumericPair::Real(lhs_r, rhs_r))
805        }
806        (NumericInput::Complex(lhs_c), NumericInput::Complex(rhs_c)) => {
807            Ok(NumericPair::Complex(lhs_c, rhs_c))
808        }
809        (NumericInput::Complex(lhs_c), NumericInput::Real(rhs_r)) => {
810            let rhs_c = promote_real_tensor(&rhs_r)?;
811            Ok(NumericPair::Complex(lhs_c, rhs_c))
812        }
813        (NumericInput::Real(lhs_r), NumericInput::Complex(rhs_c)) => {
814            let lhs_c = promote_real_tensor(&lhs_r)?;
815            Ok(NumericPair::Complex(lhs_c, rhs_c))
816        }
817    }
818}
819
820async fn coerce_numeric(value: Value) -> BuiltinResult<NumericInput> {
821    match value {
822        Value::Tensor(tensor) => {
823            let tensor = tensor::integer_tensor_to_f64(tensor).map_err(builtin_error)?;
824            ensure_matrix_shape(NAME, &tensor.shape)?;
825            Ok(NumericInput::Real(tensor))
826        }
827        Value::LogicalArray(logical) => {
828            let tensor = tensor::logical_to_tensor(&logical).map_err(builtin_error)?;
829            ensure_matrix_shape(NAME, &tensor.shape)?;
830            Ok(NumericInput::Real(tensor))
831        }
832        Value::Num(n) => {
833            let tensor = Tensor::new(vec![n], vec![1, 1]).map_err(builtin_error)?;
834            Ok(NumericInput::Real(tensor))
835        }
836        Value::Int(i) => {
837            let tensor = Tensor::new(vec![i.to_f64()], vec![1, 1]).map_err(builtin_error)?;
838            Ok(NumericInput::Real(tensor))
839        }
840        Value::Bool(b) => {
841            let tensor =
842                Tensor::new(vec![if b { 1.0 } else { 0.0 }], vec![1, 1]).map_err(builtin_error)?;
843            Ok(NumericInput::Real(tensor))
844        }
845        Value::Complex(re, im) => {
846            let tensor = ComplexTensor::new(vec![(re, im)], vec![1, 1]).map_err(builtin_error)?;
847            Ok(NumericInput::Complex(tensor))
848        }
849        Value::ComplexTensor(ct) => {
850            ensure_matrix_shape(NAME, &ct.shape)?;
851            Ok(NumericInput::Complex(ct))
852        }
853        Value::GpuTensor(handle) => {
854            let tensor = gpu_helpers::gather_tensor_async(&handle)
855                .await
856                .map_err(map_control_flow)?;
857            let tensor = tensor::integer_tensor_to_f64(tensor).map_err(builtin_error)?;
858            ensure_matrix_shape(NAME, &tensor.shape)?;
859            Ok(NumericInput::Real(tensor))
860        }
861        other => Err(builtin_error(format!(
862            "{NAME}: unsupported input type {:?}; convert to numeric values first",
863            other
864        ))),
865    }
866}
867
868fn contains_complex(value: &Value) -> bool {
869    matches!(value, Value::Complex(_, _) | Value::ComplexTensor(_))
870}
871
872fn is_scalar_handle(handle: &GpuTensorHandle) -> bool {
873    crate::builtins::common::shape::is_scalar_shape(&handle.shape)
874}
875
876struct PreparedOperand {
877    handle: GpuTensorHandle,
878    owned: bool,
879}
880
881impl PreparedOperand {
882    fn borrowed(handle: &GpuTensorHandle) -> Self {
883        Self {
884            handle: handle.clone(),
885            owned: false,
886        }
887    }
888
889    fn owned(handle: GpuTensorHandle) -> Self {
890        Self {
891            handle,
892            owned: true,
893        }
894    }
895
896    fn handle(&self) -> &GpuTensorHandle {
897        &self.handle
898    }
899}
900
901fn prepare_gpu_operand(
902    value: &Value,
903    provider: &'static dyn AccelProvider,
904) -> BuiltinResult<Option<PreparedOperand>> {
905    match value {
906        Value::GpuTensor(handle) => {
907            if handle.device_id != provider.device_id()
908                || gpu_helpers::exact_provider_for_handle(handle)
909                    .is_none_or(|owner| !std::ptr::eq(owner, provider))
910                || is_scalar_handle(handle)
911            {
912                Ok(None)
913            } else {
914                Ok(Some(PreparedOperand::borrowed(handle)))
915            }
916        }
917        Value::Tensor(tensor) => {
918            if tensor::is_scalar_tensor(tensor) {
919                Ok(None)
920            } else {
921                let uploaded = upload_tensor(provider, tensor)?;
922                Ok(Some(PreparedOperand::owned(uploaded)))
923            }
924        }
925        Value::LogicalArray(logical) => {
926            if logical.data.len() == 1 {
927                Ok(None)
928            } else {
929                let tensor = tensor::logical_to_tensor(logical).map_err(builtin_error)?;
930                let uploaded = upload_tensor(provider, &tensor)?;
931                Ok(Some(PreparedOperand::owned(uploaded)))
932            }
933        }
934        _ => Ok(None),
935    }
936}
937
938fn upload_tensor(
939    provider: &'static dyn AccelProvider,
940    tensor: &Tensor,
941) -> BuiltinResult<GpuTensorHandle> {
942    // The current provider view is floating; materialize that transfer
943    // boundary from authoritative host storage without consulting a mirror.
944    let values = tensor::tensor_values_f64_cow(tensor);
945    let view = HostTensorView {
946        data: values.as_ref(),
947        shape: &tensor.shape,
948    };
949    provider
950        .upload(&view)
951        .map_err(|e| builtin_error(format!("{NAME}: {e}")))
952}
953
954fn release_operand(provider: &'static dyn AccelProvider, operand: &mut PreparedOperand) {
955    if operand.owned {
956        let _ = provider.free(&operand.handle);
957        operand.owned = false;
958    }
959}
960
961fn solve_real(lhs: Tensor, rhs: Tensor, options: &SolveOptions) -> BuiltinResult<(Tensor, f64)> {
962    let mut lhs_effective = lhs;
963    let mut rhs_effective = rhs;
964    let mut lower = options.lower;
965    let mut upper = options.upper;
966
967    if options.transposed {
968        lhs_effective = transpose_tensor(&lhs_effective);
969        if options.conjugate {
970            conjugate_in_place(&mut lhs_effective);
971        }
972        if lower || upper {
973            std::mem::swap(&mut lower, &mut upper);
974        }
975    }
976
977    rhs_effective = normalize_rhs_tensor(rhs_effective, lhs_effective.rows())?;
978
979    if lower {
980        ensure_square(lhs_effective.rows(), lhs_effective.cols())?;
981        let (solution, rcond) = forward_substitution_real(&lhs_effective, &rhs_effective)?;
982        enforce_rcond(options, rcond)?;
983        return Ok((solution, rcond));
984    }
985
986    if upper {
987        ensure_square(lhs_effective.rows(), lhs_effective.cols())?;
988        let (solution, rcond) = backward_substitution_real(&lhs_effective, &rhs_effective)?;
989        enforce_rcond(options, rcond)?;
990        return Ok((solution, rcond));
991    }
992
993    let (solution, rcond, rank) = solve_general_real(&lhs_effective, &rhs_effective)?;
994    enforce_rcond(options, rcond)?;
995    Ok((
996        solution,
997        if lhs_effective.rows() == lhs_effective.cols() {
998            rcond
999        } else {
1000            rank
1001        },
1002    ))
1003}
1004
1005fn solve_complex(
1006    lhs: ComplexTensor,
1007    rhs: ComplexTensor,
1008    options: &SolveOptions,
1009) -> BuiltinResult<(ComplexTensor, f64)> {
1010    let mut lhs_effective = lhs;
1011    let mut rhs_effective = rhs;
1012    let mut lower = options.lower;
1013    let mut upper = options.upper;
1014
1015    if options.transposed {
1016        lhs_effective = transpose_complex(&lhs_effective);
1017        if options.conjugate {
1018            conjugate_complex_in_place(&mut lhs_effective);
1019        }
1020        if lower || upper {
1021            std::mem::swap(&mut lower, &mut upper);
1022        }
1023    }
1024
1025    rhs_effective = normalize_rhs_complex(rhs_effective, lhs_effective.rows)?;
1026
1027    if lower {
1028        ensure_square(lhs_effective.rows, lhs_effective.cols)?;
1029        let (solution, rcond) = forward_substitution_complex(&lhs_effective, &rhs_effective)?;
1030        enforce_rcond(options, rcond)?;
1031        return Ok((solution, rcond));
1032    }
1033
1034    if upper {
1035        ensure_square(lhs_effective.rows, lhs_effective.cols)?;
1036        let (solution, rcond) = backward_substitution_complex(&lhs_effective, &rhs_effective)?;
1037        enforce_rcond(options, rcond)?;
1038        return Ok((solution, rcond));
1039    }
1040
1041    let (solution, rcond, rank) = solve_general_complex(&lhs_effective, &rhs_effective)?;
1042    enforce_rcond(options, rcond)?;
1043    Ok((
1044        solution,
1045        if lhs_effective.rows == lhs_effective.cols {
1046            rcond
1047        } else {
1048            rank
1049        },
1050    ))
1051}
1052
1053fn forward_substitution_real(lhs: &Tensor, rhs: &Tensor) -> BuiltinResult<(Tensor, f64)> {
1054    let n = lhs.rows();
1055    let lhs_values = tensor::tensor_values_f64_cow(lhs);
1056    let mut solution = tensor::tensor_values_f64(rhs);
1057    let nrhs = solution.len() / n;
1058    let mut min_diag = f64::INFINITY;
1059    let mut max_diag = 0.0_f64;
1060
1061    for col in 0..nrhs {
1062        for i in 0..n {
1063            let diag = lhs_values[i + i * n];
1064            let diag_abs = diag.abs();
1065            min_diag = min_diag.min(diag_abs);
1066            max_diag = max_diag.max(diag_abs);
1067            if diag_abs == 0.0 {
1068                return Err(builtin_error(
1069                    "linsolve: matrix is singular to working precision.",
1070                ));
1071            }
1072            let mut accum = 0.0;
1073            for j in 0..i {
1074                accum += lhs_values[i + j * n] * solution[j + col * n];
1075            }
1076            let rhs_value = solution[i + col * n] - accum;
1077            solution[i + col * n] = rhs_value / diag;
1078        }
1079    }
1080
1081    let rcond = diagonal_rcond(min_diag, max_diag);
1082    let tensor = real_solution_tensor(solution, rhs.shape.clone(), real_solution_dtype(lhs, rhs))?;
1083    Ok((tensor, rcond))
1084}
1085
1086fn backward_substitution_real(lhs: &Tensor, rhs: &Tensor) -> BuiltinResult<(Tensor, f64)> {
1087    let n = lhs.rows();
1088    let lhs_values = tensor::tensor_values_f64_cow(lhs);
1089    let mut solution = tensor::tensor_values_f64(rhs);
1090    let nrhs = solution.len() / n;
1091    let mut min_diag = f64::INFINITY;
1092    let mut max_diag = 0.0_f64;
1093
1094    for col in 0..nrhs {
1095        for row_rev in 0..n {
1096            let i = n - 1 - row_rev;
1097            let diag = lhs_values[i + i * n];
1098            let diag_abs = diag.abs();
1099            min_diag = min_diag.min(diag_abs);
1100            max_diag = max_diag.max(diag_abs);
1101            if diag_abs == 0.0 {
1102                return Err(builtin_error(
1103                    "linsolve: matrix is singular to working precision.",
1104                ));
1105            }
1106            let mut accum = 0.0;
1107            for j in (i + 1)..n {
1108                accum += lhs_values[i + j * n] * solution[j + col * n];
1109            }
1110            let rhs_value = solution[i + col * n] - accum;
1111            solution[i + col * n] = rhs_value / diag;
1112        }
1113    }
1114
1115    let rcond = diagonal_rcond(min_diag, max_diag);
1116    let tensor = real_solution_tensor(solution, rhs.shape.clone(), real_solution_dtype(lhs, rhs))?;
1117    Ok((tensor, rcond))
1118}
1119
1120fn forward_substitution_complex(
1121    lhs: &ComplexTensor,
1122    rhs: &ComplexTensor,
1123) -> BuiltinResult<(ComplexTensor, f64)> {
1124    let n = lhs.rows;
1125    let nrhs = rhs.materialize_f64().len() / n;
1126    let lhs_data: Vec<Complex64> = lhs
1127        .materialize_f64()
1128        .iter()
1129        .map(|&(re, im)| Complex64::new(re, im))
1130        .collect();
1131    let mut solution: Vec<Complex64> = rhs
1132        .materialize_f64()
1133        .iter()
1134        .map(|&(re, im)| Complex64::new(re, im))
1135        .collect();
1136    let mut min_diag = f64::INFINITY;
1137    let mut max_diag = 0.0_f64;
1138
1139    for col in 0..nrhs {
1140        for i in 0..n {
1141            let diag = lhs_data[i + i * n];
1142            let diag_abs = diag.norm();
1143            min_diag = min_diag.min(diag_abs);
1144            max_diag = max_diag.max(diag_abs);
1145            if diag_abs == 0.0 {
1146                return Err(builtin_error(
1147                    "linsolve: matrix is singular to working precision.",
1148                ));
1149            }
1150            let mut accum = Complex64::new(0.0, 0.0);
1151            for j in 0..i {
1152                accum += lhs_data[i + j * n] * solution[j + col * n];
1153            }
1154            let rhs_value = solution[i + col * n] - accum;
1155            solution[i + col * n] = rhs_value / diag;
1156        }
1157    }
1158
1159    let rcond = diagonal_rcond(min_diag, max_diag);
1160    let tensor = ComplexTensor::new(
1161        solution.iter().map(|c| (c.re, c.im)).collect(),
1162        rhs.shape.clone(),
1163    )
1164    .map_err(|e| builtin_error(format!("{NAME}: {e}")))?;
1165    Ok((tensor, rcond))
1166}
1167
1168fn backward_substitution_complex(
1169    lhs: &ComplexTensor,
1170    rhs: &ComplexTensor,
1171) -> BuiltinResult<(ComplexTensor, f64)> {
1172    let n = lhs.rows;
1173    let nrhs = rhs.materialize_f64().len() / n;
1174    let lhs_data: Vec<Complex64> = lhs
1175        .materialize_f64()
1176        .iter()
1177        .map(|&(re, im)| Complex64::new(re, im))
1178        .collect();
1179    let mut solution: Vec<Complex64> = rhs
1180        .materialize_f64()
1181        .iter()
1182        .map(|&(re, im)| Complex64::new(re, im))
1183        .collect();
1184    let mut min_diag = f64::INFINITY;
1185    let mut max_diag = 0.0_f64;
1186
1187    for col in 0..nrhs {
1188        for row_rev in 0..n {
1189            let i = n - 1 - row_rev;
1190            let diag = lhs_data[i + i * n];
1191            let diag_abs = diag.norm();
1192            min_diag = min_diag.min(diag_abs);
1193            max_diag = max_diag.max(diag_abs);
1194            if diag_abs == 0.0 {
1195                return Err(builtin_error(
1196                    "linsolve: matrix is singular to working precision.",
1197                ));
1198            }
1199            let mut accum = Complex64::new(0.0, 0.0);
1200            for j in (i + 1)..n {
1201                accum += lhs_data[i + j * n] * solution[j + col * n];
1202            }
1203            let rhs_value = solution[i + col * n] - accum;
1204            solution[i + col * n] = rhs_value / diag;
1205        }
1206    }
1207
1208    let rcond = diagonal_rcond(min_diag, max_diag);
1209    let tensor = ComplexTensor::new(
1210        solution.iter().map(|c| (c.re, c.im)).collect(),
1211        rhs.shape.clone(),
1212    )
1213    .map_err(|e| builtin_error(format!("{NAME}: {e}")))?;
1214    Ok((tensor, rcond))
1215}
1216
1217fn solve_general_real(lhs: &Tensor, rhs: &Tensor) -> BuiltinResult<(Tensor, f64, f64)> {
1218    let lhs_values = tensor::tensor_values_f64_cow(lhs);
1219    let rhs_values = tensor::tensor_values_f64_cow(rhs);
1220    let a = DMatrix::from_column_slice(lhs.rows(), lhs.cols(), lhs_values.as_ref());
1221    let b = DMatrix::from_column_slice(rhs.rows(), rhs.cols(), rhs_values.as_ref());
1222    let svd = SVD::new(a.clone(), true, true);
1223    let rcond = singular_value_rcond(svd.singular_values.as_slice());
1224    let tol = compute_svd_tolerance(svd.singular_values.as_slice(), lhs.rows(), lhs.cols());
1225    let rank = svd
1226        .singular_values
1227        .iter()
1228        .filter(|value| **value > tol)
1229        .count() as f64;
1230    let solution = svd
1231        .solve(&b, tol)
1232        .map_err(|e| builtin_error(format!("{NAME}: {e}")))?;
1233    let tensor = matrix_real_to_tensor(solution, real_solution_dtype(lhs, rhs))?;
1234    Ok((tensor, rcond, rank))
1235}
1236
1237fn solve_general_complex(
1238    lhs: &ComplexTensor,
1239    rhs: &ComplexTensor,
1240) -> BuiltinResult<(ComplexTensor, f64, f64)> {
1241    let a_data: Vec<Complex64> = lhs
1242        .materialize_f64()
1243        .iter()
1244        .map(|&(re, im)| Complex64::new(re, im))
1245        .collect();
1246    let b_data: Vec<Complex64> = rhs
1247        .materialize_f64()
1248        .iter()
1249        .map(|&(re, im)| Complex64::new(re, im))
1250        .collect();
1251    let a = DMatrix::from_column_slice(lhs.rows, lhs.cols, &a_data);
1252    let b = DMatrix::from_column_slice(rhs.rows, rhs.cols, &b_data);
1253    let svd = SVD::new(a.clone(), true, true);
1254    let rcond = singular_value_rcond(svd.singular_values.as_slice());
1255    let tol = compute_svd_tolerance(svd.singular_values.as_slice(), lhs.rows, lhs.cols);
1256    let rank = svd
1257        .singular_values
1258        .iter()
1259        .filter(|value| **value > tol)
1260        .count() as f64;
1261    let solution = svd
1262        .solve(&b, tol)
1263        .map_err(|e| builtin_error(format!("{NAME}: {e}")))?;
1264    let tensor = matrix_complex_to_tensor(solution)?;
1265    Ok((tensor, rcond, rank))
1266}
1267
1268fn normalize_rhs_tensor(rhs: Tensor, expected_rows: usize) -> BuiltinResult<Tensor> {
1269    if rhs.rows() == expected_rows {
1270        return Ok(rhs);
1271    }
1272    if rhs.shape.len() == 1 && rhs.shape[0] == expected_rows {
1273        return rhs
1274            .reshape(vec![expected_rows, 1])
1275            .map_err(|e| builtin_error(format!("{NAME}: {e}")));
1276    }
1277    if tensor::tensor_element_len(&rhs) == 0 && expected_rows == 0 {
1278        return Ok(rhs);
1279    }
1280    Err(builtin_error("Matrix dimensions must agree."))
1281}
1282
1283fn normalize_rhs_complex(rhs: ComplexTensor, expected_rows: usize) -> BuiltinResult<ComplexTensor> {
1284    if rhs.rows == expected_rows {
1285        return Ok(rhs);
1286    }
1287    if rhs.shape.len() == 1 && rhs.shape[0] == expected_rows {
1288        return ComplexTensor::new(rhs.materialize_f64(), vec![expected_rows, 1])
1289            .map_err(|e| builtin_error(format!("{NAME}: {e}")));
1290    }
1291    if rhs.materialize_f64().is_empty() && expected_rows == 0 {
1292        return Ok(rhs);
1293    }
1294    Err(builtin_error("Matrix dimensions must agree."))
1295}
1296
1297fn enforce_rcond(options: &SolveOptions, rcond: f64) -> BuiltinResult<()> {
1298    if let Some(threshold) = options.rcond {
1299        if rcond < threshold {
1300            return Err(builtin_error(
1301                "linsolve: matrix is singular to working precision.",
1302            ));
1303        }
1304    }
1305    Ok(())
1306}
1307
1308fn compute_svd_tolerance(singular_values: &[f64], rows: usize, cols: usize) -> f64 {
1309    let max_sv = singular_values
1310        .iter()
1311        .copied()
1312        .fold(0.0_f64, |acc, value| acc.max(value.abs()));
1313    let max_dim = rows.max(cols) as f64;
1314    f64::EPSILON * max_dim * max_sv.max(1.0)
1315}
1316
1317fn matrix_real_to_tensor(matrix: DMatrix<f64>, dtype: NumericDType) -> BuiltinResult<Tensor> {
1318    let rows = matrix.nrows();
1319    let cols = matrix.ncols();
1320    real_solution_tensor(matrix.as_slice().to_vec(), vec![rows, cols], dtype)
1321}
1322
1323fn real_solution_dtype(lhs: &Tensor, rhs: &Tensor) -> NumericDType {
1324    if lhs.numeric_dtype() == NumericDType::F32 && rhs.numeric_dtype() == NumericDType::F32 {
1325        NumericDType::F32
1326    } else {
1327        NumericDType::F64
1328    }
1329}
1330
1331fn real_solution_tensor(
1332    values: Vec<f64>,
1333    shape: Vec<usize>,
1334    dtype: NumericDType,
1335) -> BuiltinResult<Tensor> {
1336    let tensor = match dtype {
1337        NumericDType::F32 => Tensor::from_f32(
1338            values.into_iter().map(|value| value as f32).collect(),
1339            shape,
1340        ),
1341        NumericDType::F64 => Tensor::new(values, shape),
1342        _ => Err(format!(
1343            "linsolve: unsupported real solution class {}",
1344            dtype.class_name()
1345        )),
1346    };
1347    tensor.map_err(|e| builtin_error(format!("{NAME}: {e}")))
1348}
1349
1350fn matrix_complex_to_tensor(matrix: DMatrix<Complex64>) -> BuiltinResult<ComplexTensor> {
1351    let rows = matrix.nrows();
1352    let cols = matrix.ncols();
1353    let data: Vec<(f64, f64)> = matrix.as_slice().iter().map(|c| (c.re, c.im)).collect();
1354    ComplexTensor::new(data, vec![rows, cols]).map_err(|e| builtin_error(format!("{NAME}: {e}")))
1355}
1356
1357fn promote_real_tensor(tensor: &Tensor) -> BuiltinResult<ComplexTensor> {
1358    let values = tensor::tensor_values_f64_cow(tensor);
1359    let data: Vec<(f64, f64)> = values.iter().map(|&re| (re, 0.0)).collect();
1360    ComplexTensor::new(data, tensor.shape.clone())
1361        .map_err(|e| builtin_error(format!("{NAME}: {e}")))
1362}
1363
1364fn ensure_matrix_shape(name: &str, shape: &[usize]) -> BuiltinResult<()> {
1365    if is_effectively_matrix(shape) {
1366        Ok(())
1367    } else {
1368        Err(builtin_error(format!(
1369            "{name}: inputs must be 2-D matrices or vectors"
1370        )))
1371    }
1372}
1373
1374fn is_effectively_matrix(shape: &[usize]) -> bool {
1375    match shape.len() {
1376        0..=2 => true,
1377        _ => shape.iter().skip(2).all(|&dim| dim == 1),
1378    }
1379}
1380
1381fn ensure_square(rows: usize, cols: usize) -> BuiltinResult<()> {
1382    if rows == cols {
1383        Ok(())
1384    } else {
1385        Err(builtin_error(
1386            "linsolve: triangular solves require a square coefficient matrix.",
1387        ))
1388    }
1389}
1390
1391fn transpose_tensor(tensor: &Tensor) -> Tensor {
1392    let rows = tensor.rows();
1393    let cols = tensor.cols();
1394    let mut indices = vec![0usize; tensor::tensor_element_len(tensor)];
1395    for r in 0..rows {
1396        for c in 0..cols {
1397            indices[c + r * cols] = r + c * rows;
1398        }
1399    }
1400    let storage = tensor
1401        .clone()
1402        .into_numeric_storage()
1403        .expect("validated tensor storage");
1404    Tensor::from_numeric_storage(
1405        storage
1406            .gather(&indices)
1407            .expect("transpose indices in bounds"),
1408        vec![cols, rows],
1409    )
1410    .expect("transpose tensor shape matches storage")
1411}
1412
1413fn transpose_complex(tensor: &ComplexTensor) -> ComplexTensor {
1414    let rows = tensor.rows;
1415    let cols = tensor.cols;
1416    let mut data = vec![(0.0, 0.0); tensor.materialize_f64().len()];
1417    for r in 0..rows {
1418        for c in 0..cols {
1419            data[c + r * cols] = tensor.materialize_f64()[r + c * rows];
1420        }
1421    }
1422    ComplexTensor::new(data, vec![cols, rows]).expect("transpose_complex valid")
1423}
1424
1425fn conjugate_in_place(_tensor: &mut Tensor) {
1426    // Real-valued matrices are unaffected by conjugation.
1427}
1428
1429fn conjugate_complex_in_place(tensor: &mut ComplexTensor) {
1430    let shape = tensor.shape.clone();
1431    let storage = match tensor.clone().into_complex_storage() {
1432        ComplexStorage::F64(mut values) => {
1433            for value in &mut values {
1434                value.1 = -value.1;
1435            }
1436            ComplexStorage::F64(values)
1437        }
1438        ComplexStorage::F32(mut values) => {
1439            for value in &mut values {
1440                value.1 = -value.1;
1441            }
1442            ComplexStorage::F32(values)
1443        }
1444        ComplexStorage::Integer(storage) => ComplexStorage::Integer(
1445            IntegerComplexStorage::new(
1446                storage.real,
1447                conjugate_integer_imaginary_storage(storage.imag),
1448            )
1449            .expect("complex integer component classes remain matched"),
1450        ),
1451    };
1452    *tensor = ComplexTensor::from_complex_storage(storage, shape)
1453        .expect("conjugated complex storage retains shape");
1454}
1455
1456#[cfg(test)]
1457pub(crate) mod tests {
1458    use super::*;
1459    use futures::executor::block_on;
1460    use runmat_accelerate_api::HostTensorView;
1461    use runmat_builtins::{ResolveContext, Type};
1462    use runmat_value::{CharArray, IntegerStorage, NumericStorage, StructValue};
1463    fn unwrap_error(err: crate::RuntimeError) -> crate::RuntimeError {
1464        err
1465    }
1466
1467    fn approx_eq(actual: f64, expected: f64) {
1468        assert!((actual - expected).abs() < 1e-7);
1469    }
1470
1471    fn evaluate_args(a: Value, b: Value, rest: &[Value]) -> Result<LinsolveEval, RuntimeError> {
1472        let _extensions = crate::compatibility::push_runmat_extensions_enabled(true);
1473        block_on(super::evaluate_args(a, b, rest))
1474    }
1475
1476    #[test]
1477    fn linsolve_type_uses_rhs_columns() {
1478        let out = left_divide_type(
1479            &[
1480                Type::Tensor {
1481                    shape: Some(vec![Some(2), Some(2)]),
1482                },
1483                Type::Tensor {
1484                    shape: Some(vec![Some(2), Some(3)]),
1485                },
1486            ],
1487            &ResolveContext::new(Vec::new()),
1488        );
1489        assert_eq!(
1490            out,
1491            Type::Tensor {
1492                shape: Some(vec![Some(2), Some(3)])
1493            }
1494        );
1495    }
1496
1497    #[test]
1498    fn linsolve_descriptor_signatures_cover_core_forms() {
1499        let labels: Vec<&str> = LINSOLVE_DESCRIPTOR
1500            .signatures
1501            .iter()
1502            .map(|signature| signature.label)
1503            .collect();
1504        assert!(labels.contains(&"X = linsolve(A, B)"));
1505        assert!(labels.contains(&"X = linsolve(A, B, opts)"));
1506        assert!(labels.contains(&"[X, R] = linsolve(A, B)"));
1507        assert!(labels.contains(&"[X, R] = linsolve(A, B, opts)"));
1508    }
1509
1510    #[test]
1511    fn linsolve_descriptor_errors_have_stable_codes() {
1512        let codes: Vec<&str> = LINSOLVE_DESCRIPTOR
1513            .errors
1514            .iter()
1515            .map(|err| err.code)
1516            .collect();
1517        assert!(codes.contains(&"RM.LINSOLVE.INVALID_ARGUMENT"));
1518        assert!(codes.contains(&"RM.LINSOLVE.INVALID_INPUT"));
1519        assert!(codes.contains(&"RM.LINSOLVE.INTERNAL"));
1520    }
1521
1522    use crate::builtins::common::test_support;
1523    use runmat_accelerate_api::ProviderTelemetry;
1524
1525    fn linsolve_builtin(lhs: Value, rhs: Value, rest: Vec<Value>) -> BuiltinResult<Value> {
1526        let _extensions = crate::compatibility::push_runmat_extensions_enabled(true);
1527        block_on(super::linsolve_builtin(lhs, rhs, rest))
1528    }
1529
1530    #[test]
1531    fn linsolve_integer_extension_is_rejected_in_matlab_mode() {
1532        let _matlab = crate::compatibility::push_runmat_extensions_enabled(false);
1533        let lhs = Tensor::new_integer(IntegerStorage::I32(vec![1]), vec![1, 1]).unwrap();
1534        let error = block_on(super::linsolve_builtin(
1535            Value::Tensor(lhs),
1536            Value::Num(1.0),
1537            Vec::new(),
1538        ))
1539        .expect_err("integer linsolve is a RunMat extension");
1540        assert_eq!(
1541            error.identifier(),
1542            LINSOLVE_INTEGER_INPUT_EXTENSION.error_identifier
1543        );
1544    }
1545
1546    #[test]
1547    fn linsolve_option_extensions_are_independently_gated() {
1548        let _matlab = crate::compatibility::push_runmat_extensions_enabled(false);
1549
1550        let mut integer_option = StructValue::new();
1551        integer_option
1552            .fields
1553            .insert("LT".to_string(), Value::Int(IntValue::I8(1)));
1554        let error = match options_from_rest(&[Value::Struct(integer_option)]) {
1555            Err(error) => error,
1556            Ok(_) => panic!("typed structural option must be gated"),
1557        };
1558        assert_eq!(
1559            error.identifier(),
1560            LINSOLVE_INTEGER_OPTION_EXTENSION.error_identifier
1561        );
1562
1563        let mut text_transa = StructValue::new();
1564        text_transa
1565            .fields
1566            .insert("TRANSA".to_string(), Value::from("T"));
1567        let error = match options_from_rest(&[Value::Struct(text_transa)]) {
1568            Err(error) => error,
1569            Ok(_) => panic!("text TRANSA must be gated"),
1570        };
1571        assert_eq!(
1572            error.identifier(),
1573            LINSOLVE_TEXT_TRANSA_EXTENSION.error_identifier
1574        );
1575
1576        let mut rcond = StructValue::new();
1577        rcond.fields.insert("RCOND".to_string(), Value::Num(0.1));
1578        let error = match options_from_rest(&[Value::Struct(rcond)]) {
1579            Err(error) => error,
1580            Ok(_) => panic!("RCOND must be gated"),
1581        };
1582        assert_eq!(
1583            error.identifier(),
1584            LINSOLVE_RCOND_EXTENSION.error_identifier
1585        );
1586    }
1587
1588    #[test]
1589    fn linsolve_documented_logical_transa_does_not_require_extension() {
1590        let _matlab = crate::compatibility::push_runmat_extensions_enabled(false);
1591        let mut options = StructValue::new();
1592        options
1593            .fields
1594            .insert("TRANSA".to_string(), Value::Bool(true));
1595        let parsed = options_from_rest(&[Value::Struct(options)]).expect("logical TRANSA");
1596        assert!(parsed.transposed);
1597        assert!(parsed.conjugate);
1598    }
1599
1600    fn evaluate(lhs: Value, rhs: Value, options: SolveOptions) -> BuiltinResult<LinsolveEval> {
1601        block_on(super::evaluate(lhs, rhs, options))
1602    }
1603
1604    fn fallback_count(telemetry: &ProviderTelemetry, reason: &str) -> u64 {
1605        telemetry
1606            .solve_fallbacks
1607            .iter()
1608            .find(|entry| entry.reason == reason)
1609            .map(|entry| entry.count)
1610            .unwrap_or(0)
1611    }
1612
1613    #[cfg(feature = "wgpu")]
1614    fn kernel_launch_count(telemetry: &ProviderTelemetry, kernel: &str) -> usize {
1615        telemetry
1616            .kernel_launches
1617            .iter()
1618            .filter(|entry| entry.kernel == kernel)
1619            .count()
1620    }
1621
1622    fn clear_accel_provider_state() {
1623        runmat_accelerate_api::set_thread_provider(None);
1624        runmat_accelerate_api::clear_provider();
1625    }
1626
1627    fn host_linsolve_real(
1628        a: &Tensor,
1629        b: &Tensor,
1630        options: ProviderLinsolveOptions,
1631    ) -> (Tensor, f64) {
1632        super::linsolve_host_real_for_provider(a, b, &options).expect("host linsolve")
1633    }
1634
1635    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1636    #[test]
1637    fn linsolve_basic_square() {
1638        let _accel_guard = test_support::accel_test_lock();
1639        clear_accel_provider_state();
1640        let a = Tensor::new(vec![2.0, 1.0, 1.0, 2.0], vec![2, 2]).unwrap();
1641        let b = Tensor::new(vec![4.0, 5.0], vec![2, 1]).unwrap();
1642        let result =
1643            linsolve_builtin(Value::Tensor(a), Value::Tensor(b), Vec::new()).expect("linsolve");
1644        let t = test_support::gather(result).expect("gather");
1645        assert_eq!(t.shape, vec![2, 1]);
1646        approx_eq(t.materialize_f64()[0], 1.0);
1647        approx_eq(t.materialize_f64()[1], 2.0);
1648    }
1649
1650    #[test]
1651    fn linsolve_cpu_preserves_native_single_for_general_and_triangular_solutions() {
1652        let _accel_guard = test_support::accel_test_lock();
1653        clear_accel_provider_state();
1654
1655        let a = Tensor::from_f32(vec![2.0, 1.0, 1.0, 2.0], vec![2, 2]).unwrap();
1656        let b = Tensor::from_f32(vec![4.0, 5.0], vec![2, 1]).unwrap();
1657        let result =
1658            linsolve_builtin(Value::Tensor(a), Value::Tensor(b), Vec::new()).expect("linsolve");
1659        let tensor = test_support::gather(result).expect("gather");
1660        assert_eq!(
1661            tensor.into_numeric_storage().unwrap(),
1662            NumericStorage::F32(vec![1.0, 2.0])
1663        );
1664
1665        let lower = Tensor::from_f32(vec![2.0, 1.0, 0.0, 3.0], vec![2, 2]).unwrap();
1666        let rhs = Tensor::from_f32(vec![2.0, 7.0], vec![2, 1]).unwrap();
1667        let mut options = StructValue::new();
1668        options.fields.insert("LT".to_string(), Value::Bool(true));
1669        let result = linsolve_builtin(
1670            Value::Tensor(lower),
1671            Value::Tensor(rhs),
1672            vec![Value::Struct(options)],
1673        )
1674        .expect("lower triangular linsolve");
1675        let tensor = test_support::gather(result).expect("gather");
1676        assert_eq!(
1677            tensor.into_numeric_storage().unwrap(),
1678            NumericStorage::F32(vec![1.0, 2.0])
1679        );
1680    }
1681
1682    #[test]
1683    fn linsolve_reads_typed_integer_tensor_storage_exactly() {
1684        let _accel_guard = test_support::accel_test_lock();
1685        clear_accel_provider_state();
1686        let a = Tensor::new_integer(IntegerStorage::U64(vec![2, 1, 1, 2]), vec![2, 2])
1687            .expect("integer lhs");
1688        let b =
1689            Tensor::new_integer(IntegerStorage::U64(vec![4, 5]), vec![2, 1]).expect("integer rhs");
1690
1691        let result =
1692            linsolve_builtin(Value::Tensor(a), Value::Tensor(b), Vec::new()).expect("linsolve");
1693        let tensor = test_support::gather(result).expect("gather");
1694        assert_eq!(tensor.shape, vec![2, 1]);
1695        approx_eq(tensor.materialize_f64()[0], 1.0);
1696        approx_eq(tensor.materialize_f64()[1], 2.0);
1697    }
1698
1699    #[test]
1700    fn linsolve_complex_promotion_reads_typed_integer_storage_exactly() {
1701        let _accel_guard = test_support::accel_test_lock();
1702        clear_accel_provider_state();
1703
1704        let real_lhs = Tensor::new_integer(IntegerStorage::I64(vec![1, 0, 0, 1]), vec![2, 2])
1705            .expect("integer lhs");
1706        let complex_rhs = ComplexTensor::new(vec![(3.0, 4.0), (5.0, -6.0)], vec![2, 1]).unwrap();
1707        let result = linsolve_builtin(
1708            Value::Tensor(real_lhs),
1709            Value::ComplexTensor(complex_rhs),
1710            Vec::new(),
1711        )
1712        .expect("linsolve");
1713        let Value::ComplexTensor(out) = result else {
1714            panic!("expected complex tensor output");
1715        };
1716        assert_eq!(out.shape, vec![2, 1]);
1717        approx_eq(out.materialize_f64()[0].0, 3.0);
1718        approx_eq(out.materialize_f64()[0].1, 4.0);
1719        approx_eq(out.materialize_f64()[1].0, 5.0);
1720        approx_eq(out.materialize_f64()[1].1, -6.0);
1721
1722        let complex_lhs = ComplexTensor::new(
1723            vec![(1.0, 0.0), (0.0, 0.0), (0.0, 0.0), (1.0, 0.0)],
1724            vec![2, 2],
1725        )
1726        .unwrap();
1727        let real_rhs =
1728            Tensor::new_integer(IntegerStorage::U64(vec![7, 11]), vec![2, 1]).expect("integer rhs");
1729        let result = linsolve_builtin(
1730            Value::ComplexTensor(complex_lhs),
1731            Value::Tensor(real_rhs),
1732            Vec::new(),
1733        )
1734        .expect("linsolve");
1735        let Value::ComplexTensor(out) = result else {
1736            panic!("expected complex tensor output");
1737        };
1738        assert_eq!(out.shape, vec![2, 1]);
1739        approx_eq(out.materialize_f64()[0].0, 7.0);
1740        approx_eq(out.materialize_f64()[0].1, 0.0);
1741        approx_eq(out.materialize_f64()[1].0, 11.0);
1742        approx_eq(out.materialize_f64()[1].1, 0.0);
1743    }
1744
1745    #[test]
1746    fn linsolve_provider_host_helper_reads_typed_integer_storage_exactly() {
1747        let a = Tensor::new_integer(IntegerStorage::U64(vec![2, 1, 1, 2]), vec![2, 2])
1748            .expect("integer lhs");
1749        let b =
1750            Tensor::new_integer(IntegerStorage::U64(vec![4, 5]), vec![2, 1]).expect("integer rhs");
1751
1752        let (solution, _rcond) =
1753            linsolve_host_real_for_provider(&a, &b, &ProviderLinsolveOptions::default())
1754                .expect("provider helper");
1755        assert_eq!(solution.shape, vec![2, 1]);
1756        approx_eq(solution.materialize_f64()[0], 1.0);
1757        approx_eq(solution.materialize_f64()[1], 2.0);
1758    }
1759
1760    #[test]
1761    fn linsolve_general_real_reads_typed_integer_storage_exactly() {
1762        let a = Tensor::new_integer(IntegerStorage::I16(vec![1, 2, 1, 0, 0, 1]), vec![3, 2])
1763            .expect("integer lhs");
1764        let b = Tensor::new_integer(IntegerStorage::I16(vec![3, 2, 1]), vec![3, 1]).expect("rhs");
1765
1766        let (solution, _rcond) =
1767            linsolve_host_real_for_provider(&a, &b, &ProviderLinsolveOptions::default())
1768                .expect("provider helper");
1769
1770        assert_eq!(solution.shape, vec![2, 1]);
1771        approx_eq(solution.materialize_f64()[0], 7.0 / 5.0);
1772        approx_eq(solution.materialize_f64()[1], -2.0 / 5.0);
1773        assert!(solution.integer_storage().is_none());
1774    }
1775
1776    #[test]
1777    fn linsolve_transa_reads_typed_integer_storage_exactly() {
1778        let _accel_guard = test_support::accel_test_lock();
1779        clear_accel_provider_state();
1780        let a = Tensor::new_integer(
1781            IntegerStorage::I16(vec![3, 1, 0, 0, 4, 2, 0, 0, 5]),
1782            vec![3, 3],
1783        )
1784        .expect("integer lhs");
1785        let b = Tensor::new_integer(IntegerStorage::I16(vec![5, 14, 23]), vec![3, 1])
1786            .expect("integer rhs");
1787        let mut opts = StructValue::new();
1788        opts.fields.insert("LT".to_string(), Value::Bool(true));
1789        opts.fields.insert(
1790            "TRANSA".to_string(),
1791            Value::CharArray(CharArray::new_row("T")),
1792        );
1793
1794        let result = linsolve_builtin(
1795            Value::Tensor(a),
1796            Value::Tensor(b),
1797            vec![Value::Struct(opts)],
1798        )
1799        .expect("linsolve");
1800        let tensor = test_support::gather(result).expect("gather");
1801
1802        assert_eq!(tensor.shape, vec![3, 1]);
1803        let expected_a = Tensor::new(
1804            vec![3.0, 1.0, 0.0, 0.0, 4.0, 2.0, 0.0, 0.0, 5.0],
1805            vec![3, 3],
1806        )
1807        .expect("expected lhs");
1808        let expected_b = Tensor::new(vec![5.0, 14.0, 23.0], vec![3, 1]).expect("expected rhs");
1809        let expected_a_transposed = transpose_tensor(&expected_a);
1810        let (expected_tensor, _) = host_linsolve_real(
1811            &expected_a_transposed,
1812            &expected_b,
1813            ProviderLinsolveOptions::default(),
1814        );
1815        for (actual, expected) in tensor
1816            .materialize_f64()
1817            .iter()
1818            .zip(expected_tensor.materialize_f64().iter())
1819        {
1820            approx_eq(*actual, *expected);
1821        }
1822        assert!(tensor.integer_storage().is_none());
1823    }
1824
1825    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1826    #[test]
1827    fn linsolve_lower_triangular_hint() {
1828        let _accel_guard = test_support::accel_test_lock();
1829        clear_accel_provider_state();
1830        let a = Tensor::new(
1831            vec![3.0, -1.0, 4.0, 0.0, 2.0, 1.0, 0.0, 0.0, 5.0],
1832            vec![3, 3],
1833        )
1834        .unwrap();
1835        let b = Tensor::new(vec![9.0, 1.0, 19.0], vec![3, 1]).unwrap();
1836        let mut opts = StructValue::new();
1837        opts.fields.insert("LT".to_string(), Value::Bool(true));
1838        let result = linsolve_builtin(
1839            Value::Tensor(a),
1840            Value::Tensor(b),
1841            vec![Value::Struct(opts)],
1842        )
1843        .expect("linsolve");
1844        let tensor = test_support::gather(result).expect("gather");
1845        assert_eq!(tensor.shape, vec![3, 1]);
1846        approx_eq(tensor.materialize_f64()[0], 3.0);
1847        approx_eq(tensor.materialize_f64()[1], 2.0);
1848        approx_eq(tensor.materialize_f64()[2], 1.0);
1849    }
1850
1851    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1852    #[test]
1853    fn linsolve_transposed_triangular_hint() {
1854        let _accel_guard = test_support::accel_test_lock();
1855        clear_accel_provider_state();
1856        let a = Tensor::new(
1857            vec![3.0, 1.0, 0.0, 0.0, 4.0, 2.0, 0.0, 0.0, 5.0],
1858            vec![3, 3],
1859        )
1860        .unwrap();
1861        let b = Tensor::new(vec![5.0, 14.0, 23.0], vec![3, 1]).unwrap();
1862        let mut opts = StructValue::new();
1863        opts.fields.insert("LT".to_string(), Value::Bool(true));
1864        opts.fields.insert(
1865            "TRANSA".to_string(),
1866            Value::CharArray(CharArray::new_row("T")),
1867        );
1868
1869        let result = linsolve_builtin(
1870            Value::Tensor(a.clone()),
1871            Value::Tensor(b.clone()),
1872            vec![Value::Struct(opts)],
1873        )
1874        .expect("linsolve");
1875        let tensor = test_support::gather(result).expect("gather");
1876        assert_eq!(tensor.shape, vec![3, 1]);
1877
1878        let a_transposed = transpose_tensor(&a);
1879        let (expected_tensor, _) =
1880            host_linsolve_real(&a_transposed, &b, ProviderLinsolveOptions::default());
1881
1882        for (actual, expected) in tensor
1883            .materialize_f64()
1884            .iter()
1885            .zip(expected_tensor.materialize_f64().iter())
1886        {
1887            approx_eq(*actual, *expected);
1888        }
1889    }
1890
1891    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1892    #[test]
1893    fn linsolve_complex_inputs_match_residual() {
1894        let a = ComplexTensor::new(
1895            vec![(2.0, 1.0), (-1.0, 0.0), (1.0, -2.0), (3.0, -2.0)],
1896            vec![2, 2],
1897        )
1898        .unwrap();
1899        let b = ComplexTensor::new(vec![(1.0, 0.0), (4.0, 1.0)], vec![2, 1]).unwrap();
1900        let result = linsolve_builtin(
1901            Value::ComplexTensor(a.clone()),
1902            Value::ComplexTensor(b.clone()),
1903            Vec::new(),
1904        )
1905        .expect("linsolve");
1906        let Value::ComplexTensor(out) = result else {
1907            panic!("expected complex tensor result");
1908        };
1909
1910        let mat_a: Vec<Complex64> = a
1911            .materialize_f64()
1912            .iter()
1913            .map(|&(re, im)| Complex64::new(re, im))
1914            .collect();
1915        let mat_b: Vec<Complex64> = b
1916            .materialize_f64()
1917            .iter()
1918            .map(|&(re, im)| Complex64::new(re, im))
1919            .collect();
1920        let mat_x: Vec<Complex64> = out
1921            .materialize_f64()
1922            .iter()
1923            .map(|&(re, im)| Complex64::new(re, im))
1924            .collect();
1925        let a_mat = DMatrix::from_column_slice(a.rows, a.cols, &mat_a);
1926        let b_mat = DMatrix::from_column_slice(b.rows, b.cols, &mat_b);
1927        let x_mat = DMatrix::from_column_slice(out.rows, out.cols, &mat_x);
1928        let residual = a_mat * x_mat - b_mat;
1929        assert!(residual.norm() < 1e-10, "residual={}", residual.norm());
1930    }
1931
1932    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1933    #[test]
1934    fn linsolve_complex_conjugate_transpose_matches_explicit_reference() {
1935        let a = ComplexTensor::new(
1936            vec![(2.0, 1.0), (0.0, -1.0), (1.0, 2.0), (3.0, 0.5)],
1937            vec![2, 2],
1938        )
1939        .unwrap();
1940        let b = ComplexTensor::new(vec![(1.0, -1.0), (2.0, 0.5)], vec![2, 1]).unwrap();
1941
1942        let mut opts = StructValue::new();
1943        opts.fields.insert(
1944            "TRANSA".to_string(),
1945            Value::CharArray(CharArray::new_row("C")),
1946        );
1947        let result = linsolve_builtin(
1948            Value::ComplexTensor(a.clone()),
1949            Value::ComplexTensor(b.clone()),
1950            vec![Value::Struct(opts)],
1951        )
1952        .expect("linsolve");
1953        let Value::ComplexTensor(out) = result else {
1954            panic!("expected complex tensor result");
1955        };
1956
1957        let mut a_conj_t = transpose_complex(&a);
1958        conjugate_complex_in_place(&mut a_conj_t);
1959        let reference = evaluate(
1960            Value::ComplexTensor(a_conj_t),
1961            Value::ComplexTensor(b.clone()),
1962            SolveOptions::default(),
1963        )
1964        .expect("reference");
1965        let Value::ComplexTensor(expected) = reference.solution() else {
1966            panic!("expected complex tensor reference");
1967        };
1968
1969        assert_eq!(out.shape, expected.shape);
1970        for ((out_re, out_im), (exp_re, exp_im)) in out
1971            .materialize_f64()
1972            .iter()
1973            .zip(expected.materialize_f64().iter())
1974        {
1975            assert!(
1976                (out_re - exp_re).abs() < 1e-10,
1977                "out_re={out_re} exp_re={exp_re}"
1978            );
1979            assert!(
1980                (out_im - exp_im).abs() < 1e-10,
1981                "out_im={out_im} exp_im={exp_im}"
1982            );
1983        }
1984    }
1985
1986    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1987    #[test]
1988    fn linsolve_rcond_enforced() {
1989        let _accel_guard = test_support::accel_test_lock();
1990        clear_accel_provider_state();
1991        let a = Tensor::new(vec![1.0, 1.0, 1.0, 1.0 + 1e-12], vec![2, 2]).unwrap();
1992        let b = Tensor::new(vec![2.0, 2.0 + 1e-12], vec![2, 1]).unwrap();
1993        let mut opts = StructValue::new();
1994        opts.fields.insert("RCOND".to_string(), Value::Num(1e-3));
1995        let err = unwrap_error(
1996            linsolve_builtin(
1997                Value::Tensor(a),
1998                Value::Tensor(b),
1999                vec![Value::Struct(opts)],
2000            )
2001            .expect_err("singular matrix must fail"),
2002        );
2003        assert!(
2004            err.message().contains("singular to working precision"),
2005            "unexpected error message: {err}"
2006        );
2007        assert_eq!(err.identifier(), LINSOLVE_ERROR_INVALID_INPUT.identifier);
2008    }
2009
2010    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2011    #[test]
2012    fn linsolve_options_read_integer_tensor_storage() {
2013        let _accel_guard = test_support::accel_test_lock();
2014        clear_accel_provider_state();
2015        let a = Tensor::new(vec![2.0, 0.0, 1.0, 3.0], vec![2, 2]).unwrap();
2016        let b = Tensor::new(vec![2.0, 7.0], vec![2, 1]).unwrap();
2017        let upper = Tensor::new_integer(IntegerStorage::U8(vec![1]), vec![1, 1]).unwrap();
2018        let rcond = Tensor::new_integer(IntegerStorage::U8(vec![0]), vec![1, 1]).unwrap();
2019        let mut opts = StructValue::new();
2020        opts.fields.insert("UT".to_string(), Value::Tensor(upper));
2021        opts.fields
2022            .insert("RCOND".to_string(), Value::Tensor(rcond));
2023        let result = linsolve_builtin(
2024            Value::Tensor(a),
2025            Value::Tensor(b),
2026            vec![Value::Struct(opts)],
2027        )
2028        .expect("linsolve");
2029        let Value::Tensor(out) = result else {
2030            panic!("expected tensor output");
2031        };
2032        approx_eq(out.materialize_f64()[0], -1.0 / 6.0);
2033        approx_eq(out.materialize_f64()[1], 7.0 / 3.0);
2034    }
2035
2036    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2037    #[test]
2038    fn linsolve_unknown_option_identifier() {
2039        let a = Tensor::new(vec![1.0, 0.0, 0.0, 1.0], vec![2, 2]).unwrap();
2040        let b = Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap();
2041        let mut opts = StructValue::new();
2042        opts.fields.insert("UNKNOWN".to_string(), Value::Bool(true));
2043        let err = unwrap_error(
2044            linsolve_builtin(
2045                Value::Tensor(a),
2046                Value::Tensor(b),
2047                vec![Value::Struct(opts)],
2048            )
2049            .expect_err("unknown option should fail"),
2050        );
2051        assert!(err.message().contains("unknown option"));
2052        assert_eq!(err.identifier(), LINSOLVE_ERROR_INVALID_ARGUMENT.identifier);
2053    }
2054
2055    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2056    #[test]
2057    fn linsolve_output_count_limit_identifier() {
2058        let a = Tensor::new(vec![1.0], vec![1, 1]).unwrap();
2059        let b = Tensor::new(vec![2.0], vec![1, 1]).unwrap();
2060        let _guard = crate::output_count::push_output_count(Some(3));
2061        let err = unwrap_error(
2062            linsolve_builtin(Value::Tensor(a), Value::Tensor(b), Vec::new())
2063                .expect_err("three outputs should fail"),
2064        );
2065        assert!(err.message().contains("at most two outputs"));
2066        assert_eq!(err.identifier(), LINSOLVE_ERROR_INVALID_ARGUMENT.identifier);
2067    }
2068
2069    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2070    #[test]
2071    fn linsolve_recovers_rcond_output() {
2072        let _accel_guard = test_support::accel_test_lock();
2073        clear_accel_provider_state();
2074        let a = Tensor::new(vec![1.0, 0.0, 0.0, 1.0], vec![2, 2]).unwrap();
2075        let b = Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap();
2076        let eval = evaluate_args(Value::Tensor(a.clone()), Value::Tensor(b.clone()), &[])
2077            .expect("evaluate");
2078        let solution_tensor = match eval.solution() {
2079            Value::Tensor(sol) => sol.clone(),
2080            Value::GpuTensor(handle) => {
2081                test_support::gather(Value::GpuTensor(handle.clone())).expect("gather solution")
2082            }
2083            other => panic!("unexpected solution value {other:?}"),
2084        };
2085        assert_eq!(solution_tensor.shape, vec![2, 1]);
2086        approx_eq(solution_tensor.materialize_f64()[0], 1.0);
2087        approx_eq(solution_tensor.materialize_f64()[1], 2.0);
2088
2089        let rcond_value = match eval.reciprocal_condition() {
2090            Value::Num(r) => r,
2091            Value::GpuTensor(handle) => {
2092                let gathered =
2093                    test_support::gather(Value::GpuTensor(handle.clone())).expect("gather rcond");
2094                gathered.materialize_f64()[0]
2095            }
2096            other => panic!("unexpected rcond value {other:?}"),
2097        };
2098        approx_eq(rcond_value, 1.0);
2099    }
2100
2101    #[test]
2102    fn linsolve_rectangular_second_output_reports_rank() {
2103        let _accel_guard = test_support::accel_test_lock();
2104        clear_accel_provider_state();
2105        let a = Tensor::new(vec![1.0, 2.0, 3.0, 2.0, 4.0, 6.0], vec![3, 2]).unwrap();
2106        let b = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
2107
2108        let eval = evaluate_args(Value::Tensor(a), Value::Tensor(b), &[]).expect("evaluate");
2109
2110        assert_eq!(eval.reciprocal_condition(), Value::Num(1.0));
2111    }
2112
2113    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2114    #[test]
2115    fn gpu_round_trip_matches_cpu() {
2116        test_support::with_test_provider(|provider| {
2117            let a = Tensor::new(vec![2.0, 1.0, 1.0, 3.0], vec![2, 2]).unwrap();
2118            let b = Tensor::new(vec![4.0, 5.0], vec![2, 1]).unwrap();
2119
2120            let cpu = linsolve_builtin(
2121                Value::Tensor(a.clone()),
2122                Value::Tensor(b.clone()),
2123                Vec::new(),
2124            )
2125            .expect("cpu linsolve");
2126            let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
2127
2128            let view_a = HostTensorView {
2129                data: &a.materialize_f64(),
2130                shape: &a.shape,
2131            };
2132            let view_b = HostTensorView {
2133                data: &b.materialize_f64(),
2134                shape: &b.shape,
2135            };
2136            let ha = provider.upload(&view_a).expect("upload A");
2137            let hb = provider.upload(&view_b).expect("upload B");
2138
2139            let gpu_value = linsolve_builtin(
2140                Value::GpuTensor(ha.clone()),
2141                Value::GpuTensor(hb.clone()),
2142                Vec::new(),
2143            )
2144            .expect("gpu linsolve");
2145            let gathered = test_support::gather(gpu_value).expect("gather");
2146            let _ = provider.free(&ha);
2147            let _ = provider.free(&hb);
2148
2149            assert_eq!(gathered.shape, cpu_tensor.shape);
2150            for (gpu, cpu) in gathered
2151                .materialize_f64()
2152                .iter()
2153                .zip(cpu_tensor.materialize_f64().iter())
2154            {
2155                assert!((gpu - cpu).abs() < 1e-12);
2156            }
2157        });
2158    }
2159
2160    #[test]
2161    fn host_inputs_auto_promote_into_provider_solve_path() {
2162        test_support::with_test_provider(|provider| {
2163            provider.reset_telemetry();
2164            let a = Tensor::new(vec![2.0, 1.0, 1.0, 3.0], vec![2, 2]).unwrap();
2165            let b = Tensor::new(vec![4.0, 5.0], vec![2, 1]).unwrap();
2166            let _ = linsolve_builtin(Value::Tensor(a), Value::Tensor(b), Vec::new())
2167                .expect("host linsolve");
2168            let telemetry = provider.telemetry_snapshot();
2169            assert!(telemetry.linsolve.count >= 1);
2170            assert!(fallback_count(&telemetry, "linsolve:host_reupload") >= 1);
2171            assert!(telemetry.upload_bytes > 0);
2172            assert!(telemetry.download_bytes > 0);
2173        });
2174    }
2175
2176    #[test]
2177    fn typed_integer_host_inputs_keep_the_double_solve_boundary_on_host() {
2178        test_support::with_test_provider(|provider| {
2179            provider.reset_telemetry();
2180            let a = Tensor::new_integer(IntegerStorage::U64(vec![2, 1, 1, 2]), vec![2, 2])
2181                .expect("integer lhs");
2182            let b = Tensor::new_integer(IntegerStorage::U64(vec![4, 5]), vec![2, 1])
2183                .expect("integer rhs");
2184
2185            let result = linsolve_builtin(Value::Tensor(a), Value::Tensor(b), Vec::new())
2186                .expect("host provider linsolve");
2187            let tensor = test_support::gather(result).expect("gather");
2188            assert_eq!(tensor.shape, vec![2, 1]);
2189            approx_eq(tensor.materialize_f64()[0], 1.0);
2190            approx_eq(tensor.materialize_f64()[1], 2.0);
2191
2192            let telemetry = provider.telemetry_snapshot();
2193            assert_eq!(telemetry.linsolve.count, 0);
2194            assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
2195        });
2196    }
2197
2198    #[test]
2199    fn provider_telemetry_records_gpu_host_reupload_path() {
2200        test_support::with_test_provider(|provider| {
2201            provider.reset_telemetry();
2202            let a = Tensor::new(vec![2.0, 1.0, 1.0, 3.0], vec![2, 2]).unwrap();
2203            let b = Tensor::new(vec![4.0, 5.0], vec![2, 1]).unwrap();
2204            let ha = provider
2205                .upload(&HostTensorView {
2206                    data: &a.materialize_f64(),
2207                    shape: &a.shape,
2208                })
2209                .expect("upload A");
2210            let hb = provider
2211                .upload(&HostTensorView {
2212                    data: &b.materialize_f64(),
2213                    shape: &b.shape,
2214                })
2215                .expect("upload B");
2216
2217            let _ = linsolve_builtin(
2218                Value::GpuTensor(ha.clone()),
2219                Value::GpuTensor(hb.clone()),
2220                Vec::new(),
2221            )
2222            .expect("gpu linsolve");
2223
2224            let telemetry = provider.telemetry_snapshot();
2225            assert_eq!(telemetry.linsolve.count, 1);
2226            assert!(telemetry.upload_bytes > 0);
2227            assert!(telemetry.download_bytes > 0);
2228            assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 1);
2229
2230            let _ = provider.free(&ha);
2231            let _ = provider.free(&hb);
2232        });
2233    }
2234
2235    #[test]
2236    fn scalar_gpu_inputs_fall_back_without_provider_solve_dispatch() {
2237        test_support::with_test_provider(|provider| {
2238            provider.reset_telemetry();
2239            let a = Tensor::new(vec![2.0], vec![1, 1]).unwrap();
2240            let b = Tensor::new(vec![6.0], vec![1, 1]).unwrap();
2241            let ha = provider
2242                .upload(&HostTensorView {
2243                    data: &a.materialize_f64(),
2244                    shape: &a.shape,
2245                })
2246                .expect("upload A");
2247            let hb = provider
2248                .upload(&HostTensorView {
2249                    data: &b.materialize_f64(),
2250                    shape: &b.shape,
2251                })
2252                .expect("upload B");
2253
2254            let result = linsolve_builtin(
2255                Value::GpuTensor(ha.clone()),
2256                Value::GpuTensor(hb.clone()),
2257                Vec::new(),
2258            )
2259            .expect("fallback linsolve");
2260            let gathered = test_support::gather(result).expect("gather fallback");
2261            assert_eq!(gathered.materialize_f64(), vec![3.0]);
2262
2263            let telemetry = provider.telemetry_snapshot();
2264            assert_eq!(telemetry.linsolve.count, 0);
2265            assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
2266            assert!(telemetry.download_bytes > 0);
2267
2268            let _ = provider.free(&ha);
2269            let _ = provider.free(&hb);
2270        });
2271    }
2272
2273    #[cfg(feature = "wgpu")]
2274    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2275    #[test]
2276    fn wgpu_square_linsolve_avoids_host_reupload_fallback() {
2277        let _accel_guard = test_support::accel_test_lock();
2278        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2279            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2280        ) else {
2281            return;
2282        };
2283        if provider.precision() != runmat_accelerate_api::ProviderPrecision::F32 {
2284            return;
2285        }
2286        let a = Tensor::new(vec![3.0, 1.0, 2.0, 4.0], vec![2, 2]).unwrap();
2287        let b = Tensor::new(vec![7.0, 8.0], vec![2, 1]).unwrap();
2288
2289        let cpu = linsolve_builtin(
2290            Value::Tensor(a.clone()),
2291            Value::Tensor(b.clone()),
2292            Vec::new(),
2293        )
2294        .expect("cpu linsolve");
2295        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
2296        provider.reset_telemetry();
2297
2298        let ha = provider
2299            .upload(&HostTensorView {
2300                data: &a.materialize_f64(),
2301                shape: &a.shape,
2302            })
2303            .expect("upload A");
2304        let hb = provider
2305            .upload(&HostTensorView {
2306                data: &b.materialize_f64(),
2307                shape: &b.shape,
2308            })
2309            .expect("upload B");
2310
2311        let _output_guard = crate::output_count::push_output_count(Some(1));
2312        let gpu_value = linsolve_builtin(
2313            Value::GpuTensor(ha.clone()),
2314            Value::GpuTensor(hb.clone()),
2315            Vec::new(),
2316        )
2317        .expect("gpu square linsolve");
2318        let gpu_solution = match gpu_value {
2319            Value::OutputList(mut outputs) => outputs.remove(0),
2320            other => other,
2321        };
2322        let gathered = test_support::gather(gpu_solution).expect("gather");
2323        let _ = provider.free(&ha);
2324        let _ = provider.free(&hb);
2325
2326        assert_eq!(gathered.shape, cpu_tensor.shape);
2327        for (gpu, cpu) in gathered
2328            .materialize_f64()
2329            .iter()
2330            .zip(cpu_tensor.materialize_f64().iter())
2331        {
2332            assert!((gpu - cpu).abs() < 1e-4);
2333        }
2334
2335        let telemetry = provider.telemetry_snapshot();
2336        assert_eq!(telemetry.linsolve.count, 1);
2337        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
2338        assert_eq!(kernel_launch_count(&telemetry, "linsolve_posdef_chol"), 0);
2339        assert_eq!(kernel_launch_count(&telemetry, "linsolve_tall_qr"), 1);
2340    }
2341
2342    #[cfg(feature = "wgpu")]
2343    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2344    #[test]
2345    fn wgpu_square_linsolve_uses_device_path_without_output_count() {
2346        let _accel_guard = test_support::accel_test_lock();
2347        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2348            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2349        ) else {
2350            return;
2351        };
2352        if provider.precision() != runmat_accelerate_api::ProviderPrecision::F32 {
2353            return;
2354        }
2355        let a = Tensor::new(vec![3.0, 1.0, 2.0, 4.0], vec![2, 2]).unwrap();
2356        let b = Tensor::new(vec![7.0, 8.0], vec![2, 1]).unwrap();
2357
2358        let cpu = linsolve_builtin(
2359            Value::Tensor(a.clone()),
2360            Value::Tensor(b.clone()),
2361            Vec::new(),
2362        )
2363        .expect("cpu linsolve");
2364        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
2365        provider.reset_telemetry();
2366
2367        let ha = provider
2368            .upload(&HostTensorView {
2369                data: &a.materialize_f64(),
2370                shape: &a.shape,
2371            })
2372            .expect("upload A");
2373        let hb = provider
2374            .upload(&HostTensorView {
2375                data: &b.materialize_f64(),
2376                shape: &b.shape,
2377            })
2378            .expect("upload B");
2379
2380        let gpu_value = linsolve_builtin(
2381            Value::GpuTensor(ha.clone()),
2382            Value::GpuTensor(hb.clone()),
2383            Vec::new(),
2384        )
2385        .expect("gpu square linsolve");
2386        let gathered = test_support::gather(gpu_value).expect("gather");
2387        let _ = provider.free(&ha);
2388        let _ = provider.free(&hb);
2389
2390        assert_eq!(gathered.shape, cpu_tensor.shape);
2391        for (gpu, cpu) in gathered
2392            .materialize_f64()
2393            .iter()
2394            .zip(cpu_tensor.materialize_f64().iter())
2395        {
2396            assert!((gpu - cpu).abs() < 1e-4, "gpu={gpu} cpu={cpu}");
2397        }
2398
2399        let telemetry = provider.telemetry_snapshot();
2400        assert_eq!(telemetry.linsolve.count, 1);
2401        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
2402        assert_eq!(kernel_launch_count(&telemetry, "linsolve_tall_qr"), 1);
2403    }
2404
2405    #[cfg(feature = "wgpu")]
2406    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2407    #[test]
2408    fn wgpu_square_linsolve_recovers_rcond_output_on_device() {
2409        let _accel_guard = test_support::accel_test_lock();
2410        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2411            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2412        ) else {
2413            return;
2414        };
2415        if provider.precision() != runmat_accelerate_api::ProviderPrecision::F32 {
2416            return;
2417        }
2418        let a = Tensor::new(vec![3.0, 1.0, 2.0, 4.0], vec![2, 2]).unwrap();
2419        let b = Tensor::new(vec![7.0, 8.0], vec![2, 1]).unwrap();
2420
2421        let (_, cpu_rcond) = host_linsolve_real(&a, &b, ProviderLinsolveOptions::default());
2422        provider.reset_telemetry();
2423
2424        let ha = provider
2425            .upload(&HostTensorView {
2426                data: &a.materialize_f64(),
2427                shape: &a.shape,
2428            })
2429            .expect("upload A");
2430        let hb = provider
2431            .upload(&HostTensorView {
2432                data: &b.materialize_f64(),
2433                shape: &b.shape,
2434            })
2435            .expect("upload B");
2436
2437        let _output_guard = crate::output_count::push_output_count(Some(2));
2438        let gpu_value = linsolve_builtin(
2439            Value::GpuTensor(ha.clone()),
2440            Value::GpuTensor(hb.clone()),
2441            Vec::new(),
2442        )
2443        .expect("gpu square linsolve");
2444        let outputs = match gpu_value {
2445            Value::OutputList(outputs) => outputs,
2446            other => panic!("expected output list, got {other:?}"),
2447        };
2448        assert_eq!(outputs.len(), 2);
2449        let gathered = test_support::gather(outputs[0].clone()).expect("gather");
2450        let gpu_rcond = match &outputs[1] {
2451            Value::Num(value) => *value,
2452            other => panic!("unexpected gpu rcond {other:?}"),
2453        };
2454        let _ = provider.free(&ha);
2455        let _ = provider.free(&hb);
2456
2457        assert_eq!(gathered.shape, vec![2, 1]);
2458        assert!(
2459            (gpu_rcond - cpu_rcond).abs() < 1e-4,
2460            "gpu={gpu_rcond} cpu={cpu_rcond}"
2461        );
2462
2463        let telemetry = provider.telemetry_snapshot();
2464        assert_eq!(telemetry.linsolve.count, 1);
2465        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
2466        assert_eq!(kernel_launch_count(&telemetry, "linsolve_tall_qr"), 1);
2467    }
2468
2469    #[cfg(feature = "wgpu")]
2470    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2471    #[test]
2472    fn wgpu_square_linsolve_with_rcond_option_stays_on_device() {
2473        let _accel_guard = test_support::accel_test_lock();
2474        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2475            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2476        ) else {
2477            return;
2478        };
2479        if provider.precision() != runmat_accelerate_api::ProviderPrecision::F32 {
2480            return;
2481        }
2482
2483        let a = Tensor::new(vec![3.0, 1.0, 2.0, 4.0], vec![2, 2]).unwrap();
2484        let b = Tensor::new(vec![7.0, 8.0], vec![2, 1]).unwrap();
2485        let mut cpu_opts = StructValue::new();
2486        cpu_opts
2487            .fields
2488            .insert("RCOND".to_string(), Value::Num(0.05));
2489        let cpu = linsolve_builtin(
2490            Value::Tensor(a.clone()),
2491            Value::Tensor(b.clone()),
2492            vec![Value::Struct(cpu_opts)],
2493        )
2494        .expect("cpu linsolve");
2495        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
2496        provider.reset_telemetry();
2497
2498        let ha = provider
2499            .upload(&HostTensorView {
2500                data: &a.materialize_f64(),
2501                shape: &a.shape,
2502            })
2503            .expect("upload A");
2504        let hb = provider
2505            .upload(&HostTensorView {
2506                data: &b.materialize_f64(),
2507                shape: &b.shape,
2508            })
2509            .expect("upload B");
2510
2511        let _output_guard = crate::output_count::push_output_count(Some(1));
2512        let mut gpu_opts = StructValue::new();
2513        gpu_opts
2514            .fields
2515            .insert("RCOND".to_string(), Value::Num(0.05));
2516        let gpu_value = linsolve_builtin(
2517            Value::GpuTensor(ha.clone()),
2518            Value::GpuTensor(hb.clone()),
2519            vec![Value::Struct(gpu_opts)],
2520        )
2521        .expect("gpu square linsolve");
2522        let gpu_solution = match gpu_value {
2523            Value::OutputList(mut outputs) => outputs.remove(0),
2524            other => other,
2525        };
2526        let gathered = test_support::gather(gpu_solution).expect("gather");
2527        let _ = provider.free(&ha);
2528        let _ = provider.free(&hb);
2529
2530        assert_eq!(gathered.shape, cpu_tensor.shape);
2531        for (gpu, cpu) in gathered
2532            .materialize_f64()
2533            .iter()
2534            .zip(cpu_tensor.materialize_f64().iter())
2535        {
2536            assert!((gpu - cpu).abs() < 1e-4, "gpu={gpu} cpu={cpu}");
2537        }
2538
2539        let telemetry = provider.telemetry_snapshot();
2540        assert_eq!(telemetry.linsolve.count, 1);
2541        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
2542        assert_eq!(kernel_launch_count(&telemetry, "linsolve_tall_qr"), 1);
2543    }
2544
2545    #[cfg(feature = "wgpu")]
2546    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2547    #[test]
2548    fn wgpu_tall_linsolve_avoids_host_reupload_fallback() {
2549        let _accel_guard = test_support::accel_test_lock();
2550        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2551            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2552        ) else {
2553            return;
2554        };
2555        if provider.precision() != runmat_accelerate_api::ProviderPrecision::F32 {
2556            return;
2557        }
2558        let a = Tensor::new(vec![1.0, 0.0, 1.0, 0.0, 1.0, 1.0], vec![3, 2]).unwrap();
2559        let b = Tensor::new(vec![1.0, 2.0, 2.0], vec![3, 1]).unwrap();
2560
2561        let cpu = linsolve_builtin(
2562            Value::Tensor(a.clone()),
2563            Value::Tensor(b.clone()),
2564            Vec::new(),
2565        )
2566        .expect("cpu linsolve");
2567        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
2568        provider.reset_telemetry();
2569
2570        let ha = provider
2571            .upload(&HostTensorView {
2572                data: &a.materialize_f64(),
2573                shape: &a.shape,
2574            })
2575            .expect("upload A");
2576        let hb = provider
2577            .upload(&HostTensorView {
2578                data: &b.materialize_f64(),
2579                shape: &b.shape,
2580            })
2581            .expect("upload B");
2582
2583        let _output_guard = crate::output_count::push_output_count(Some(1));
2584        let gpu_value = linsolve_builtin(
2585            Value::GpuTensor(ha.clone()),
2586            Value::GpuTensor(hb.clone()),
2587            Vec::new(),
2588        )
2589        .expect("gpu tall linsolve");
2590        let gpu_solution = match gpu_value {
2591            Value::OutputList(mut outputs) => outputs.remove(0),
2592            other => other,
2593        };
2594        let gathered = test_support::gather(gpu_solution).expect("gather");
2595        let _ = provider.free(&ha);
2596        let _ = provider.free(&hb);
2597
2598        assert_eq!(gathered.shape, cpu_tensor.shape);
2599        for (gpu, cpu) in gathered
2600            .materialize_f64()
2601            .iter()
2602            .zip(cpu_tensor.materialize_f64().iter())
2603        {
2604            assert!((gpu - cpu).abs() < 1e-4, "gpu={gpu} cpu={cpu}");
2605        }
2606
2607        let telemetry = provider.telemetry_snapshot();
2608        assert_eq!(telemetry.linsolve.count, 1);
2609        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
2610    }
2611
2612    #[cfg(feature = "wgpu")]
2613    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2614    #[test]
2615    fn wgpu_posdef_linsolve_avoids_host_reupload_fallback() {
2616        let _accel_guard = test_support::accel_test_lock();
2617        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2618            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2619        ) else {
2620            return;
2621        };
2622        if provider.precision() != runmat_accelerate_api::ProviderPrecision::F32 {
2623            return;
2624        }
2625        let a = Tensor::new(vec![4.0, 1.0, 1.0, 3.0], vec![2, 2]).unwrap();
2626        let b = Tensor::new(vec![7.0, 8.0], vec![2, 1]).unwrap();
2627
2628        let mut cpu_opts = StructValue::new();
2629        cpu_opts
2630            .fields
2631            .insert("POSDEF".to_string(), Value::Bool(true));
2632        let cpu = linsolve_builtin(
2633            Value::Tensor(a.clone()),
2634            Value::Tensor(b.clone()),
2635            vec![Value::Struct(cpu_opts)],
2636        )
2637        .expect("cpu linsolve");
2638        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
2639        let (_, cpu_rcond) = host_linsolve_real(
2640            &a,
2641            &b,
2642            ProviderLinsolveOptions {
2643                posdef: true,
2644                ..Default::default()
2645            },
2646        );
2647        provider.reset_telemetry();
2648
2649        let ha = provider
2650            .upload(&HostTensorView {
2651                data: &a.materialize_f64(),
2652                shape: &a.shape,
2653            })
2654            .expect("upload A");
2655        let hb = provider
2656            .upload(&HostTensorView {
2657                data: &b.materialize_f64(),
2658                shape: &b.shape,
2659            })
2660            .expect("upload B");
2661
2662        let _output_guard = crate::output_count::push_output_count(Some(2));
2663        let mut gpu_opts = StructValue::new();
2664        gpu_opts
2665            .fields
2666            .insert("POSDEF".to_string(), Value::Bool(true));
2667        let gpu_value = linsolve_builtin(
2668            Value::GpuTensor(ha.clone()),
2669            Value::GpuTensor(hb.clone()),
2670            vec![Value::Struct(gpu_opts)],
2671        )
2672        .expect("gpu posdef linsolve");
2673        let mut outputs = match gpu_value {
2674            Value::OutputList(outputs) => outputs,
2675            other => panic!("expected output list, got {other:?}"),
2676        };
2677        let gpu_rcond = match outputs.remove(1) {
2678            Value::Num(value) => value,
2679            other => panic!("unexpected rcond value {other:?}"),
2680        };
2681        let gpu_solution = outputs.remove(0);
2682        let gathered = test_support::gather(gpu_solution).expect("gather");
2683        let _ = provider.free(&ha);
2684        let _ = provider.free(&hb);
2685
2686        assert_eq!(gathered.shape, cpu_tensor.shape);
2687        for (gpu, cpu) in gathered
2688            .materialize_f64()
2689            .iter()
2690            .zip(cpu_tensor.materialize_f64().iter())
2691        {
2692            assert!((gpu - cpu).abs() < 1e-4, "gpu={gpu} cpu={cpu}");
2693        }
2694        assert!(
2695            (gpu_rcond - cpu_rcond).abs() < 1e-4,
2696            "gpu={gpu_rcond} cpu={cpu_rcond}"
2697        );
2698
2699        let telemetry = provider.telemetry_snapshot();
2700        assert_eq!(telemetry.linsolve.count, 1);
2701        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
2702        assert_eq!(kernel_launch_count(&telemetry, "linsolve_posdef_chol"), 1);
2703        assert_eq!(kernel_launch_count(&telemetry, "linsolve_tall_qr"), 0);
2704    }
2705
2706    #[cfg(feature = "wgpu")]
2707    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2708    #[test]
2709    fn wgpu_transposed_posdef_linsolve_uses_cholesky_path() {
2710        let _accel_guard = test_support::accel_test_lock();
2711        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2712            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2713        ) else {
2714            return;
2715        };
2716        if provider.precision() != runmat_accelerate_api::ProviderPrecision::F32 {
2717            return;
2718        }
2719        let a = Tensor::new(vec![6.0, 2.0, 2.0, 5.0], vec![2, 2]).unwrap();
2720        let b = Tensor::new(vec![8.0, 9.0], vec![2, 1]).unwrap();
2721
2722        let mut cpu_opts = StructValue::new();
2723        cpu_opts
2724            .fields
2725            .insert("POSDEF".to_string(), Value::Bool(true));
2726        cpu_opts.fields.insert(
2727            "TRANSA".to_string(),
2728            Value::CharArray(CharArray::new_row("T")),
2729        );
2730        let cpu = linsolve_builtin(
2731            Value::Tensor(a.clone()),
2732            Value::Tensor(b.clone()),
2733            vec![Value::Struct(cpu_opts)],
2734        )
2735        .expect("cpu linsolve");
2736        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
2737        provider.reset_telemetry();
2738
2739        let ha = provider
2740            .upload(&HostTensorView {
2741                data: &a.materialize_f64(),
2742                shape: &a.shape,
2743            })
2744            .expect("upload A");
2745        let hb = provider
2746            .upload(&HostTensorView {
2747                data: &b.materialize_f64(),
2748                shape: &b.shape,
2749            })
2750            .expect("upload B");
2751
2752        let _output_guard = crate::output_count::push_output_count(Some(1));
2753        let mut gpu_opts = StructValue::new();
2754        gpu_opts
2755            .fields
2756            .insert("POSDEF".to_string(), Value::Bool(true));
2757        gpu_opts.fields.insert(
2758            "TRANSA".to_string(),
2759            Value::CharArray(CharArray::new_row("T")),
2760        );
2761        let gpu_value = linsolve_builtin(
2762            Value::GpuTensor(ha.clone()),
2763            Value::GpuTensor(hb.clone()),
2764            vec![Value::Struct(gpu_opts)],
2765        )
2766        .expect("gpu transposed posdef linsolve");
2767        let gpu_solution = match gpu_value {
2768            Value::OutputList(mut outputs) => outputs.remove(0),
2769            other => other,
2770        };
2771        let gathered = test_support::gather(gpu_solution).expect("gather");
2772        let _ = provider.free(&ha);
2773        let _ = provider.free(&hb);
2774
2775        assert_eq!(gathered.shape, cpu_tensor.shape);
2776        for (gpu, cpu) in gathered
2777            .materialize_f64()
2778            .iter()
2779            .zip(cpu_tensor.materialize_f64().iter())
2780        {
2781            assert!((gpu - cpu).abs() < 1e-4, "gpu={gpu} cpu={cpu}");
2782        }
2783
2784        let telemetry = provider.telemetry_snapshot();
2785        assert_eq!(telemetry.linsolve.count, 1);
2786        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
2787        assert_eq!(kernel_launch_count(&telemetry, "linsolve_posdef_chol"), 1);
2788        assert_eq!(kernel_launch_count(&telemetry, "linsolve_tall_qr"), 0);
2789    }
2790
2791    #[cfg(feature = "wgpu")]
2792    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2793    #[test]
2794    fn wgpu_symmetric_linsolve_avoids_host_reupload_fallback() {
2795        let _accel_guard = test_support::accel_test_lock();
2796        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2797            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2798        ) else {
2799            return;
2800        };
2801        if provider.precision() != runmat_accelerate_api::ProviderPrecision::F32 {
2802            return;
2803        }
2804        let a = Tensor::new(vec![5.0, 2.0, 2.0, 6.0], vec![2, 2]).unwrap();
2805        let b = Tensor::new(vec![9.0, 8.0], vec![2, 1]).unwrap();
2806
2807        let mut cpu_opts = StructValue::new();
2808        cpu_opts.fields.insert("SYM".to_string(), Value::Bool(true));
2809        let cpu = linsolve_builtin(
2810            Value::Tensor(a.clone()),
2811            Value::Tensor(b.clone()),
2812            vec![Value::Struct(cpu_opts)],
2813        )
2814        .expect("cpu linsolve");
2815        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
2816        provider.reset_telemetry();
2817
2818        let ha = provider
2819            .upload(&HostTensorView {
2820                data: &a.materialize_f64(),
2821                shape: &a.shape,
2822            })
2823            .expect("upload A");
2824        let hb = provider
2825            .upload(&HostTensorView {
2826                data: &b.materialize_f64(),
2827                shape: &b.shape,
2828            })
2829            .expect("upload B");
2830
2831        let _output_guard = crate::output_count::push_output_count(Some(1));
2832        let mut gpu_opts = StructValue::new();
2833        gpu_opts.fields.insert("SYM".to_string(), Value::Bool(true));
2834        let gpu_value = linsolve_builtin(
2835            Value::GpuTensor(ha.clone()),
2836            Value::GpuTensor(hb.clone()),
2837            vec![Value::Struct(gpu_opts)],
2838        )
2839        .expect("gpu symmetric linsolve");
2840        let gpu_solution = match gpu_value {
2841            Value::OutputList(mut outputs) => outputs.remove(0),
2842            other => other,
2843        };
2844        let gathered = test_support::gather(gpu_solution).expect("gather");
2845        let _ = provider.free(&ha);
2846        let _ = provider.free(&hb);
2847
2848        assert_eq!(gathered.shape, cpu_tensor.shape);
2849        for (gpu, cpu) in gathered
2850            .materialize_f64()
2851            .iter()
2852            .zip(cpu_tensor.materialize_f64().iter())
2853        {
2854            assert!((gpu - cpu).abs() < 1e-4, "gpu={gpu} cpu={cpu}");
2855        }
2856
2857        let telemetry = provider.telemetry_snapshot();
2858        assert_eq!(telemetry.linsolve.count, 1);
2859        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
2860    }
2861
2862    #[cfg(feature = "wgpu")]
2863    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2864    #[test]
2865    fn wgpu_transposed_square_linsolve_avoids_host_reupload_fallback() {
2866        let _accel_guard = test_support::accel_test_lock();
2867        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2868            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2869        ) else {
2870            return;
2871        };
2872        if provider.precision() != runmat_accelerate_api::ProviderPrecision::F32 {
2873            return;
2874        }
2875        let a = Tensor::new(vec![3.0, 1.0, 2.0, 4.0], vec![2, 2]).unwrap();
2876        let b = Tensor::new(vec![5.0, 14.0], vec![2, 1]).unwrap();
2877
2878        let mut cpu_opts = StructValue::new();
2879        cpu_opts.fields.insert(
2880            "TRANSA".to_string(),
2881            Value::CharArray(CharArray::new_row("T")),
2882        );
2883        let cpu = linsolve_builtin(
2884            Value::Tensor(a.clone()),
2885            Value::Tensor(b.clone()),
2886            vec![Value::Struct(cpu_opts)],
2887        )
2888        .expect("cpu linsolve");
2889        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
2890        provider.reset_telemetry();
2891
2892        let ha = provider
2893            .upload(&HostTensorView {
2894                data: &a.materialize_f64(),
2895                shape: &a.shape,
2896            })
2897            .expect("upload A");
2898        let hb = provider
2899            .upload(&HostTensorView {
2900                data: &b.materialize_f64(),
2901                shape: &b.shape,
2902            })
2903            .expect("upload B");
2904
2905        let _output_guard = crate::output_count::push_output_count(Some(1));
2906        let mut gpu_opts = StructValue::new();
2907        gpu_opts.fields.insert(
2908            "TRANSA".to_string(),
2909            Value::CharArray(CharArray::new_row("T")),
2910        );
2911        let gpu_value = linsolve_builtin(
2912            Value::GpuTensor(ha.clone()),
2913            Value::GpuTensor(hb.clone()),
2914            vec![Value::Struct(gpu_opts)],
2915        )
2916        .expect("gpu transposed square linsolve");
2917        let gpu_solution = match gpu_value {
2918            Value::OutputList(mut outputs) => outputs.remove(0),
2919            other => other,
2920        };
2921        let gathered = test_support::gather(gpu_solution).expect("gather");
2922        let _ = provider.free(&ha);
2923        let _ = provider.free(&hb);
2924
2925        assert_eq!(gathered.shape, cpu_tensor.shape);
2926        for (gpu, cpu) in gathered
2927            .materialize_f64()
2928            .iter()
2929            .zip(cpu_tensor.materialize_f64().iter())
2930        {
2931            assert!((gpu - cpu).abs() < 1e-4, "gpu={gpu} cpu={cpu}");
2932        }
2933
2934        let telemetry = provider.telemetry_snapshot();
2935        assert_eq!(telemetry.linsolve.count, 1);
2936        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
2937    }
2938
2939    #[cfg(feature = "wgpu")]
2940    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2941    #[test]
2942    fn wgpu_conjugate_square_linsolve_avoids_host_reupload_fallback_for_real_inputs() {
2943        let _accel_guard = test_support::accel_test_lock();
2944        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2945            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2946        ) else {
2947            return;
2948        };
2949        if provider.precision() != runmat_accelerate_api::ProviderPrecision::F32 {
2950            return;
2951        }
2952
2953        let a = Tensor::new(vec![3.0, 1.0, 2.0, 4.0], vec![2, 2]).unwrap();
2954        let b = Tensor::new(vec![5.0, 14.0], vec![2, 1]).unwrap();
2955        let mut cpu_opts = StructValue::new();
2956        cpu_opts.fields.insert(
2957            "TRANSA".to_string(),
2958            Value::CharArray(CharArray::new_row("C")),
2959        );
2960        let cpu = linsolve_builtin(
2961            Value::Tensor(a.clone()),
2962            Value::Tensor(b.clone()),
2963            vec![Value::Struct(cpu_opts)],
2964        )
2965        .expect("cpu linsolve");
2966        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
2967        provider.reset_telemetry();
2968
2969        let ha = provider
2970            .upload(&HostTensorView {
2971                data: &a.materialize_f64(),
2972                shape: &a.shape,
2973            })
2974            .expect("upload A");
2975        let hb = provider
2976            .upload(&HostTensorView {
2977                data: &b.materialize_f64(),
2978                shape: &b.shape,
2979            })
2980            .expect("upload B");
2981
2982        let _output_guard = crate::output_count::push_output_count(Some(1));
2983        let mut gpu_opts = StructValue::new();
2984        gpu_opts.fields.insert(
2985            "TRANSA".to_string(),
2986            Value::CharArray(CharArray::new_row("C")),
2987        );
2988        let gpu_value = linsolve_builtin(
2989            Value::GpuTensor(ha.clone()),
2990            Value::GpuTensor(hb.clone()),
2991            vec![Value::Struct(gpu_opts)],
2992        )
2993        .expect("gpu conjugate square linsolve");
2994        let gpu_solution = match gpu_value {
2995            Value::OutputList(mut outputs) => outputs.remove(0),
2996            other => other,
2997        };
2998        let gathered = test_support::gather(gpu_solution).expect("gather");
2999        let _ = provider.free(&ha);
3000        let _ = provider.free(&hb);
3001
3002        assert_eq!(gathered.shape, cpu_tensor.shape);
3003        for (gpu, cpu) in gathered
3004            .materialize_f64()
3005            .iter()
3006            .zip(cpu_tensor.materialize_f64().iter())
3007        {
3008            assert!((gpu - cpu).abs() < 1e-4, "gpu={gpu} cpu={cpu}");
3009        }
3010
3011        let telemetry = provider.telemetry_snapshot();
3012        assert_eq!(telemetry.linsolve.count, 1);
3013        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
3014        assert_eq!(kernel_launch_count(&telemetry, "linsolve_tall_qr"), 1);
3015    }
3016
3017    #[cfg(feature = "wgpu")]
3018    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
3019    #[test]
3020    fn wgpu_transposed_rectangular_linsolve_avoids_host_reupload_fallback() {
3021        let _accel_guard = test_support::accel_test_lock();
3022        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
3023            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
3024        ) else {
3025            return;
3026        };
3027        if provider.precision() != runmat_accelerate_api::ProviderPrecision::F32 {
3028            return;
3029        }
3030        let a = Tensor::new(vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0], vec![2, 3]).unwrap();
3031        let b = Tensor::new(vec![1.0, 2.0, 2.0], vec![3, 1]).unwrap();
3032
3033        let mut cpu_opts = StructValue::new();
3034        cpu_opts.fields.insert(
3035            "TRANSA".to_string(),
3036            Value::CharArray(CharArray::new_row("T")),
3037        );
3038        cpu_opts
3039            .fields
3040            .insert("RECT".to_string(), Value::Bool(true));
3041        let cpu = linsolve_builtin(
3042            Value::Tensor(a.clone()),
3043            Value::Tensor(b.clone()),
3044            vec![Value::Struct(cpu_opts)],
3045        )
3046        .expect("cpu linsolve");
3047        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
3048        provider.reset_telemetry();
3049
3050        let ha = provider
3051            .upload(&HostTensorView {
3052                data: &a.materialize_f64(),
3053                shape: &a.shape,
3054            })
3055            .expect("upload A");
3056        let hb = provider
3057            .upload(&HostTensorView {
3058                data: &b.materialize_f64(),
3059                shape: &b.shape,
3060            })
3061            .expect("upload B");
3062
3063        let _output_guard = crate::output_count::push_output_count(Some(1));
3064        let mut gpu_opts = StructValue::new();
3065        gpu_opts.fields.insert(
3066            "TRANSA".to_string(),
3067            Value::CharArray(CharArray::new_row("T")),
3068        );
3069        gpu_opts
3070            .fields
3071            .insert("RECT".to_string(), Value::Bool(true));
3072        let gpu_value = linsolve_builtin(
3073            Value::GpuTensor(ha.clone()),
3074            Value::GpuTensor(hb.clone()),
3075            vec![Value::Struct(gpu_opts)],
3076        )
3077        .expect("gpu transposed rectangular linsolve");
3078        let gpu_solution = match gpu_value {
3079            Value::OutputList(mut outputs) => outputs.remove(0),
3080            other => other,
3081        };
3082        let gathered = test_support::gather(gpu_solution).expect("gather");
3083        let _ = provider.free(&ha);
3084        let _ = provider.free(&hb);
3085
3086        assert_eq!(gathered.shape, cpu_tensor.shape);
3087        for (gpu, cpu) in gathered
3088            .materialize_f64()
3089            .iter()
3090            .zip(cpu_tensor.materialize_f64().iter())
3091        {
3092            assert!((gpu - cpu).abs() < 1e-4, "gpu={gpu} cpu={cpu}");
3093        }
3094
3095        let telemetry = provider.telemetry_snapshot();
3096        assert_eq!(telemetry.linsolve.count, 1);
3097        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
3098    }
3099
3100    #[cfg(feature = "wgpu")]
3101    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
3102    #[test]
3103    fn wgpu_triangular_hint_avoids_host_reupload_fallback() {
3104        let _accel_guard = test_support::accel_test_lock();
3105        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
3106            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
3107        ) else {
3108            return;
3109        };
3110        let a = Tensor::new(
3111            vec![3.0, -1.0, 4.0, 0.0, 2.0, 1.0, 0.0, 0.0, 5.0],
3112            vec![3, 3],
3113        )
3114        .unwrap();
3115        let b = Tensor::new(vec![9.0, 1.0, 19.0], vec![3, 1]).unwrap();
3116
3117        let cpu = linsolve_builtin(Value::Tensor(a.clone()), Value::Tensor(b.clone()), {
3118            let mut opts = StructValue::new();
3119            opts.fields.insert("LT".to_string(), Value::Bool(true));
3120            vec![Value::Struct(opts)]
3121        })
3122        .expect("cpu linsolve");
3123        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
3124        provider.reset_telemetry();
3125
3126        let ha = provider
3127            .upload(&HostTensorView {
3128                data: &a.materialize_f64(),
3129                shape: &a.shape,
3130            })
3131            .expect("upload A");
3132        let hb = provider
3133            .upload(&HostTensorView {
3134                data: &b.materialize_f64(),
3135                shape: &b.shape,
3136            })
3137            .expect("upload B");
3138
3139        let _output_guard = crate::output_count::push_output_count(Some(1));
3140        let mut opts = StructValue::new();
3141        opts.fields.insert("LT".to_string(), Value::Bool(true));
3142        let gpu_value = linsolve_builtin(
3143            Value::GpuTensor(ha.clone()),
3144            Value::GpuTensor(hb.clone()),
3145            vec![Value::Struct(opts)],
3146        )
3147        .expect("gpu triangular linsolve");
3148        let gpu_solution = match gpu_value {
3149            Value::OutputList(mut outputs) => outputs.remove(0),
3150            other => other,
3151        };
3152        let gathered = test_support::gather(gpu_solution).expect("gather");
3153        let _ = provider.free(&ha);
3154        let _ = provider.free(&hb);
3155
3156        assert_eq!(gathered.shape, cpu_tensor.shape);
3157        for (gpu, cpu) in gathered
3158            .materialize_f64()
3159            .iter()
3160            .zip(cpu_tensor.materialize_f64().iter())
3161        {
3162            assert!((gpu - cpu).abs() < 1e-5);
3163        }
3164
3165        let telemetry = provider.telemetry_snapshot();
3166        assert_eq!(telemetry.linsolve.count, 1);
3167        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
3168    }
3169
3170    #[cfg(feature = "wgpu")]
3171    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
3172    #[test]
3173    fn wgpu_transposed_triangular_hint_avoids_host_reupload_fallback() {
3174        let _accel_guard = test_support::accel_test_lock();
3175        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
3176            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
3177        ) else {
3178            return;
3179        };
3180        let a = Tensor::new(
3181            vec![3.0, 1.0, 0.0, 0.0, 4.0, 2.0, 0.0, 0.0, 5.0],
3182            vec![3, 3],
3183        )
3184        .unwrap();
3185        let b = Tensor::new(vec![5.0, 14.0, 23.0], vec![3, 1]).unwrap();
3186
3187        let mut cpu_opts = StructValue::new();
3188        cpu_opts.fields.insert("LT".to_string(), Value::Bool(true));
3189        cpu_opts.fields.insert(
3190            "TRANSA".to_string(),
3191            Value::CharArray(CharArray::new_row("T")),
3192        );
3193        let cpu = linsolve_builtin(
3194            Value::Tensor(a.clone()),
3195            Value::Tensor(b.clone()),
3196            vec![Value::Struct(cpu_opts)],
3197        )
3198        .expect("cpu linsolve");
3199        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
3200        provider.reset_telemetry();
3201
3202        let ha = provider
3203            .upload(&HostTensorView {
3204                data: &a.materialize_f64(),
3205                shape: &a.shape,
3206            })
3207            .expect("upload A");
3208        let hb = provider
3209            .upload(&HostTensorView {
3210                data: &b.materialize_f64(),
3211                shape: &b.shape,
3212            })
3213            .expect("upload B");
3214
3215        let _output_guard = crate::output_count::push_output_count(Some(1));
3216        let mut gpu_opts = StructValue::new();
3217        gpu_opts.fields.insert("LT".to_string(), Value::Bool(true));
3218        gpu_opts.fields.insert(
3219            "TRANSA".to_string(),
3220            Value::CharArray(CharArray::new_row("T")),
3221        );
3222        let gpu_value = linsolve_builtin(
3223            Value::GpuTensor(ha.clone()),
3224            Value::GpuTensor(hb.clone()),
3225            vec![Value::Struct(gpu_opts)],
3226        )
3227        .expect("gpu transposed triangular linsolve");
3228        let gpu_solution = match gpu_value {
3229            Value::OutputList(mut outputs) => outputs.remove(0),
3230            other => other,
3231        };
3232        let gathered = test_support::gather(gpu_solution).expect("gather");
3233        let _ = provider.free(&ha);
3234        let _ = provider.free(&hb);
3235
3236        assert_eq!(gathered.shape, cpu_tensor.shape);
3237        for (gpu, cpu) in gathered
3238            .materialize_f64()
3239            .iter()
3240            .zip(cpu_tensor.materialize_f64().iter())
3241        {
3242            assert!((gpu - cpu).abs() < 1e-5);
3243        }
3244
3245        let telemetry = provider.telemetry_snapshot();
3246        assert_eq!(telemetry.linsolve.count, 1);
3247        assert_eq!(fallback_count(&telemetry, "linsolve:host_reupload"), 0);
3248    }
3249
3250    #[cfg(feature = "wgpu")]
3251    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
3252    #[test]
3253    fn wgpu_round_trip_matches_cpu() {
3254        let _accel_guard = test_support::accel_test_lock();
3255        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
3256            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
3257        ) else {
3258            return;
3259        };
3260        let tol = match provider.precision() {
3261            runmat_accelerate_api::ProviderPrecision::F64 => 1e-12,
3262            runmat_accelerate_api::ProviderPrecision::F32 => 1e-5,
3263        };
3264
3265        let a = Tensor::new(vec![3.0, 1.0, 2.0, 4.0], vec![2, 2]).unwrap();
3266        let b = Tensor::new(vec![7.0, 8.0], vec![2, 1]).unwrap();
3267
3268        let cpu = linsolve_builtin(
3269            Value::Tensor(a.clone()),
3270            Value::Tensor(b.clone()),
3271            Vec::new(),
3272        )
3273        .expect("cpu linsolve");
3274        let cpu_tensor = test_support::gather(cpu).expect("cpu gather");
3275
3276        let view_a = HostTensorView {
3277            data: &a.materialize_f64(),
3278            shape: &a.shape,
3279        };
3280        let view_b = HostTensorView {
3281            data: &b.materialize_f64(),
3282            shape: &b.shape,
3283        };
3284        let ha = provider.upload(&view_a).expect("upload A");
3285        let hb = provider.upload(&view_b).expect("upload B");
3286        let gpu_value = linsolve_builtin(
3287            Value::GpuTensor(ha.clone()),
3288            Value::GpuTensor(hb.clone()),
3289            Vec::new(),
3290        )
3291        .expect("gpu linsolve");
3292        let gathered = test_support::gather(gpu_value).expect("gather");
3293        let _ = provider.free(&ha);
3294        let _ = provider.free(&hb);
3295
3296        assert_eq!(gathered.shape, cpu_tensor.shape);
3297        for (gpu, cpu) in gathered
3298            .materialize_f64()
3299            .iter()
3300            .zip(cpu_tensor.materialize_f64().iter())
3301        {
3302            assert!((gpu - cpu).abs() < tol);
3303        }
3304    }
3305}