1use 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
331pub 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
366pub 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#[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 pub fn solution(&self) -> Value {
392 self.solution.clone()
393 }
394
395 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
455pub 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 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 }
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}