Skip to main content

runmat_runtime/builtins/math/reduction/
gradient.rs

1//! MATLAB-compatible `gradient` builtin with scalar and coordinate-vector spacing support.
2
3use 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        // Return early before the `push(1)` padding loop: that loop would give a
663        // shape like [1] or [1,1] whose product is 1 ≠ 0, violating Tensor's
664        // invariant. Use the normalised shape directly, falling back to [0,0] if
665        // matlab_gradient_shape returned an empty vec (untyped empty tensor).
666        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        // Same fix as gradient_real_tensor_host: avoid padding the shape with 1s
752        // before the early return, which would produce product ≠ 0 for empty data.
753        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}