1use runmat_accelerate_api::{GpuTensorHandle, GpuTensorStorage};
4use runmat_builtins::{
5 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
6 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
7 ComplexTensor, ResolveContext, Tensor, Type, Value,
8};
9use runmat_macros::runtime_builtin;
10
11use crate::builtins::common::gpu_helpers;
12use crate::builtins::common::random_args::complex_tensor_into_value;
13use crate::builtins::common::spec::{
14 BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
15 ProviderHook, ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
16};
17use crate::builtins::common::tensor;
18use crate::builtins::math::type_resolvers::numeric_unary_type;
19use crate::{build_runtime_error, BuiltinResult, RuntimeError};
20
21const NAME: &str = "gradient";
22
23fn gradient_type(args: &[Type], ctx: &ResolveContext) -> Type {
24 numeric_unary_type(args, ctx)
25}
26
27const GRADIENT_OUTPUT_G: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
28 name: "G",
29 ty: BuiltinParamType::NumericArray,
30 arity: BuiltinParamArity::Required,
31 default: None,
32 description: "Primary gradient component.",
33}];
34
35const GRADIENT_OUTPUT_GS: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
36 name: "Gi",
37 ty: BuiltinParamType::NumericArray,
38 arity: BuiltinParamArity::Variadic,
39 default: None,
40 description: "Gradient components ordered by MATLAB axis semantics.",
41}];
42
43const GRADIENT_INPUTS_F: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
44 name: "F",
45 ty: BuiltinParamType::Any,
46 arity: BuiltinParamArity::Required,
47 default: None,
48 description: "Input scalar or array.",
49}];
50
51const GRADIENT_INPUTS_F_H: [BuiltinParamDescriptor; 2] = [
52 BuiltinParamDescriptor {
53 name: "F",
54 ty: BuiltinParamType::Any,
55 arity: BuiltinParamArity::Required,
56 default: None,
57 description: "Input scalar or array.",
58 },
59 BuiltinParamDescriptor {
60 name: "h",
61 ty: BuiltinParamType::Any,
62 arity: BuiltinParamArity::Optional,
63 default: Some("1"),
64 description: "Scalar spacing shared across all output dimensions, or a coordinate vector for vector inputs.",
65 },
66];
67
68const GRADIENT_INPUTS_F_HS: [BuiltinParamDescriptor; 2] = [
69 BuiltinParamDescriptor {
70 name: "F",
71 ty: BuiltinParamType::Any,
72 arity: BuiltinParamArity::Required,
73 default: None,
74 description: "Input scalar or array.",
75 },
76 BuiltinParamDescriptor {
77 name: "h_i",
78 ty: BuiltinParamType::Any,
79 arity: BuiltinParamArity::Variadic,
80 default: None,
81 description:
82 "Per-dimension scalar or coordinate-vector spacings (one per gradient dimension).",
83 },
84];
85
86const GRADIENT_SIGNATURES: [BuiltinSignatureDescriptor; 4] = [
87 BuiltinSignatureDescriptor {
88 label: "G = gradient(F)",
89 inputs: &GRADIENT_INPUTS_F,
90 outputs: &GRADIENT_OUTPUT_G,
91 },
92 BuiltinSignatureDescriptor {
93 label: "G = gradient(F, h)",
94 inputs: &GRADIENT_INPUTS_F_H,
95 outputs: &GRADIENT_OUTPUT_G,
96 },
97 BuiltinSignatureDescriptor {
98 label: "[G1, G2, ...] = gradient(F)",
99 inputs: &GRADIENT_INPUTS_F,
100 outputs: &GRADIENT_OUTPUT_GS,
101 },
102 BuiltinSignatureDescriptor {
103 label: "[G1, G2, ...] = gradient(F, h1, h2, ...)",
104 inputs: &GRADIENT_INPUTS_F_HS,
105 outputs: &GRADIENT_OUTPUT_GS,
106 },
107];
108
109const GRADIENT_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
110 code: "RM.GRADIENT.INVALID_ARGUMENT",
111 identifier: Some("RunMat:gradient:InvalidArgument"),
112 when: "Output-count or spacing argument grammar is invalid.",
113 message: "gradient: invalid argument",
114};
115
116const GRADIENT_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
117 code: "RM.GRADIENT.INVALID_INPUT",
118 identifier: Some("RunMat:gradient:InvalidInput"),
119 when: "Input value cannot be converted to a supported gradient domain.",
120 message: "gradient: invalid input",
121};
122
123const GRADIENT_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
124 code: "RM.GRADIENT.INTERNAL",
125 identifier: Some("RunMat:gradient:Internal"),
126 when: "Gradient execution fails due to gather, conversion, allocation, or indexing operations.",
127 message: "gradient: internal failure",
128};
129
130const GRADIENT_ERRORS: [BuiltinErrorDescriptor; 3] = [
131 GRADIENT_ERROR_INVALID_ARGUMENT,
132 GRADIENT_ERROR_INVALID_INPUT,
133 GRADIENT_ERROR_INTERNAL,
134];
135
136pub const GRADIENT_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
137 signatures: &GRADIENT_SIGNATURES,
138 output_mode: BuiltinOutputMode::ByRequestedOutputCount,
139 completion_policy: BuiltinCompletionPolicy::Public,
140 errors: &GRADIENT_ERRORS,
141};
142
143fn gradient_descriptor_error_with_message(
144 message: impl Into<String>,
145 error: &'static BuiltinErrorDescriptor,
146) -> RuntimeError {
147 let mut builder = build_runtime_error(message).with_builtin(NAME);
148 if let Some(identifier) = error.identifier {
149 builder = builder.with_identifier(identifier);
150 }
151 builder.build()
152}
153
154fn gradient_descriptor_error_with_detail(
155 error: &'static BuiltinErrorDescriptor,
156 detail: impl AsRef<str>,
157) -> RuntimeError {
158 gradient_descriptor_error_with_message(format!("{}: {}", error.message, detail.as_ref()), error)
159}
160
161fn gradient_invalid_argument(detail: impl AsRef<str>) -> RuntimeError {
162 gradient_descriptor_error_with_detail(&GRADIENT_ERROR_INVALID_ARGUMENT, detail)
163}
164
165fn gradient_invalid_input(detail: impl AsRef<str>) -> RuntimeError {
166 gradient_descriptor_error_with_detail(&GRADIENT_ERROR_INVALID_INPUT, detail)
167}
168
169fn gradient_internal_error(detail: impl AsRef<str>) -> RuntimeError {
170 gradient_descriptor_error_with_detail(&GRADIENT_ERROR_INTERNAL, detail)
171}
172
173#[derive(Clone, Debug, PartialEq)]
174enum GradientSpacing {
175 Scalar(f64),
176 Coordinates(Vec<f64>),
177}
178
179impl GradientSpacing {
180 fn is_scalar(&self) -> bool {
181 matches!(self, Self::Scalar(_))
182 }
183}
184
185#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::math::reduction::gradient")]
186pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
187 name: "gradient",
188 op_kind: GpuOpKind::Custom("numerical-gradient"),
189 supported_precisions: &[ScalarType::F32, ScalarType::F64],
190 broadcast: BroadcastSemantics::Matlab,
191 provider_hooks: &[
192 ProviderHook::Custom("gradient_dim"),
193 ProviderHook::Custom("gradient_dim_with_coordinates"),
194 ],
195 constant_strategy: ConstantStrategy::InlineLiteral,
196 residency: ResidencyPolicy::NewHandle,
197 nan_mode: ReductionNaN::Include,
198 two_pass_threshold: None,
199 workgroup_size: None,
200 accepts_nan_mode: false,
201 notes:
202 "Providers may keep scalar-spacing gradients on device via `gradient_dim` and coordinate-vector spacing via `gradient_dim_with_coordinates`.",
203};
204
205#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::math::reduction::gradient")]
206pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
207 name: "gradient",
208 shape: ShapeRequirements::Any,
209 constant_strategy: ConstantStrategy::InlineLiteral,
210 elementwise: None,
211 reduction: None,
212 emits_nan: false,
213 notes: "Gradient preserves input shape and uses edge-aware finite differences, so providers expose it through a custom sink hook.",
214};
215
216#[runtime_builtin(
217 name = "gradient",
218 category = "math/reduction",
219 summary = "Compute numerical gradients.",
220 keywords = "gradient,numerical gradient,finite difference,vector field,gpu",
221 accel = "gradient",
222 type_resolver(gradient_type),
223 descriptor(crate::builtins::math::reduction::gradient::GRADIENT_DESCRIPTOR),
224 builtin_path = "crate::builtins::math::reduction::gradient"
225)]
226async fn gradient_builtin(value: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
227 let requested_outputs = crate::output_count::current_output_count().unwrap_or(1);
228 if requested_outputs == 0 {
229 return Ok(Value::OutputList(Vec::new()));
230 }
231
232 let available_outputs = gradient_output_dims(value_shape(&value), value_len(&value));
233 if requested_outputs > available_outputs.len() {
234 return Err(gradient_invalid_argument(format!(
235 "gradient: requested {requested_outputs} outputs, but input supports at most {}",
236 available_outputs.len()
237 )));
238 }
239
240 let dim_lengths =
241 gradient_dim_lengths(value_shape(&value), value_len(&value), &available_outputs);
242 let spacings = parse_spacings(&rest, &available_outputs, &dim_lengths).await?;
243 let outputs =
244 evaluate_gradient_outputs(value, &available_outputs[..requested_outputs], &spacings)
245 .await?;
246
247 if crate::output_count::current_output_count().is_some() {
248 return Ok(Value::OutputList(outputs));
249 }
250
251 Ok(outputs
252 .into_iter()
253 .next()
254 .expect("single-output gradient result"))
255}
256
257async fn evaluate_gradient_outputs(
258 value: Value,
259 requested_dims: &[usize],
260 all_spacings: &[GradientSpacing],
261) -> BuiltinResult<Vec<Value>> {
262 if let Value::GpuTensor(handle) = value {
263 return gradient_gpu_outputs(handle, requested_dims, all_spacings).await;
264 }
265
266 evaluate_host_gradient_outputs(value, requested_dims, all_spacings)
267}
268
269fn evaluate_host_gradient_outputs(
270 value: Value,
271 requested_dims: &[usize],
272 all_spacings: &[GradientSpacing],
273) -> BuiltinResult<Vec<Value>> {
274 match value {
275 Value::Tensor(tensor) => {
276 let mut outputs = Vec::with_capacity(requested_dims.len());
277 for &dim in requested_dims {
278 let spacing = spacing_for_dim(dim, requested_dims, all_spacings);
279 outputs.push(tensor::tensor_into_value(
280 gradient_real_tensor_host_with_spacing(tensor.clone(), dim, spacing)?,
281 ));
282 }
283 Ok(outputs)
284 }
285 Value::LogicalArray(logical) => {
286 let tensor = tensor::logical_to_tensor(&logical).map_err(gradient_invalid_input)?;
287 let mut outputs = Vec::with_capacity(requested_dims.len());
288 for &dim in requested_dims {
289 let spacing = spacing_for_dim(dim, requested_dims, all_spacings);
290 outputs.push(tensor::tensor_into_value(
291 gradient_real_tensor_host_with_spacing(tensor.clone(), dim, spacing)?,
292 ));
293 }
294 Ok(outputs)
295 }
296 Value::Num(_) | Value::Int(_) | Value::Bool(_) => {
297 let tensor =
298 tensor::value_into_tensor_for(NAME, value).map_err(gradient_invalid_input)?;
299 let mut outputs = Vec::with_capacity(requested_dims.len());
300 for &dim in requested_dims {
301 let spacing = spacing_for_dim(dim, requested_dims, all_spacings);
302 outputs.push(tensor::tensor_into_value(
303 gradient_real_tensor_host_with_spacing(tensor.clone(), dim, spacing)?,
304 ));
305 }
306 Ok(outputs)
307 }
308 Value::Complex(re, im) => {
309 let tensor = ComplexTensor {
310 data: vec![(re, im)],
311 shape: vec![1, 1],
312 rows: 1,
313 cols: 1,
314 };
315 let mut outputs = Vec::with_capacity(requested_dims.len());
316 for &dim in requested_dims {
317 let spacing = spacing_for_dim(dim, requested_dims, all_spacings);
318 outputs.push(complex_tensor_into_value(
319 gradient_complex_tensor_host_with_spacing(tensor.clone(), dim, spacing)?,
320 ));
321 }
322 Ok(outputs)
323 }
324 Value::ComplexTensor(tensor) => {
325 let mut outputs = Vec::with_capacity(requested_dims.len());
326 for &dim in requested_dims {
327 let spacing = spacing_for_dim(dim, requested_dims, all_spacings);
328 outputs.push(complex_tensor_into_value(
329 gradient_complex_tensor_host_with_spacing(tensor.clone(), dim, spacing)?,
330 ));
331 }
332 Ok(outputs)
333 }
334 other => Err(gradient_invalid_input(format!(
335 "gradient: unsupported input type {:?}; expected numeric or logical data",
336 other
337 ))),
338 }
339}
340
341async fn gradient_gpu_outputs(
342 handle: GpuTensorHandle,
343 requested_dims: &[usize],
344 all_spacings: &[GradientSpacing],
345) -> BuiltinResult<Vec<Value>> {
346 let complex_storage =
347 runmat_accelerate_api::handle_storage(&handle) == GpuTensorStorage::ComplexInterleaved;
348
349 if let Some(provider) =
350 runmat_accelerate_api::provider_for_handle(&handle).or_else(runmat_accelerate_api::provider)
351 {
352 let _guard = runmat_accelerate_api::ThreadProviderGuard::set(Some(provider));
353 let mut outputs = Vec::with_capacity(requested_dims.len());
354 for &dim in requested_dims {
355 let spacing = spacing_for_dim(dim, requested_dims, all_spacings);
356 let device_result = match spacing {
357 GradientSpacing::Scalar(spacing) => {
358 provider.gradient_dim(&handle, dim.saturating_sub(1), *spacing)
359 }
360 GradientSpacing::Coordinates(coordinates) => {
361 let shape = vec![coordinates.len(), 1];
362 let coord_handle =
363 match provider.upload(&runmat_accelerate_api::HostTensorView {
364 data: coordinates,
365 shape: &shape,
366 }) {
367 Ok(handle) => handle,
368 Err(_) => {
369 let gathered =
370 gpu_helpers::gather_value_async(&Value::GpuTensor(handle))
371 .await?;
372 return evaluate_host_gradient_outputs(
373 gathered,
374 requested_dims,
375 all_spacings,
376 );
377 }
378 };
379 let result = provider.gradient_dim_with_coordinates(
380 &handle,
381 dim.saturating_sub(1),
382 &coord_handle,
383 );
384 let _ = provider.free(&coord_handle);
385 result
386 }
387 };
388 match device_result {
389 Ok(device_result) => {
390 if complex_storage
391 || runmat_accelerate_api::handle_storage(&device_result)
392 == GpuTensorStorage::ComplexInterleaved
393 {
394 outputs.push(gpu_helpers::complex_gpu_value(device_result));
395 } else {
396 outputs.push(gpu_helpers::resident_gpu_value(device_result));
397 }
398 }
399 Err(_) => {
400 let gathered =
401 gpu_helpers::gather_value_async(&Value::GpuTensor(handle)).await?;
402 return evaluate_host_gradient_outputs(gathered, requested_dims, all_spacings);
403 }
404 }
405 }
406 return Ok(outputs);
407 }
408
409 let gathered = gpu_helpers::gather_value_async(&Value::GpuTensor(handle)).await?;
410 evaluate_host_gradient_outputs(gathered, requested_dims, all_spacings)
411}
412
413fn spacing_for_dim<'a>(
414 dim: usize,
415 available_dims: &[usize],
416 spacings: &'a [GradientSpacing],
417) -> &'a GradientSpacing {
418 let index = available_dims
419 .iter()
420 .position(|candidate| *candidate == dim)
421 .expect("spacing lookup requires matching dimension");
422 &spacings[index]
423}
424
425async fn parse_spacings(
426 args: &[Value],
427 available_dims: &[usize],
428 dim_lengths: &[usize],
429) -> BuiltinResult<Vec<GradientSpacing>> {
430 match args.len() {
431 0 => Ok(vec![GradientSpacing::Scalar(1.0); available_dims.len()]),
432 1 => {
433 let spacing = parse_spacing_argument(&args[0], dim_lengths[0]).await?;
434 if spacing.is_scalar() {
435 Ok(vec![spacing; available_dims.len()])
436 } else if available_dims.len() == 1 {
437 Ok(vec![spacing])
438 } else {
439 Err(gradient_invalid_argument(
440 "gradient: coordinate-vector spacing for arrays requires one spacing argument per gradient dimension",
441 ))
442 }
443 }
444 count if count == available_dims.len() => {
445 let mut spacings = Vec::with_capacity(args.len());
446 for (value, &dim_len) in args.iter().zip(dim_lengths.iter()) {
447 spacings.push(parse_spacing_argument(value, dim_len).await?);
448 }
449 Ok(spacings)
450 }
451 _ => Err(gradient_invalid_argument(format!(
452 "gradient: expected 0, 1, or {} scalar/coordinate-vector spacing arguments",
453 available_dims.len()
454 ))),
455 }
456}
457
458async fn parse_spacing_argument(value: &Value, dim_len: usize) -> BuiltinResult<GradientSpacing> {
459 if let Value::GpuTensor(_) = value {
460 let gathered = gpu_helpers::gather_value_async(value).await?;
461 return parse_host_spacing_argument(&gathered, dim_len);
462 }
463 parse_host_spacing_argument(value, dim_len)
464}
465
466fn parse_host_spacing_argument(value: &Value, dim_len: usize) -> BuiltinResult<GradientSpacing> {
467 let tensor =
468 tensor::value_into_tensor_for(NAME, value.clone()).map_err(gradient_invalid_argument)?;
469 if tensor.data.is_empty() {
470 return Err(gradient_invalid_argument(
471 "gradient: empty spacing arguments are not supported",
472 ));
473 }
474
475 if tensor.data.len() == 1 {
476 let spacing = tensor.data[0];
477 validate_scalar_spacing(spacing)?;
478 return Ok(GradientSpacing::Scalar(spacing));
479 }
480
481 validate_coordinate_spacing(&tensor.data, dim_len)?;
482 Ok(GradientSpacing::Coordinates(tensor.data))
483}
484
485fn validate_scalar_spacing(spacing: f64) -> BuiltinResult<()> {
486 if !spacing.is_finite() {
487 return Err(gradient_invalid_argument(
488 "gradient: spacing must be finite",
489 ));
490 }
491 if spacing == 0.0 {
492 return Err(gradient_invalid_argument(
493 "gradient: spacing must be nonzero",
494 ));
495 }
496 Ok(())
497}
498
499fn validate_coordinate_spacing(coords: &[f64], dim_len: usize) -> BuiltinResult<()> {
500 if coords.len() != dim_len {
501 return Err(gradient_invalid_argument(format!(
502 "gradient: coordinate-vector spacing length {} does not match dimension length {dim_len}",
503 coords.len()
504 )));
505 }
506
507 if coords.iter().any(|coord| !coord.is_finite()) {
508 return Err(gradient_invalid_argument(
509 "gradient: coordinate-vector spacing must be finite",
510 ));
511 }
512
513 if coords.len() <= 1 {
514 return Ok(());
515 }
516
517 if coords[1] == coords[0] {
518 return Err(gradient_invalid_argument(
519 "gradient: coordinate-vector spacing points must be distinct",
520 ));
521 }
522
523 for k in 1..coords.len() {
524 if coords[k] == coords[k - 1] {
525 return Err(gradient_invalid_argument(
526 "gradient: coordinate-vector spacing points must be distinct",
527 ));
528 }
529 }
530
531 for k in 1..coords.len() - 1 {
532 if coords[k + 1] == coords[k - 1] {
533 return Err(gradient_invalid_argument(
534 "gradient: coordinate-vector spacing cannot produce zero finite-difference denominator",
535 ));
536 }
537 }
538 Ok(())
539}
540
541fn value_shape(value: &Value) -> &[usize] {
542 match value {
543 Value::Tensor(tensor) => &tensor.shape,
544 Value::LogicalArray(logical) => &logical.shape,
545 Value::ComplexTensor(tensor) => &tensor.shape,
546 Value::GpuTensor(handle) => &handle.shape,
547 _ => &[],
548 }
549}
550
551fn value_len(value: &Value) -> usize {
552 match value {
553 Value::Tensor(tensor) => tensor.data.len(),
554 Value::LogicalArray(logical) => logical.data.len(),
555 Value::ComplexTensor(tensor) => tensor.data.len(),
556 Value::GpuTensor(handle) => product(&handle.shape),
557 _ => 1,
558 }
559}
560
561pub fn matlab_gradient_shape(shape: &[usize], len: usize) -> Vec<usize> {
562 if shape.is_empty() {
563 if len == 0 {
564 Vec::new()
565 } else {
566 vec![1, 1]
567 }
568 } else if shape.len() == 1 {
569 if shape[0] == 1 {
570 vec![1, 1]
571 } else {
572 vec![1, shape[0]]
573 }
574 } else {
575 shape.to_vec()
576 }
577}
578
579fn gradient_output_dims(shape: &[usize], len: usize) -> Vec<usize> {
580 let normalized_shape = matlab_gradient_shape(shape, len);
581 let mut ext_shape = if normalized_shape.is_empty() {
582 if len == 0 {
583 vec![0, 0]
584 } else {
585 vec![1, 1]
586 }
587 } else {
588 normalized_shape
589 };
590 if ext_shape.len() == 1 {
591 ext_shape.push(1);
592 }
593
594 if ext_shape.len() <= 2 {
595 let rows = ext_shape.first().copied().unwrap_or(1);
596 let cols = ext_shape.get(1).copied().unwrap_or(1);
597 if rows == 1 && cols == 1 {
598 vec![1]
599 } else if rows == 1 {
600 vec![2]
601 } else if cols == 1 {
602 vec![1]
603 } else {
604 vec![2, 1]
605 }
606 } else {
607 let mut dims = vec![2, 1];
608 for dim in 3..=ext_shape.len() {
609 dims.push(dim);
610 }
611 dims
612 }
613}
614
615fn gradient_dim_lengths(shape: &[usize], len: usize, dims: &[usize]) -> Vec<usize> {
616 let mut ext_shape = matlab_gradient_shape(shape, len);
617 if ext_shape.is_empty() {
618 ext_shape = if len == 0 { vec![0, 0] } else { vec![1, 1] };
619 }
620
621 let max_dim = dims.iter().copied().max().unwrap_or(1);
622 while ext_shape.len() < max_dim {
623 ext_shape.push(1);
624 }
625
626 dims.iter()
627 .map(|dim| ext_shape[dim.saturating_sub(1)])
628 .collect()
629}
630
631pub fn gradient_real_tensor_host(
632 tensor: Tensor,
633 dim: usize,
634 spacing: f64,
635) -> BuiltinResult<Tensor> {
636 let spacing = GradientSpacing::Scalar(spacing);
637 gradient_real_tensor_host_with_spacing(tensor, dim, &spacing)
638}
639
640#[allow(dead_code)]
641pub fn gradient_real_tensor_host_with_coordinates(
642 tensor: Tensor,
643 dim: usize,
644 coordinates: Vec<f64>,
645) -> BuiltinResult<Tensor> {
646 let spacing = GradientSpacing::Coordinates(coordinates);
647 gradient_real_tensor_host_with_spacing(tensor, dim, &spacing)
648}
649
650fn gradient_real_tensor_host_with_spacing(
651 tensor: Tensor,
652 dim: usize,
653 spacing: &GradientSpacing,
654) -> BuiltinResult<Tensor> {
655 let Tensor {
656 data, shape, dtype, ..
657 } = tensor;
658 let dim_index = dim.saturating_sub(1);
659 let mut shape = matlab_gradient_shape(&shape, data.len());
660
661 if data.is_empty() {
662 let empty_shape = if shape.is_empty() { vec![0, 0] } else { shape };
667 return Tensor::new_with_dtype(Vec::new(), empty_shape, dtype)
668 .map_err(|e| gradient_internal_error(format!("gradient: {e}")));
669 }
670
671 while shape.len() <= dim_index {
672 shape.push(1);
673 }
674
675 let mut ext_shape = shape.clone();
676 while ext_shape.len() <= dim_index {
677 ext_shape.push(1);
678 }
679 let len_dim = ext_shape[dim_index];
680 let stride_before = if dim_index == 0 {
681 1usize
682 } else {
683 product(&ext_shape[..dim_index]).max(1)
684 };
685 let stride_after = if dim_index + 1 >= ext_shape.len() {
686 1usize
687 } else {
688 product(&ext_shape[dim_index + 1..]).max(1)
689 };
690
691 let mut out = vec![0.0; data.len()];
692 if len_dim > 1 {
693 let block = stride_before
694 .checked_mul(len_dim)
695 .ok_or_else(|| gradient_internal_error("gradient: block size overflow"))?;
696 for after in 0..stride_after {
697 let base = after
698 .checked_mul(block)
699 .ok_or_else(|| gradient_internal_error("gradient: indexing overflow"))?;
700 for before in 0..stride_before {
701 for k in 0..len_dim {
702 let idx = base + before + k * stride_before;
703 out[idx] = if k == 0 {
704 (data[idx + stride_before] - data[idx])
705 / spacing_denominator(spacing, k, len_dim)
706 } else if k + 1 == len_dim {
707 (data[idx] - data[idx - stride_before])
708 / spacing_denominator(spacing, k, len_dim)
709 } else {
710 (data[idx + stride_before] - data[idx - stride_before])
711 / spacing_denominator(spacing, k, len_dim)
712 };
713 }
714 }
715 }
716 }
717
718 Tensor::new_with_dtype(out, shape, dtype)
719 .map_err(|e| gradient_internal_error(format!("gradient: {e}")))
720}
721
722pub fn gradient_complex_tensor_host(
723 tensor: ComplexTensor,
724 dim: usize,
725 spacing: f64,
726) -> BuiltinResult<ComplexTensor> {
727 let spacing = GradientSpacing::Scalar(spacing);
728 gradient_complex_tensor_host_with_spacing(tensor, dim, &spacing)
729}
730
731#[allow(dead_code)]
732pub fn gradient_complex_tensor_host_with_coordinates(
733 tensor: ComplexTensor,
734 dim: usize,
735 coordinates: Vec<f64>,
736) -> BuiltinResult<ComplexTensor> {
737 let spacing = GradientSpacing::Coordinates(coordinates);
738 gradient_complex_tensor_host_with_spacing(tensor, dim, &spacing)
739}
740
741fn gradient_complex_tensor_host_with_spacing(
742 tensor: ComplexTensor,
743 dim: usize,
744 spacing: &GradientSpacing,
745) -> BuiltinResult<ComplexTensor> {
746 let ComplexTensor { data, shape, .. } = tensor;
747 let dim_index = dim.saturating_sub(1);
748 let mut shape = matlab_gradient_shape(&shape, data.len());
749
750 if data.is_empty() {
751 let empty_shape = if shape.is_empty() { vec![0, 0] } else { shape };
754 return ComplexTensor::new(Vec::new(), empty_shape)
755 .map_err(|e| gradient_internal_error(format!("gradient: {e}")));
756 }
757
758 while shape.len() <= dim_index {
759 shape.push(1);
760 }
761
762 let mut ext_shape = shape.clone();
763 while ext_shape.len() <= dim_index {
764 ext_shape.push(1);
765 }
766 let len_dim = ext_shape[dim_index];
767 let stride_before = if dim_index == 0 {
768 1usize
769 } else {
770 product(&ext_shape[..dim_index]).max(1)
771 };
772 let stride_after = if dim_index + 1 >= ext_shape.len() {
773 1usize
774 } else {
775 product(&ext_shape[dim_index + 1..]).max(1)
776 };
777
778 let mut out = vec![(0.0, 0.0); data.len()];
779 if len_dim > 1 {
780 let block = stride_before
781 .checked_mul(len_dim)
782 .ok_or_else(|| gradient_internal_error("gradient: block size overflow"))?;
783 for after in 0..stride_after {
784 let base = after
785 .checked_mul(block)
786 .ok_or_else(|| gradient_internal_error("gradient: indexing overflow"))?;
787 for before in 0..stride_before {
788 for k in 0..len_dim {
789 let idx = base + before + k * stride_before;
790 out[idx] = if k == 0 {
791 scale_complex(
792 sub_complex(data[idx + stride_before], data[idx]),
793 1.0 / spacing_denominator(spacing, k, len_dim),
794 )
795 } else if k + 1 == len_dim {
796 scale_complex(
797 sub_complex(data[idx], data[idx - stride_before]),
798 1.0 / spacing_denominator(spacing, k, len_dim),
799 )
800 } else {
801 scale_complex(
802 sub_complex(data[idx + stride_before], data[idx - stride_before]),
803 1.0 / spacing_denominator(spacing, k, len_dim),
804 )
805 };
806 }
807 }
808 }
809 }
810
811 ComplexTensor::new(out, shape).map_err(|e| gradient_internal_error(format!("gradient: {e}")))
812}
813
814fn spacing_denominator(spacing: &GradientSpacing, k: usize, len_dim: usize) -> f64 {
815 match spacing {
816 GradientSpacing::Scalar(spacing) => {
817 if k == 0 || k + 1 == len_dim {
818 *spacing
819 } else {
820 2.0 * spacing
821 }
822 }
823 GradientSpacing::Coordinates(coords) => {
824 if k == 0 {
825 coords[1] - coords[0]
826 } else if k + 1 == len_dim {
827 coords[len_dim - 1] - coords[len_dim - 2]
828 } else {
829 coords[k + 1] - coords[k - 1]
830 }
831 }
832 }
833}
834
835fn sub_complex(lhs: (f64, f64), rhs: (f64, f64)) -> (f64, f64) {
836 (lhs.0 - rhs.0, lhs.1 - rhs.1)
837}
838
839fn scale_complex(value: (f64, f64), scale: f64) -> (f64, f64) {
840 (value.0 * scale, value.1 * scale)
841}
842
843fn product(dims: &[usize]) -> usize {
844 dims.iter()
845 .copied()
846 .fold(1usize, |acc, value| acc.saturating_mul(value))
847}
848
849#[cfg(test)]
850mod tests {
851 use super::*;
852 use crate::builtins::common::test_support;
853 use futures::executor::block_on;
854 #[cfg(feature = "wgpu")]
855 use runmat_accelerate_api::AccelProvider;
856 use runmat_accelerate_api::HostTensorView;
857 use runmat_builtins::{NumericDType, Tensor};
858
859 fn gradient_builtin(value: Value, rest: Vec<Value>) -> BuiltinResult<Value> {
860 block_on(super::gradient_builtin(value, rest))
861 }
862
863 #[test]
864 fn gradient_descriptor_signatures_cover_core_forms() {
865 let labels: Vec<&str> = GRADIENT_DESCRIPTOR
866 .signatures
867 .iter()
868 .map(|sig| sig.label)
869 .collect();
870 assert!(labels.contains(&"G = gradient(F)"));
871 assert!(labels.contains(&"G = gradient(F, h)"));
872 assert!(labels.contains(&"[G1, G2, ...] = gradient(F)"));
873 assert!(labels.contains(&"[G1, G2, ...] = gradient(F, h1, h2, ...)"));
874 }
875
876 #[test]
877 fn gradient_descriptor_errors_have_stable_codes() {
878 assert!(GRADIENT_DESCRIPTOR
879 .errors
880 .iter()
881 .any(|error| error.code == GRADIENT_ERROR_INVALID_ARGUMENT.code));
882 assert!(GRADIENT_DESCRIPTOR
883 .errors
884 .iter()
885 .any(|error| error.code == GRADIENT_ERROR_INVALID_INPUT.code));
886 assert!(GRADIENT_DESCRIPTOR
887 .errors
888 .iter()
889 .any(|error| error.code == GRADIENT_ERROR_INTERNAL.code));
890 }
891
892 #[test]
893 fn gradient_row_vector_returns_horizontal_derivative() {
894 let tensor = Tensor::new(vec![1.0, 4.0, 9.0], vec![1, 3]).unwrap();
895 let result = gradient_builtin(Value::Tensor(tensor), Vec::new()).expect("gradient");
896 assert_eq!(
897 result,
898 Value::Tensor(Tensor::new(vec![3.0, 4.0, 5.0], vec![1, 3]).unwrap())
899 );
900 }
901
902 #[test]
903 fn gradient_one_dimensional_tensor_is_treated_as_row_vector() {
904 let tensor = Tensor::new(vec![1.0, 4.0, 9.0], vec![3]).unwrap();
905 let result =
906 gradient_builtin(Value::Tensor(tensor), vec![Value::Num(2.0)]).expect("gradient");
907 match result {
908 Value::Tensor(out) => {
909 assert_eq!(out.shape, vec![1, 3]);
910 assert_eq!(out.data, vec![1.5, 2.0, 2.5]);
911 }
912 other => panic!("expected tensor, got {other:?}"),
913 }
914 }
915
916 #[test]
917 fn gradient_matrix_outputs_follow_matlab_order() {
918 let tensor = Tensor::new(vec![1.0, 3.0, 2.0, 4.0], vec![2, 2]).unwrap();
919 let _guard = crate::output_count::push_output_count(Some(2));
920 let result = gradient_builtin(Value::Tensor(tensor), Vec::new()).expect("gradient");
921 match result {
922 Value::OutputList(outputs) => {
923 let fx = test_support::gather(outputs[0].clone()).expect("fx");
924 let fy = test_support::gather(outputs[1].clone()).expect("fy");
925 assert_eq!(fx.data, vec![1.0, 1.0, 1.0, 1.0]);
926 assert_eq!(fy.data, vec![2.0, 2.0, 2.0, 2.0]);
927 }
928 other => panic!("expected output list, got {other:?}"),
929 }
930 }
931
932 #[test]
933 fn gradient_scalar_spacing_scales_output() {
934 let tensor = Tensor::new(vec![1.0, 4.0, 9.0], vec![1, 3]).unwrap();
935 let result =
936 gradient_builtin(Value::Tensor(tensor), vec![Value::Num(2.0)]).expect("gradient");
937 match result {
938 Value::Tensor(out) => assert_eq!(out.data, vec![1.5, 2.0, 2.5]),
939 other => panic!("expected tensor, got {other:?}"),
940 }
941 }
942
943 #[test]
944 fn gradient_preserves_single_precision_host_tensor() {
945 let tensor =
946 Tensor::new_with_dtype(vec![1.0, 4.0, 9.0], vec![1, 3], NumericDType::F32).unwrap();
947 let result = gradient_builtin(Value::Tensor(tensor), Vec::new()).expect("gradient");
948 match result {
949 Value::Tensor(out) => assert_eq!(out.dtype, NumericDType::F32),
950 other => panic!("expected tensor, got {other:?}"),
951 }
952 }
953
954 #[test]
955 fn gradient_complex_host_supported() {
956 let tensor =
957 ComplexTensor::new(vec![(1.0, 1.0), (4.0, 3.0), (9.0, 6.0)], vec![1, 3]).unwrap();
958 let result = gradient_builtin(Value::ComplexTensor(tensor), Vec::new()).expect("gradient");
959 match result {
960 Value::ComplexTensor(out) => {
961 assert_eq!(out.data, vec![(3.0, 2.0), (4.0, 2.5), (5.0, 3.0)]);
962 }
963 other => panic!("expected complex tensor, got {other:?}"),
964 }
965 }
966
967 #[test]
968 fn gradient_coordinate_vector_spacing_for_row_vector() {
969 let tensor = Tensor::new(vec![1.0, 4.0, 9.0], vec![1, 3]).unwrap();
970 let spacing = Tensor::new(vec![0.0, 1.0, 3.0], vec![1, 3]).unwrap();
971 let result = gradient_builtin(Value::Tensor(tensor), vec![Value::Tensor(spacing)])
972 .expect("gradient");
973 match result {
974 Value::Tensor(out) => {
975 assert_eq!(out.shape, vec![1, 3]);
976 assert_eq!(out.data, vec![3.0, 8.0 / 3.0, 2.5]);
977 }
978 other => panic!("expected tensor, got {other:?}"),
979 }
980 }
981
982 #[test]
983 fn gradient_mixed_scalar_and_coordinate_vector_spacing_for_matrix() {
984 let tensor = Tensor::new(vec![0.0, 20.0, 1.0, 21.0, 9.0, 29.0], vec![2, 3]).unwrap();
985 let x = Tensor::new(vec![0.0, 1.0, 3.0], vec![1, 3]).unwrap();
986 let _guard = crate::output_count::push_output_count(Some(2));
987 let result = gradient_builtin(
988 Value::Tensor(tensor),
989 vec![Value::Tensor(x), Value::Num(2.0)],
990 )
991 .expect("gradient");
992 match result {
993 Value::OutputList(outputs) => {
994 let fx = test_support::gather(outputs[0].clone()).expect("fx");
995 let fy = test_support::gather(outputs[1].clone()).expect("fy");
996 assert_eq!(fx.shape, vec![2, 3]);
997 assert_eq!(fx.data, vec![1.0, 1.0, 3.0, 3.0, 4.0, 4.0]);
998 assert_eq!(fy.shape, vec![2, 3]);
999 assert_eq!(fy.data, vec![10.0, 10.0, 10.0, 10.0, 10.0, 10.0]);
1000 }
1001 other => panic!("expected output list, got {other:?}"),
1002 }
1003 }
1004
1005 #[test]
1006 fn gradient_complex_coordinate_vector_spacing() {
1007 let tensor =
1008 ComplexTensor::new(vec![(1.0, 1.0), (4.0, 3.0), (9.0, 7.0)], vec![1, 3]).unwrap();
1009 let spacing = Tensor::new(vec![0.0, 1.0, 3.0], vec![1, 3]).unwrap();
1010 let result = gradient_builtin(Value::ComplexTensor(tensor), vec![Value::Tensor(spacing)])
1011 .expect("gradient");
1012 match result {
1013 Value::ComplexTensor(out) => {
1014 assert_eq!(out.shape, vec![1, 3]);
1015 assert_eq!(out.data, vec![(3.0, 2.0), (8.0 / 3.0, 2.0), (2.5, 2.0)]);
1016 }
1017 other => panic!("expected complex tensor, got {other:?}"),
1018 }
1019 }
1020
1021 #[test]
1022 fn gradient_rejects_coordinate_vector_length_mismatch() {
1023 let tensor = Tensor::new(vec![1.0, 4.0, 9.0], vec![1, 3]).unwrap();
1024 let spacing = Tensor::new(vec![0.0, 1.0], vec![1, 2]).unwrap();
1025 let err =
1026 gradient_builtin(Value::Tensor(tensor), vec![Value::Tensor(spacing)]).unwrap_err();
1027 assert_eq!(err.identifier(), GRADIENT_ERROR_INVALID_ARGUMENT.identifier);
1028 assert!(err.message().contains("length"));
1029 }
1030
1031 #[test]
1032 fn gradient_allows_nonmonotonic_coordinate_vector_spacing() {
1033 let tensor = Tensor::new(vec![1.0, 4.0, 9.0], vec![1, 3]).unwrap();
1034 let spacing = Tensor::new(vec![0.0, 1.0, 0.5], vec![1, 3]).unwrap();
1035 let result = gradient_builtin(Value::Tensor(tensor), vec![Value::Tensor(spacing)])
1036 .expect("gradient");
1037 match result {
1038 Value::Tensor(out) => assert_eq!(out.data, vec![3.0, 16.0, -10.0]),
1039 other => panic!("expected tensor, got {other:?}"),
1040 }
1041 }
1042
1043 #[test]
1044 fn gradient_rejects_zero_coordinate_denominator() {
1045 let tensor = Tensor::new(vec![1.0, 4.0, 9.0], vec![1, 3]).unwrap();
1046 let spacing = Tensor::new(vec![0.0, 1.0, 0.0], vec![1, 3]).unwrap();
1047 let err =
1048 gradient_builtin(Value::Tensor(tensor), vec![Value::Tensor(spacing)]).unwrap_err();
1049 assert_eq!(err.identifier(), GRADIENT_ERROR_INVALID_ARGUMENT.identifier);
1050 assert!(err.message().contains("denominator"));
1051 }
1052
1053 #[test]
1054 fn gradient_rejects_too_many_outputs() {
1055 let tensor = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
1056 let _guard = crate::output_count::push_output_count(Some(2));
1057 let err = gradient_builtin(Value::Tensor(tensor), Vec::new()).unwrap_err();
1058 assert_eq!(err.identifier(), GRADIENT_ERROR_INVALID_ARGUMENT.identifier);
1059 assert!(err.message().contains("requested 2 outputs"));
1060 }
1061
1062 #[test]
1063 #[cfg(feature = "wgpu")]
1064 fn gradient_gpu_scalar_spacing_matches_cpu_and_stays_resident() {
1065 let _guard = test_support::accel_test_lock();
1066 let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
1067 runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
1068 ) else {
1069 return;
1070 };
1071 let host =
1072 Tensor::new_with_dtype(vec![1.0, 4.0, 9.0], vec![1, 3], NumericDType::F32).unwrap();
1073 let view = HostTensorView {
1074 data: &host.data,
1075 shape: &host.shape,
1076 };
1077 let handle = provider.upload(&view).expect("upload");
1078 let result =
1079 gradient_builtin(Value::GpuTensor(handle), vec![Value::Num(2.0)]).expect("gradient");
1080 match result {
1081 Value::GpuTensor(out) => {
1082 let gathered = test_support::gather(Value::GpuTensor(out)).expect("gather");
1083 assert_eq!(gathered.data, vec![1.5, 2.0, 2.5]);
1084 assert_eq!(gathered.dtype, NumericDType::F32);
1085 }
1086 other => panic!("expected gpu tensor, got {other:?}"),
1087 }
1088 }
1089
1090 #[test]
1091 #[cfg(feature = "wgpu")]
1092 fn gradient_gpu_coordinate_spacing_matches_cpu_and_stays_resident() {
1093 let _guard = test_support::accel_test_lock();
1094 let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
1095 runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
1096 ) else {
1097 return;
1098 };
1099 let host =
1100 Tensor::new_with_dtype(vec![1.0, 4.0, 9.0], vec![1, 3], NumericDType::F32).unwrap();
1101 let view = HostTensorView {
1102 data: &host.data,
1103 shape: &host.shape,
1104 };
1105 let handle = provider.upload(&view).expect("upload");
1106 let spacing = Tensor::new(vec![0.0, 1.0, 3.0], vec![1, 3]).unwrap();
1107 let result = gradient_builtin(Value::GpuTensor(handle), vec![Value::Tensor(spacing)])
1108 .expect("gradient");
1109 match result {
1110 Value::GpuTensor(out) => {
1111 let gathered = test_support::gather(Value::GpuTensor(out)).expect("gather");
1112 assert_eq!(gathered.shape, vec![1, 3]);
1113 assert_eq!(gathered.dtype, NumericDType::F32);
1114 let expected = [3.0, 8.0 / 3.0, 2.5];
1115 for (idx, (actual, expected)) in gathered.data.iter().zip(expected).enumerate() {
1116 assert!(
1117 (*actual - expected).abs() < 1.0e-5,
1118 "gradient mismatch at {idx}: actual={actual} expected={expected}"
1119 );
1120 }
1121 }
1122 other => panic!("expected gpu tensor, got {other:?}"),
1123 }
1124 }
1125
1126 #[test]
1127 #[cfg(feature = "wgpu")]
1128 fn gradient_gpu_one_dimensional_shape_matches_matlab_row_vector_semantics() {
1129 let _guard = test_support::accel_test_lock();
1130 let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
1131 runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
1132 ) else {
1133 return;
1134 };
1135 let data = [1.0, 4.0, 9.0];
1136 let shape = [3usize];
1137 let view = HostTensorView {
1138 data: &data,
1139 shape: &shape,
1140 };
1141 let handle = provider.upload(&view).expect("upload");
1142 let result =
1143 gradient_builtin(Value::GpuTensor(handle), vec![Value::Num(2.0)]).expect("gradient");
1144 let gathered = test_support::gather(result).expect("gather");
1145 assert_eq!(gathered.shape, vec![1, 3]);
1146 assert_eq!(gathered.data, vec![1.5, 2.0, 2.5]);
1147 }
1148
1149 #[test]
1150 #[cfg(feature = "wgpu")]
1151 fn gradient_gpu_multi_output_uses_output_list() {
1152 let _guard = test_support::accel_test_lock();
1153 let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
1154 runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
1155 ) else {
1156 return;
1157 };
1158 let host = Tensor::new(vec![1.0, 3.0, 2.0, 4.0], vec![2, 2]).unwrap();
1159 let view = HostTensorView {
1160 data: &host.data,
1161 shape: &host.shape,
1162 };
1163 let handle = provider.upload(&view).expect("upload");
1164 let _out_guard = crate::output_count::push_output_count(Some(2));
1165 let result = gradient_builtin(Value::GpuTensor(handle), Vec::new()).expect("gradient");
1166 match result {
1167 Value::OutputList(outputs) => {
1168 assert!(matches!(outputs[0], Value::GpuTensor(_)));
1169 assert!(matches!(outputs[1], Value::GpuTensor(_)));
1170 }
1171 other => panic!("expected output list, got {other:?}"),
1172 }
1173 }
1174
1175 #[test]
1176 fn gradient_gpu_coordinate_vector_spacing_stays_resident() {
1177 test_support::with_test_provider(|provider| {
1178 let host = Tensor::new(vec![1.0, 4.0, 9.0], vec![1, 3]).unwrap();
1179 let view = HostTensorView {
1180 data: &host.data,
1181 shape: &host.shape,
1182 };
1183 let handle = provider.upload(&view).expect("upload");
1184 let spacing = Tensor::new(vec![0.0, 1.0, 3.0], vec![1, 3]).unwrap();
1185 let result = gradient_builtin(Value::GpuTensor(handle), vec![Value::Tensor(spacing)])
1186 .expect("gradient");
1187 match result {
1188 Value::GpuTensor(out_handle) => {
1189 let out = test_support::gather(Value::GpuTensor(out_handle)).expect("gather");
1190 assert_eq!(out.shape, vec![1, 3]);
1191 assert_eq!(out.data, vec![3.0, 8.0 / 3.0, 2.5]);
1192 }
1193 other => panic!("expected gpu tensor, got {other:?}"),
1194 }
1195 });
1196 }
1197
1198 #[test]
1199 fn gradient_gpu_mixed_scalar_and_coordinate_outputs_stay_resident() {
1200 test_support::with_test_provider(|provider| {
1201 let host = Tensor::new(vec![1.0, 3.0, 2.0, 4.0], vec![2, 2]).unwrap();
1202 let view = HostTensorView {
1203 data: &host.data,
1204 shape: &host.shape,
1205 };
1206 let handle = provider.upload(&view).expect("upload");
1207 let spacing = Tensor::new(vec![0.0, 2.0], vec![2, 1]).unwrap();
1208 let _out_guard = crate::output_count::push_output_count(Some(2));
1209 let result = gradient_builtin(
1210 Value::GpuTensor(handle),
1211 vec![Value::Tensor(spacing), Value::Num(2.0)],
1212 )
1213 .expect("gradient");
1214 match result {
1215 Value::OutputList(outputs) => {
1216 assert!(matches!(outputs[0], Value::GpuTensor(_)));
1217 assert!(matches!(outputs[1], Value::GpuTensor(_)));
1218 let first = test_support::gather(outputs[0].clone()).expect("gather first");
1219 let second = test_support::gather(outputs[1].clone()).expect("gather second");
1220 assert_eq!(first.shape, vec![2, 2]);
1221 assert_eq!(first.data, vec![0.5, 0.5, 0.5, 0.5]);
1222 assert_eq!(second.shape, vec![2, 2]);
1223 assert_eq!(second.data, vec![1.0, 1.0, 1.0, 1.0]);
1224 }
1225 other => panic!("expected output list, got {other:?}"),
1226 }
1227 });
1228 }
1229
1230 #[test]
1231 fn gradient_inprocess_complex_gpu_matches_cpu_and_stays_resident() {
1232 test_support::with_test_provider(|provider| {
1233 let host = ComplexTensor::new(
1234 vec![
1235 (1.0, 1.0),
1236 (2.0, -1.0),
1237 (4.0, 3.0),
1238 (6.0, 2.0),
1239 (9.0, 6.0),
1240 (12.0, 4.0),
1241 ],
1242 vec![2, 3],
1243 )
1244 .unwrap();
1245 let expected =
1246 gradient_complex_tensor_host(host.clone(), 2, 2.0).expect("cpu gradient");
1247 let handle = gpu_helpers::upload_complex_tensor(provider, &host).expect("upload");
1248 let result = gradient_builtin(Value::GpuTensor(handle), vec![Value::Num(2.0)])
1249 .expect("gradient");
1250 let Value::GpuTensor(out_handle) = result else {
1251 panic!("expected complex gpu tensor");
1252 };
1253 assert_eq!(
1254 runmat_accelerate_api::handle_storage(&out_handle),
1255 GpuTensorStorage::ComplexInterleaved
1256 );
1257 let gathered = block_on(
1258 crate::builtins::math::fft::common::gather_gpu_complex_tensor(&out_handle, NAME),
1259 )
1260 .expect("gather complex gradient");
1261 assert_eq!(gathered.shape, expected.shape);
1262 assert_eq!(gathered.data, expected.data);
1263 });
1264 }
1265
1266 #[test]
1267 #[cfg(feature = "wgpu")]
1268 fn gradient_gpu_complex_matches_cpu_and_stays_resident() {
1269 let _guard = test_support::accel_test_lock();
1270 let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
1271 runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
1272 ) else {
1273 return;
1274 };
1275 let host = ComplexTensor::new(
1276 vec![
1277 (1.0, 1.0),
1278 (2.0, -1.0),
1279 (4.0, 3.0),
1280 (6.0, 2.0),
1281 (9.0, 6.0),
1282 (12.0, 4.0),
1283 ],
1284 vec![2, 3],
1285 )
1286 .unwrap();
1287 let expected = gradient_complex_tensor_host(host.clone(), 2, 2.0).expect("cpu gradient");
1288 let handle = gpu_helpers::upload_complex_tensor(provider, &host).expect("upload");
1289 let result =
1290 gradient_builtin(Value::GpuTensor(handle), vec![Value::Num(2.0)]).expect("gradient");
1291 let Value::GpuTensor(out_handle) = result else {
1292 panic!("expected complex gpu tensor");
1293 };
1294 assert_eq!(
1295 runmat_accelerate_api::handle_storage(&out_handle),
1296 GpuTensorStorage::ComplexInterleaved
1297 );
1298 let gathered = block_on(
1299 crate::builtins::math::fft::common::gather_gpu_complex_tensor(&out_handle, NAME),
1300 )
1301 .expect("gather complex gradient");
1302 assert_eq!(gathered.shape, expected.shape);
1303 for (idx, (actual, expected)) in gathered.data.iter().zip(expected.data.iter()).enumerate()
1304 {
1305 assert!(
1306 (actual.0 - expected.0).abs() <= 1.0e-5,
1307 "real mismatch at {idx}: actual={} expected={}",
1308 actual.0,
1309 expected.0
1310 );
1311 assert!(
1312 (actual.1 - expected.1).abs() <= 1.0e-5,
1313 "imag mismatch at {idx}: actual={} expected={}",
1314 actual.1,
1315 expected.1
1316 );
1317 }
1318 }
1319}