Skip to main content

tenferro_ad/
semantic_transform.rs

1//! Whole-program automatic differentiation over semantic SSA programs.
2
3mod core_dynamic;
4mod core_indexing;
5mod core_reductions;
6mod core_structural;
7
8use std::collections::{HashMap, HashSet};
9
10use tenferro_ops::{dim_expr::DimExpr, ShapeExtent};
11use tenferro_runtime::program::{
12    CoreSemanticOp, FrozenProgram, ProgramBuildError, ProgramFinishError, ProgramImport,
13    ProgramInputSpec, ProgramQueryError, ProgramValue, ProgramValueMetadata, SemanticOpRef,
14    SemanticProgramBuilder,
15};
16use tenferro_runtime::{CompareDir, DType, DotGeneralConfig};
17
18use crate::semantic_extension::{AdValue, SemanticAdError, SemanticExtensionRuleSet};
19use core_dynamic::{dynamic_shape_vjp, linearize_dynamic_shape};
20use core_indexing::{indexing_vjp, linearize_indexing};
21use core_reductions::{linearize_nonlinear_reduction, nonlinear_reduction_vjp};
22use core_structural::{concatenate_vjp, linearize_concatenate, pad_vjp, slice_vjp};
23
24/// Semantic AD transform role used by typed diagnostics.
25#[derive(Clone, Copy, Debug, PartialEq, Eq)]
26pub enum SemanticTransformRole {
27    /// Forward-mode linearization.
28    Jvp,
29    /// Reverse-mode transposition.
30    Vjp,
31}
32
33/// Failures produced by whole-program semantic AD.
34#[derive(Debug, thiserror::Error)]
35pub enum SemanticAdTransformError {
36    /// An activity mask did not match the corresponding ordered value list.
37    #[error("semantic {role:?} {field} expects {expected} entries, got {actual}")]
38    ActivityArity {
39        /// Transform role.
40        role: SemanticTransformRole,
41        /// Invalid mask.
42        field: &'static str,
43        /// Required entry count.
44        expected: usize,
45        /// Supplied entry count.
46        actual: usize,
47    },
48    /// An active core operation has not yet been admitted to semantic AD.
49    #[error("semantic {role:?} does not support active core operation {op}")]
50    UnsupportedCore {
51        /// Transform role.
52        role: SemanticTransformRole,
53        /// Bounded operation diagnostic.
54        op: String,
55    },
56    /// A future semantic operation variant is unknown to this transform.
57    #[error("semantic {role:?} does not support this semantic operation variant")]
58    UnsupportedOperationVariant {
59        /// Transform role.
60        role: SemanticTransformRole,
61    },
62    /// Active derivative metadata is outside the admitted exact-shape subset.
63    #[error("semantic {role:?} does not support derivative metadata: {message}")]
64    UnsupportedMetadata {
65        /// Transform role.
66        role: SemanticTransformRole,
67        /// Bounded metadata diagnostic.
68        message: String,
69    },
70    /// Source-program metadata could not be queried.
71    #[error("semantic AD source-program query failed: {0}")]
72    Query(#[from] ProgramQueryError),
73    /// Destination semantic-program construction failed.
74    #[error("semantic AD program construction failed: {0}")]
75    Build(#[from] ProgramBuildError),
76    /// An extension-owned semantic AD rule failed.
77    #[error("semantic extension AD failed: {0}")]
78    Extension(#[from] SemanticAdError),
79    /// The transformed program could not be frozen atomically.
80    #[error("semantic AD program finalization failed: {0}")]
81    Finish(#[from] ProgramFinishError),
82    /// The semantic AD transform cache could not be accessed.
83    #[error("semantic AD transform cache failed: {0}")]
84    Cache(#[source] tenferro_runtime::Error),
85}
86
87/// One frozen derivative program plus ordered derivative input/output maps.
88///
89/// Original primal inputs retain their source order. Active derivative seed
90/// inputs are appended in source order. Program outputs contain only active
91/// derivative values, also in source order; `None` records an inactive value.
92#[derive(Clone, Debug)]
93pub struct SemanticAdProgram {
94    frozen: FrozenProgram,
95    derivative_input_indices: Box<[Option<usize>]>,
96    derivative_output_indices: Box<[Option<usize>]>,
97}
98
99struct ValueShapePlan {
100    shape: Vec<DimExpr>,
101    dynamic_axes: Vec<usize>,
102}
103
104impl SemanticAdProgram {
105    /// Borrow the frozen derivative program.
106    pub const fn frozen(&self) -> &FrozenProgram {
107        &self.frozen
108    }
109
110    /// Return transformed-program input indices for ordered derivative seeds.
111    pub fn derivative_input_indices(&self) -> &[Option<usize>] {
112        &self.derivative_input_indices
113    }
114
115    /// Return transformed-program output indices for ordered derivatives.
116    pub fn derivative_output_indices(&self) -> &[Option<usize>] {
117        &self.derivative_output_indices
118    }
119
120    /// Consume this result and return the frozen derivative program.
121    pub fn into_frozen(self) -> FrozenProgram {
122        self.frozen
123    }
124
125    pub(crate) fn with_input_prefix_bindings_from(
126        &self,
127        source: &FrozenProgram,
128    ) -> Result<Self, ProgramFinishError> {
129        Ok(Self {
130            frozen: self.frozen.with_input_prefix_bindings_from(source)?,
131            derivative_input_indices: self.derivative_input_indices.clone(),
132            derivative_output_indices: self.derivative_output_indices.clone(),
133        })
134    }
135}
136
137/// Build a forward-mode derivative program.
138///
139/// `active_inputs` follows source-program input order. Each active input gets
140/// one appended tangent seed. The result maps source outputs to compact
141/// derivative-program outputs.
142///
143/// # Errors
144///
145/// Returns [`SemanticAdTransformError::ActivityArity`] for a mask-length
146/// mismatch, [`SemanticAdTransformError::Extension`] for a rejected extension
147/// rule, or the corresponding `Query`, `Build`, `Finish`, or unsupported
148/// operation/metadata variant for failures encountered during transformation.
149pub fn semantic_jvp(
150    input: &FrozenProgram,
151    active_inputs: &[bool],
152    rules: &SemanticExtensionRuleSet,
153) -> Result<SemanticAdProgram, SemanticAdTransformError> {
154    validate_activity(
155        SemanticTransformRole::Jvp,
156        "active_inputs",
157        input.program.inputs().len(),
158        active_inputs.len(),
159    )?;
160    let mut builder = SemanticProgramBuilder::new();
161    let values = import_source(input, &mut builder)?;
162    let mut tangents = HashMap::new();
163    let mut derivative_input_indices = vec![None; input.program.inputs().len()];
164    let mut next_input = input.program.inputs().len();
165    for (index, source) in input.program.inputs().iter().copied().enumerate() {
166        if active_inputs[index] {
167            let imported_source = values[&source];
168            let tangent = builder.input(ProgramInputSpec::from_metadata(
169                builder.value_metadata(imported_source)?.clone(),
170            ))?;
171            derivative_input_indices[index] = Some(next_input);
172            next_input += 1;
173            tangents.insert(source, AdValue::Value(tangent));
174        } else {
175            tangents.insert(source, AdValue::Absent);
176        }
177    }
178
179    let live = source_output_liveness(input);
180    for operation in input.program.operations() {
181        let tangent_inputs: Vec<_> = operation
182            .inputs()
183            .iter()
184            .map(|value| tangents.get(value).copied().unwrap_or(AdValue::Absent))
185            .collect();
186        let active_outputs: Vec<_> = operation
187            .outputs()
188            .iter()
189            .map(|value| live.contains(value))
190            .collect();
191        let tangent_outputs = if tangent_inputs
192            .iter()
193            .all(|value| matches!(value, AdValue::Absent))
194        {
195            vec![AdValue::Absent; operation.outputs().len()].into_boxed_slice()
196        } else {
197            match operation.op() {
198                SemanticOpRef::Extension(_) => rules
199                    .linearize_operation(
200                        operation,
201                        &mapped_values(operation.inputs(), &values),
202                        &mapped_values(operation.outputs(), &values),
203                        &tangent_inputs,
204                        &active_outputs,
205                        &mut builder,
206                    )?
207                    .tangent_outputs()
208                    .into(),
209                SemanticOpRef::Core(op) => linearize_core(
210                    op,
211                    &mapped_values(operation.inputs(), &values),
212                    &tangent_inputs,
213                    &mut builder,
214                )?,
215                _ => {
216                    return Err(SemanticAdTransformError::UnsupportedOperationVariant {
217                        role: SemanticTransformRole::Jvp,
218                    });
219                }
220            }
221        };
222        for (source, tangent) in operation.outputs().iter().copied().zip(tangent_outputs) {
223            tangents.insert(source, tangent);
224        }
225    }
226
227    let outputs = input
228        .program
229        .outputs()
230        .iter()
231        .map(|value| tangents.get(value).copied().unwrap_or(AdValue::Absent))
232        .collect();
233    finish_derivative(builder, derivative_input_indices, outputs)
234}
235
236/// Build a reverse-mode derivative program using extension semantic rules.
237///
238/// `active_inputs` selects requested source-input cotangents and
239/// `active_outputs` selects source outputs that receive appended cotangent
240/// seeds. Both masks follow source order.
241///
242/// # Errors
243///
244/// Returns [`SemanticAdTransformError::ActivityArity`] for a mask-length
245/// mismatch, [`SemanticAdTransformError::Extension`] for a rejected extension
246/// rule, or the corresponding `Query`, `Build`, `Finish`, or unsupported
247/// operation/metadata variant for failures encountered during transformation.
248pub fn semantic_vjp(
249    input: &FrozenProgram,
250    active_inputs: &[bool],
251    active_outputs: &[bool],
252    rules: &SemanticExtensionRuleSet,
253) -> Result<SemanticAdProgram, SemanticAdTransformError> {
254    semantic_vjp_with_saved_outputs(input, active_inputs, active_outputs, rules, &[])
255        .map(|(program, _)| program)
256}
257
258// Execution-only specialization: saved primal values are coefficients of the
259// VJP, not new differentiation targets. Keep the unspecialized derivative
260// program for functional gradients so higher-order AD retains their producers.
261pub(crate) fn semantic_vjp_with_saved_outputs(
262    input: &FrozenProgram,
263    active_inputs: &[bool],
264    active_outputs: &[bool],
265    rules: &SemanticExtensionRuleSet,
266    saved_outputs: &[usize],
267) -> Result<(SemanticAdProgram, Vec<usize>), SemanticAdTransformError> {
268    validate_activity(
269        SemanticTransformRole::Vjp,
270        "active_inputs",
271        input.program.inputs().len(),
272        active_inputs.len(),
273    )?;
274    validate_activity(
275        SemanticTransformRole::Vjp,
276        "active_outputs",
277        input.program.outputs().len(),
278        active_outputs.len(),
279    )?;
280    let mut builder = SemanticProgramBuilder::new();
281    let mut values = import_source(input, &mut builder)?;
282    let forward_active = requested_input_reachability(input, active_inputs);
283    let mut cotangents = HashMap::new();
284    let mut derivative_input_indices = vec![None; input.program.outputs().len()];
285    let mut next_input = input.program.inputs().len();
286    for (index, source) in input.program.outputs().iter().copied().enumerate() {
287        if active_outputs[index] {
288            let imported_source = values[&source];
289            let cotangent = builder.input(ProgramInputSpec::from_metadata(
290                builder.value_metadata(imported_source)?.clone(),
291            ))?;
292            derivative_input_indices[index] = Some(next_input);
293            next_input += 1;
294            accumulate_cotangent(&mut builder, &mut cotangents, source, cotangent)?;
295        }
296    }
297
298    let mut saved_input_indices = Vec::with_capacity(saved_outputs.len());
299    for &output_index in saved_outputs {
300        // INVARIANT: the eager caller supplies indices of the residual roots it
301        // appended to this source program, never user-provided indices.
302        let source = input.program.outputs()[output_index];
303        let imported = values[&source];
304        let saved = builder.input(ProgramInputSpec::from_metadata(
305            builder.value_metadata(imported)?.clone(),
306        ))?;
307        saved_input_indices.push(next_input);
308        next_input += 1;
309        values.insert(source, saved);
310    }
311
312    let operations: Vec<_> = input.program.operations().collect();
313    for operation in operations.into_iter().rev() {
314        let cotangent_outputs: Vec<_> = operation
315            .outputs()
316            .iter()
317            .map(|value| {
318                cotangents
319                    .get(value)
320                    .copied()
321                    .map_or(AdValue::Absent, AdValue::Value)
322            })
323            .collect();
324        if cotangent_outputs
325            .iter()
326            .all(|value| matches!(value, AdValue::Absent))
327        {
328            continue;
329        }
330        let active_operation_inputs: Vec<_> = operation
331            .inputs()
332            .iter()
333            .map(|value| forward_active.contains(value))
334            .collect();
335        if active_operation_inputs.iter().all(|active| !active) {
336            continue;
337        }
338        let cotangent_inputs = match operation.op() {
339            SemanticOpRef::Extension(op) => {
340                if rules.lookup_primal_vjp(op.family_id()).is_some() {
341                    rules.primal_vjp_operation(
342                        operation,
343                        &mapped_values(operation.inputs(), &values),
344                        &mapped_values(operation.outputs(), &values),
345                        &cotangent_outputs,
346                        &active_operation_inputs,
347                        &mut builder,
348                    )?
349                } else {
350                    let inactive_tangents = vec![AdValue::Absent; operation.inputs().len()];
351                    let active_operation_outputs: Vec<_> = cotangent_outputs
352                        .iter()
353                        .map(|value| matches!(value, AdValue::Value(_)))
354                        .collect();
355                    let linearized = rules.linearize_operation(
356                        operation,
357                        &mapped_values(operation.inputs(), &values),
358                        &mapped_values(operation.outputs(), &values),
359                        &inactive_tangents,
360                        &active_operation_outputs,
361                        &mut builder,
362                    )?;
363                    rules.linear_transpose_operation(
364                        operation,
365                        &mapped_values(operation.inputs(), &values),
366                        &mapped_values(operation.outputs(), &values),
367                        &cotangent_outputs,
368                        &active_operation_inputs,
369                        linearized.residuals(),
370                        &mut builder,
371                    )?
372                }
373            }
374            SemanticOpRef::Core(op) => vjp_core(
375                op,
376                &mapped_values(operation.inputs(), &values),
377                &mapped_values(operation.outputs(), &values),
378                &cotangent_outputs,
379                &active_operation_inputs,
380                &mut builder,
381            )?,
382            _ => {
383                return Err(SemanticAdTransformError::UnsupportedOperationVariant {
384                    role: SemanticTransformRole::Vjp,
385                });
386            }
387        };
388        for (source, cotangent) in operation.inputs().iter().copied().zip(cotangent_inputs) {
389            if let AdValue::Value(cotangent) = cotangent {
390                accumulate_cotangent(&mut builder, &mut cotangents, source, cotangent)?;
391            }
392        }
393    }
394
395    let outputs = input
396        .program
397        .inputs()
398        .iter()
399        .enumerate()
400        .map(|(index, value)| {
401            if active_inputs[index] {
402                cotangents
403                    .get(value)
404                    .copied()
405                    .map_or(AdValue::Absent, AdValue::Value)
406            } else {
407                AdValue::Absent
408            }
409        })
410        .collect();
411    Ok((
412        finish_derivative(builder, derivative_input_indices, outputs)?,
413        saved_input_indices,
414    ))
415}
416
417fn import_source(
418    input: &FrozenProgram,
419    builder: &mut SemanticProgramBuilder,
420) -> Result<HashMap<ProgramValue, ProgramValue>, SemanticAdTransformError> {
421    let mut source_values = input.program.inputs().to_vec();
422    source_values.extend(
423        input
424            .program
425            .operations()
426            .flat_map(|operation| operation.outputs().iter().copied()),
427    );
428    let imported = builder.import(ProgramImport {
429        program: input.program.as_ref(),
430        bindings: &input.bindings,
431        roots: &source_values,
432    })?;
433    Ok(source_values
434        .into_iter()
435        .zip(imported.roots().iter().copied())
436        .collect())
437}
438
439fn mapped_values(
440    source: &[ProgramValue],
441    values: &HashMap<ProgramValue, ProgramValue>,
442) -> Vec<ProgramValue> {
443    source.iter().map(|value| values[value]).collect()
444}
445
446fn source_output_liveness(input: &FrozenProgram) -> HashSet<ProgramValue> {
447    let mut live: HashSet<_> = input.program.outputs().iter().copied().collect();
448    let operations: Vec<_> = input.program.operations().collect();
449    for operation in operations.into_iter().rev() {
450        if operation
451            .outputs()
452            .iter()
453            .any(|output| live.contains(output))
454        {
455            live.extend(operation.inputs().iter().copied());
456        }
457    }
458    live
459}
460
461fn requested_input_reachability(
462    input: &FrozenProgram,
463    active_inputs: &[bool],
464) -> HashSet<ProgramValue> {
465    let mut active: HashSet<_> = input
466        .program
467        .inputs()
468        .iter()
469        .copied()
470        .zip(active_inputs.iter().copied())
471        .filter_map(|(value, is_active)| is_active.then_some(value))
472        .collect();
473    for operation in input.program.operations() {
474        if operation
475            .inputs()
476            .iter()
477            .any(|value| active.contains(value))
478        {
479            active.extend(operation.outputs().iter().copied());
480        }
481    }
482    active
483}
484
485fn accumulate_cotangent(
486    builder: &mut SemanticProgramBuilder,
487    cotangents: &mut HashMap<ProgramValue, ProgramValue>,
488    source: ProgramValue,
489    cotangent: ProgramValue,
490) -> Result<(), ProgramBuildError> {
491    let combined = if let Some(existing) = cotangents.get(&source).copied() {
492        builder.add_op(CoreSemanticOp::Add, &[existing, cotangent])?[0]
493    } else {
494        cotangent
495    };
496    cotangents.insert(source, combined);
497    Ok(())
498}
499
500fn linearize_core(
501    op: &CoreSemanticOp,
502    primal_inputs: &[ProgramValue],
503    tangent_inputs: &[AdValue],
504    builder: &mut SemanticProgramBuilder,
505) -> Result<Box<[AdValue]>, SemanticAdTransformError> {
506    let output = match op {
507        CoreSemanticOp::Add => add_ad_values(builder, tangent_inputs[0], tangent_inputs[1])?,
508        CoreSemanticOp::Sub => sub_ad_values(builder, tangent_inputs[0], tangent_inputs[1])?,
509        CoreSemanticOp::Mul => {
510            let lhs = multiply_ad_value(builder, tangent_inputs[0], primal_inputs[1])?;
511            let rhs = multiply_ad_value(builder, tangent_inputs[1], primal_inputs[0])?;
512            add_ad_values(builder, lhs, rhs)?
513        }
514        CoreSemanticOp::Div => {
515            let lhs = divide_ad_value(builder, tangent_inputs[0], primal_inputs[1])?;
516            let rhs_numerator = multiply_ad_value(builder, tangent_inputs[1], primal_inputs[0])?;
517            let denominator =
518                builder.add_op(CoreSemanticOp::Mul, &[primal_inputs[1], primal_inputs[1]])?[0];
519            let rhs = divide_ad_value(builder, rhs_numerator, denominator)?;
520            sub_ad_values(builder, lhs, rhs)?
521        }
522        CoreSemanticOp::Pow => {
523            let lhs = if matches!(tangent_inputs[0], AdValue::Value(_)) {
524                let one = one_like(builder, primal_inputs[1], SemanticTransformRole::Jvp)?;
525                let exponent_minus_one =
526                    builder.add_op(CoreSemanticOp::Sub, &[primal_inputs[1], one])?[0];
527                let power = builder
528                    .add_op(CoreSemanticOp::Pow, &[primal_inputs[0], exponent_minus_one])?[0];
529                let coefficient =
530                    builder.add_op(CoreSemanticOp::Mul, &[primal_inputs[1], power])?[0];
531                multiply_ad_value(builder, tangent_inputs[0], coefficient)?
532            } else {
533                AdValue::Absent
534            };
535            let rhs = if matches!(tangent_inputs[1], AdValue::Value(_)) {
536                let log = builder.add_op(CoreSemanticOp::Log, &[primal_inputs[0]])?[0];
537                let power =
538                    builder.add_op(CoreSemanticOp::Pow, &[primal_inputs[0], primal_inputs[1]])?[0];
539                let coefficient = builder.add_op(CoreSemanticOp::Mul, &[log, power])?[0];
540                multiply_ad_value(builder, tangent_inputs[1], coefficient)?
541            } else {
542                AdValue::Absent
543            };
544            add_ad_values(builder, lhs, rhs)?
545        }
546        CoreSemanticOp::DotGeneral { config } => {
547            linearize_dot_general(builder, primal_inputs, tangent_inputs, config)?
548        }
549        CoreSemanticOp::Abs => {
550            let input_dtype = builder.value_metadata(primal_inputs[0])?.dtype();
551            let sign = builder.add_op(CoreSemanticOp::Sign, &[primal_inputs[0]])?[0];
552            let coefficient = if is_complex_dtype(input_dtype) {
553                builder.add_op(CoreSemanticOp::Conj, &[sign])?[0]
554            } else {
555                sign
556            };
557            let tangent = multiply_ad_value(builder, tangent_inputs[0], coefficient)?;
558            convert_ad_value(builder, tangent, input_dtype, abs_output_dtype(input_dtype))?
559        }
560        CoreSemanticOp::Sign => linearize_sign(builder, primal_inputs[0], tangent_inputs[0])?,
561        CoreSemanticOp::Maximum | CoreSemanticOp::Minimum => {
562            linearize_extrema(builder, op, primal_inputs, tangent_inputs)?
563        }
564        CoreSemanticOp::Select => select_ad_values(
565            builder,
566            primal_inputs[0],
567            tangent_inputs[1],
568            tangent_inputs[2],
569        )?,
570        CoreSemanticOp::Clamp => linearize_clamp(builder, primal_inputs, tangent_inputs)?,
571        CoreSemanticOp::Neg | CoreSemanticOp::Conj => {
572            unary_ad_value(builder, op.clone(), tangent_inputs[0])?
573        }
574        CoreSemanticOp::Exp
575        | CoreSemanticOp::Log
576        | CoreSemanticOp::Sin
577        | CoreSemanticOp::Cos
578        | CoreSemanticOp::Tanh
579        | CoreSemanticOp::Sqrt
580        | CoreSemanticOp::Rsqrt
581        | CoreSemanticOp::Expm1
582        | CoreSemanticOp::Log1p
583        | CoreSemanticOp::Erf => {
584            linearize_analytic_unary(builder, op, primal_inputs[0], tangent_inputs[0])?
585        }
586        CoreSemanticOp::Transpose { .. }
587        | CoreSemanticOp::Reshape { .. }
588        | CoreSemanticOp::BroadcastInDim { .. }
589        | CoreSemanticOp::ReduceSum { .. }
590        | CoreSemanticOp::ExtractDiag { .. }
591        | CoreSemanticOp::EmbedDiag { .. }
592        | CoreSemanticOp::Tril { .. }
593        | CoreSemanticOp::Triu { .. }
594        | CoreSemanticOp::Slice(_)
595        | CoreSemanticOp::Pad(_)
596        | CoreSemanticOp::Reverse { .. } => {
597            linearize_unary_core(builder, op.clone(), primal_inputs, tangent_inputs[0])?
598        }
599        CoreSemanticOp::ReduceSumSquares { axes } => core_reductions::linearize_sum_squares(
600            builder,
601            primal_inputs[0],
602            tangent_inputs[0],
603            axes,
604        )?,
605        CoreSemanticOp::Concatenate { axis, input_count } => {
606            linearize_concatenate(builder, primal_inputs, tangent_inputs, *axis, *input_count)?
607        }
608        CoreSemanticOp::Gather(_)
609        | CoreSemanticOp::GatherDynamicSliceSizes { .. }
610        | CoreSemanticOp::Scatter(_)
611        | CoreSemanticOp::DynamicSlice { .. }
612        | CoreSemanticOp::DynamicUpdateSlice => {
613            linearize_indexing(builder, op, primal_inputs, tangent_inputs)?
614        }
615        CoreSemanticOp::DynamicTruncate { .. } | CoreSemanticOp::PadToMatch { .. } => {
616            linearize_dynamic_shape(builder, op, primal_inputs, tangent_inputs[0])?
617        }
618        CoreSemanticOp::Convert { from, to } => {
619            if is_differentiable_dtype(*from) && is_differentiable_dtype(*to) {
620                linearize_unary_core(builder, op.clone(), primal_inputs, tangent_inputs[0])?
621            } else {
622                AdValue::Absent
623            }
624        }
625        CoreSemanticOp::ReduceProd { .. }
626        | CoreSemanticOp::ReduceMax { .. }
627        | CoreSemanticOp::ReduceMin { .. } => {
628            linearize_nonlinear_reduction(builder, op, primal_inputs, tangent_inputs[0])?
629        }
630        CoreSemanticOp::Rem
631        | CoreSemanticOp::Compare(_)
632        | CoreSemanticOp::ShapeOf { .. }
633        | CoreSemanticOp::Constant { .. } => AdValue::Absent,
634        _ => return Err(unsupported_core(SemanticTransformRole::Jvp, op)),
635    };
636    Ok([output].into())
637}
638
639// These direct semantic VJPs consume outputs, unlike the lower-level linear
640// transpose adapters, whose residual contract exposes only their inputs.
641pub(crate) fn eager_core_residual_spec(
642    op: &tenferro_ops::std_tensor_op::StdTensorOp,
643) -> crate::semantic_extension::ResidualSpec {
644    use tenferro_ops::std_tensor_op::StdTensorOp;
645    match op {
646        StdTensorOp::Exp | StdTensorOp::Tanh => crate::semantic_extension::ResidualSpec::output(0),
647        _ => tenferro_ops::ad::primitive_residual_spec(op).unwrap_or_default(),
648    }
649}
650
651fn vjp_core(
652    op: &CoreSemanticOp,
653    primal_inputs: &[ProgramValue],
654    primal_outputs: &[ProgramValue],
655    cotangent_outputs: &[AdValue],
656    active_inputs: &[bool],
657    builder: &mut SemanticProgramBuilder,
658) -> Result<Box<[AdValue]>, SemanticAdTransformError> {
659    let cotangent = cotangent_outputs[0];
660    let inputs = match op {
661        CoreSemanticOp::Add => vec![
662            active_cotangent(builder, cotangent, active_inputs[0], primal_inputs[0])?,
663            active_cotangent(builder, cotangent, active_inputs[1], primal_inputs[1])?,
664        ],
665        CoreSemanticOp::Sub => {
666            let negated = unary_ad_value(builder, CoreSemanticOp::Neg, cotangent)?;
667            vec![
668                active_cotangent(builder, cotangent, active_inputs[0], primal_inputs[0])?,
669                normalize_ad_value(builder, negated, active_inputs[1], primal_inputs[1])?,
670            ]
671        }
672        CoreSemanticOp::Mul => {
673            let rhs_coefficient = conjugate_if_complex(builder, primal_inputs[1])?;
674            let lhs_coefficient = conjugate_if_complex(builder, primal_inputs[0])?;
675            let lhs = multiply_ad_value(builder, cotangent, rhs_coefficient)?;
676            let rhs = multiply_ad_value(builder, cotangent, lhs_coefficient)?;
677            vec![
678                normalize_ad_value(builder, lhs, active_inputs[0], primal_inputs[0])?,
679                normalize_ad_value(builder, rhs, active_inputs[1], primal_inputs[1])?,
680            ]
681        }
682        CoreSemanticOp::Div => {
683            let rhs_coefficient = conjugate_if_complex(builder, primal_inputs[1])?;
684            let lhs = divide_ad_value(builder, cotangent, rhs_coefficient)?;
685            let lhs_coefficient = conjugate_if_complex(builder, primal_inputs[0])?;
686            let denominator =
687                builder.add_op(CoreSemanticOp::Mul, &[rhs_coefficient, rhs_coefficient])?[0];
688            let rhs = multiply_ad_value(builder, cotangent, lhs_coefficient)?;
689            let rhs = divide_ad_value(builder, rhs, denominator)?;
690            let rhs = unary_ad_value(builder, CoreSemanticOp::Neg, rhs)?;
691            vec![
692                normalize_ad_value(builder, lhs, active_inputs[0], primal_inputs[0])?,
693                normalize_ad_value(builder, rhs, active_inputs[1], primal_inputs[1])?,
694            ]
695        }
696        CoreSemanticOp::Pow => {
697            let lhs = if active_inputs[0] {
698                let one = one_like(builder, primal_inputs[1], SemanticTransformRole::Vjp)?;
699                let exponent_minus_one =
700                    builder.add_op(CoreSemanticOp::Sub, &[primal_inputs[1], one])?[0];
701                let power = builder
702                    .add_op(CoreSemanticOp::Pow, &[primal_inputs[0], exponent_minus_one])?[0];
703                let coefficient =
704                    builder.add_op(CoreSemanticOp::Mul, &[primal_inputs[1], power])?[0];
705                let coefficient = conjugate_if_complex(builder, coefficient)?;
706                multiply_ad_value(builder, cotangent, coefficient)?
707            } else {
708                AdValue::Absent
709            };
710            let rhs = if active_inputs[1] {
711                let log = builder.add_op(CoreSemanticOp::Log, &[primal_inputs[0]])?[0];
712                let power =
713                    builder.add_op(CoreSemanticOp::Pow, &[primal_inputs[0], primal_inputs[1]])?[0];
714                let coefficient = builder.add_op(CoreSemanticOp::Mul, &[log, power])?[0];
715                let coefficient = conjugate_if_complex(builder, coefficient)?;
716                multiply_ad_value(builder, cotangent, coefficient)?
717            } else {
718                AdValue::Absent
719            };
720            vec![
721                normalize_ad_value(builder, lhs, active_inputs[0], primal_inputs[0])?,
722                normalize_ad_value(builder, rhs, active_inputs[1], primal_inputs[1])?,
723            ]
724        }
725        CoreSemanticOp::DotGeneral { config } => {
726            dot_general_vjp(builder, primal_inputs, cotangent, active_inputs, config)?
727        }
728        CoreSemanticOp::Abs => {
729            let input_dtype = builder.value_metadata(primal_inputs[0])?.dtype();
730            let output_dtype = abs_output_dtype(input_dtype);
731            let cotangent = convert_ad_value(builder, cotangent, output_dtype, input_dtype)?;
732            let sign = builder.add_op(CoreSemanticOp::Sign, &[primal_inputs[0]])?[0];
733            let cotangent = multiply_ad_value(builder, cotangent, sign)?;
734            vec![normalize_ad_value(
735                builder,
736                cotangent,
737                active_inputs[0],
738                primal_inputs[0],
739            )?]
740        }
741        CoreSemanticOp::Sign => vec![AdValue::Absent],
742        CoreSemanticOp::Maximum | CoreSemanticOp::Minimum => {
743            extrema_vjp(builder, op, primal_inputs, cotangent, active_inputs)?
744        }
745        CoreSemanticOp::Select => {
746            let (on_true, on_false) = split_select_cotangent(
747                builder,
748                primal_inputs[0],
749                cotangent,
750                active_inputs[1],
751                active_inputs[2],
752            )?;
753            vec![
754                AdValue::Absent,
755                normalize_ad_value(builder, on_true, active_inputs[1], primal_inputs[1])?,
756                normalize_ad_value(builder, on_false, active_inputs[2], primal_inputs[2])?,
757            ]
758        }
759        CoreSemanticOp::Clamp => clamp_vjp(builder, primal_inputs, cotangent, active_inputs)?,
760        CoreSemanticOp::Neg => {
761            let negated = unary_ad_value(builder, CoreSemanticOp::Neg, cotangent)?;
762            vec![normalize_ad_value(
763                builder,
764                negated,
765                active_inputs[0],
766                primal_inputs[0],
767            )?]
768        }
769        CoreSemanticOp::Conj => {
770            let conjugated = unary_ad_value(builder, CoreSemanticOp::Conj, cotangent)?;
771            vec![normalize_ad_value(
772                builder,
773                conjugated,
774                active_inputs[0],
775                primal_inputs[0],
776            )?]
777        }
778        CoreSemanticOp::Exp
779        | CoreSemanticOp::Log
780        | CoreSemanticOp::Sin
781        | CoreSemanticOp::Cos
782        | CoreSemanticOp::Tanh
783        | CoreSemanticOp::Sqrt
784        | CoreSemanticOp::Rsqrt
785        | CoreSemanticOp::Expm1
786        | CoreSemanticOp::Log1p
787        | CoreSemanticOp::Erf => {
788            let coefficient = match op {
789                CoreSemanticOp::Exp => primal_outputs[0],
790                CoreSemanticOp::Tanh => {
791                    let y = primal_outputs[0];
792                    let square = builder.add_op(CoreSemanticOp::Mul, &[y, y])?[0];
793                    let one = one_like(builder, y, SemanticTransformRole::Vjp)?;
794                    builder.add_op(CoreSemanticOp::Sub, &[one, square])?[0]
795                }
796                _ => analytic_unary_coefficient(
797                    builder,
798                    op,
799                    primal_inputs[0],
800                    SemanticTransformRole::Vjp,
801                )?,
802            };
803            let coefficient = conjugate_if_complex(builder, coefficient)?;
804            let cotangent = multiply_ad_value(builder, cotangent, coefficient)?;
805            vec![normalize_ad_value(
806                builder,
807                cotangent,
808                active_inputs[0],
809                primal_inputs[0],
810            )?]
811        }
812        CoreSemanticOp::Transpose { perm } => {
813            let transposed = unary_ad_value(
814                builder,
815                CoreSemanticOp::Transpose {
816                    perm: inverse_permutation(perm),
817                },
818                cotangent,
819            )?;
820            primary_cotangent(builder, transposed, active_inputs, primal_inputs, false)?
821        }
822        CoreSemanticOp::Reshape { .. } => {
823            let reshaped = reshape_ad_value_to_input(builder, cotangent, primal_inputs[0])?;
824            primary_cotangent(builder, reshaped, active_inputs, primal_inputs, false)?
825        }
826        CoreSemanticOp::BroadcastInDim { dims, .. } => {
827            let reduced =
828                transpose_broadcast(builder, cotangent, primal_inputs[0], dims.as_slice())?;
829            primary_cotangent(builder, reduced, active_inputs, primal_inputs, false)?
830        }
831        CoreSemanticOp::Convert { from, to } => {
832            let converted = if is_differentiable_dtype(*from) && is_differentiable_dtype(*to) {
833                unary_ad_value(
834                    builder,
835                    CoreSemanticOp::Convert {
836                        from: *to,
837                        to: *from,
838                    },
839                    cotangent,
840                )?
841            } else {
842                AdValue::Absent
843            };
844            primary_cotangent(builder, converted, active_inputs, primal_inputs, false)?
845        }
846        CoreSemanticOp::ReduceSum { axes } => {
847            let input_shape = value_shape_plan(
848                builder,
849                primal_inputs[0],
850                SemanticTransformRole::Vjp,
851                "reduce-sum input",
852            )?;
853            let dims = (0..input_shape.shape.len())
854                .filter(|axis| !axes.contains(axis))
855                .collect();
856            let broadcast = broadcast_ad_value_in_dim_to_shape(
857                builder,
858                cotangent,
859                primal_inputs[0],
860                &input_shape,
861                dims,
862            )?;
863            let broadcast = truncate_ad_value_to_dynamic_axes(
864                builder,
865                broadcast,
866                primal_inputs[0],
867                &input_shape.dynamic_axes,
868            )?;
869            primary_cotangent(builder, broadcast, active_inputs, primal_inputs, false)?
870        }
871        CoreSemanticOp::ReduceSumSquares { axes } => core_reductions::sum_squares_vjp(
872            builder,
873            primal_inputs[0],
874            cotangent,
875            active_inputs[0],
876            axes,
877        )?,
878        CoreSemanticOp::ExtractDiag { axis_a, axis_b } => {
879            let embedded = unary_ad_value(
880                builder,
881                CoreSemanticOp::EmbedDiag {
882                    axis_a: if axis_a < axis_b { *axis_a } else { axis_a - 1 },
883                    axis_b: *axis_b,
884                },
885                cotangent,
886            )?;
887            let padded = match embedded {
888                AdValue::Absent => AdValue::Absent,
889                AdValue::Value(value) => {
890                    let value = builder.add_op(
891                        CoreSemanticOp::PadToMatch { axis: *axis_a },
892                        &[value, primal_inputs[0]],
893                    )?[0];
894                    AdValue::Value(
895                        builder.add_op(
896                            CoreSemanticOp::PadToMatch { axis: *axis_b },
897                            &[value, primal_inputs[0]],
898                        )?[0],
899                    )
900                }
901            };
902            primary_cotangent(builder, padded, active_inputs, primal_inputs, false)?
903        }
904        CoreSemanticOp::EmbedDiag { axis_a, axis_b } => {
905            let source_axis = if axis_b <= axis_a {
906                axis_a + 1
907            } else {
908                *axis_a
909            };
910            let extracted = unary_ad_value(
911                builder,
912                CoreSemanticOp::ExtractDiag {
913                    axis_a: source_axis,
914                    axis_b: *axis_b,
915                },
916                cotangent,
917            )?;
918            let restored = if axis_b < axis_a {
919                let rank = builder.value_metadata(primal_inputs[0])?.shape().len();
920                let mut perm: Vec<_> = (0..rank).collect();
921                let diagonal_axis = perm.remove(*axis_b);
922                perm.insert(*axis_a, diagonal_axis);
923                unary_ad_value(builder, CoreSemanticOp::Transpose { perm }, extracted)?
924            } else {
925                extracted
926            };
927            primary_cotangent(builder, restored, active_inputs, primal_inputs, false)?
928        }
929        CoreSemanticOp::Tril { .. }
930        | CoreSemanticOp::Triu { .. }
931        | CoreSemanticOp::Reverse { .. } => {
932            let transformed = unary_ad_value(builder, op.clone(), cotangent)?;
933            primary_cotangent(builder, transformed, active_inputs, primal_inputs, false)?
934        }
935        CoreSemanticOp::Slice(config) => slice_vjp(
936            builder,
937            primal_inputs[0],
938            cotangent,
939            active_inputs[0],
940            config,
941        )?,
942        CoreSemanticOp::Pad(config) => pad_vjp(
943            builder,
944            primal_inputs[0],
945            cotangent,
946            active_inputs[0],
947            config,
948        )?,
949        CoreSemanticOp::Concatenate { axis, input_count } => concatenate_vjp(
950            builder,
951            primal_inputs,
952            cotangent,
953            active_inputs,
954            *axis,
955            *input_count,
956        )?,
957        CoreSemanticOp::Gather(_)
958        | CoreSemanticOp::GatherDynamicSliceSizes { .. }
959        | CoreSemanticOp::Scatter(_)
960        | CoreSemanticOp::DynamicSlice { .. }
961        | CoreSemanticOp::DynamicUpdateSlice => {
962            indexing_vjp(builder, op, primal_inputs, cotangent, active_inputs)?
963        }
964        CoreSemanticOp::DynamicTruncate { .. } | CoreSemanticOp::PadToMatch { .. } => {
965            dynamic_shape_vjp(builder, op, primal_inputs, cotangent, active_inputs)?
966        }
967        CoreSemanticOp::ReduceProd { .. }
968        | CoreSemanticOp::ReduceMax { .. }
969        | CoreSemanticOp::ReduceMin { .. } => {
970            nonlinear_reduction_vjp(builder, op, primal_inputs, cotangent, active_inputs[0])?
971        }
972        CoreSemanticOp::Rem | CoreSemanticOp::Compare(_) => {
973            vec![AdValue::Absent, AdValue::Absent]
974        }
975        CoreSemanticOp::ShapeOf { .. } => vec![AdValue::Absent],
976        CoreSemanticOp::Constant { .. } => Vec::new(),
977        _ => return Err(unsupported_core(SemanticTransformRole::Vjp, op)),
978    };
979    Ok(inputs.into_boxed_slice())
980}
981
982fn linearize_unary_core(
983    builder: &mut SemanticProgramBuilder,
984    op: CoreSemanticOp,
985    primal_inputs: &[ProgramValue],
986    tangent: AdValue,
987) -> Result<AdValue, ProgramBuildError> {
988    let AdValue::Value(tangent) = tangent else {
989        return Ok(AdValue::Absent);
990    };
991    let mut inputs = Vec::with_capacity(primal_inputs.len());
992    inputs.push(tangent);
993    inputs.extend_from_slice(&primal_inputs[1..]);
994    Ok(AdValue::Value(builder.add_op(op, &inputs)?[0]))
995}
996
997fn linearize_dot_general(
998    builder: &mut SemanticProgramBuilder,
999    primal_inputs: &[ProgramValue],
1000    tangent_inputs: &[AdValue],
1001    config: &DotGeneralConfig,
1002) -> Result<AdValue, SemanticAdTransformError> {
1003    validate_dot_general_metadata(builder, primal_inputs, config, SemanticTransformRole::Jvp)?;
1004    let mut terms = Vec::with_capacity(2);
1005    if let AdValue::Value(tangent) = tangent_inputs[0] {
1006        terms.push(
1007            builder.add_op(
1008                CoreSemanticOp::DotGeneral {
1009                    config: config.clone(),
1010                },
1011                &[tangent, primal_inputs[1]],
1012            )?[0],
1013        );
1014    }
1015    if let AdValue::Value(tangent) = tangent_inputs[1] {
1016        terms.push(
1017            builder.add_op(
1018                CoreSemanticOp::DotGeneral {
1019                    config: config.clone(),
1020                },
1021                &[primal_inputs[0], tangent],
1022            )?[0],
1023        );
1024    }
1025    let mut terms = terms.into_iter();
1026    let Some(mut result) = terms.next() else {
1027        return Ok(AdValue::Absent);
1028    };
1029    for term in terms {
1030        result = builder.add_op(CoreSemanticOp::Add, &[result, term])?[0];
1031    }
1032    Ok(AdValue::Value(result))
1033}
1034
1035fn dot_general_vjp(
1036    builder: &mut SemanticProgramBuilder,
1037    primal_inputs: &[ProgramValue],
1038    cotangent: AdValue,
1039    active_inputs: &[bool],
1040    config: &DotGeneralConfig,
1041) -> Result<Vec<AdValue>, SemanticAdTransformError> {
1042    let (lhs_rank, rhs_rank) =
1043        validate_dot_general_metadata(builder, primal_inputs, config, SemanticTransformRole::Vjp)?;
1044    let lhs_free = dot_general_free_dims(
1045        lhs_rank,
1046        &config.lhs_contracting_dims,
1047        &config.lhs_batch_dims,
1048        SemanticTransformRole::Vjp,
1049    )?;
1050    let rhs_free = dot_general_free_dims(
1051        rhs_rank,
1052        &config.rhs_contracting_dims,
1053        &config.rhs_batch_dims,
1054        SemanticTransformRole::Vjp,
1055    )?;
1056    let AdValue::Value(cotangent) = cotangent else {
1057        return Ok(vec![AdValue::Absent, AdValue::Absent]);
1058    };
1059    let mut result = vec![AdValue::Absent, AdValue::Absent];
1060
1061    if active_inputs[0] {
1062        let rhs = conjugate_if_complex(builder, primal_inputs[1])?;
1063        let (transpose_config, perm) =
1064            dot_general_transpose_plan_for_lhs(config, lhs_rank, rhs_rank, &lhs_free, &rhs_free)?;
1065        let value = builder.add_op(
1066            CoreSemanticOp::DotGeneral {
1067                config: transpose_config,
1068            },
1069            &[cotangent, rhs],
1070        )?[0];
1071        let value = transpose_if_needed(builder, value, &perm)?;
1072        result[0] = normalize_ad_value(builder, AdValue::Value(value), true, primal_inputs[0])?;
1073    }
1074    if active_inputs[1] {
1075        let lhs = conjugate_if_complex(builder, primal_inputs[0])?;
1076        let (transpose_config, perm) =
1077            dot_general_transpose_plan_for_rhs(config, lhs_rank, rhs_rank, &lhs_free, &rhs_free)?;
1078        let value = builder.add_op(
1079            CoreSemanticOp::DotGeneral {
1080                config: transpose_config,
1081            },
1082            &[lhs, cotangent],
1083        )?[0];
1084        let value = transpose_if_needed(builder, value, &perm)?;
1085        result[1] = normalize_ad_value(builder, AdValue::Value(value), true, primal_inputs[1])?;
1086    }
1087    Ok(result)
1088}
1089
1090fn validate_dot_general_metadata(
1091    builder: &SemanticProgramBuilder,
1092    primal_inputs: &[ProgramValue],
1093    config: &DotGeneralConfig,
1094    role: SemanticTransformRole,
1095) -> Result<(usize, usize), SemanticAdTransformError> {
1096    let lhs_rank = builder.value_metadata(primal_inputs[0])?.shape().len();
1097    let rhs_rank = builder.value_metadata(primal_inputs[1])?.shape().len();
1098    config
1099        .validate_dims_with_ranks(lhs_rank, rhs_rank)
1100        .map_err(|error| SemanticAdTransformError::UnsupportedMetadata {
1101            role,
1102            message: format!(
1103                "invalid dot_general dimensions for ranks {lhs_rank} and {rhs_rank}: {error}"
1104            ),
1105        })?;
1106    Ok((lhs_rank, rhs_rank))
1107}
1108
1109fn dot_general_free_dims(
1110    rank: usize,
1111    contracting: &[usize],
1112    batch: &[usize],
1113    role: SemanticTransformRole,
1114) -> Result<Vec<usize>, SemanticAdTransformError> {
1115    let mut bound = vec![false; rank];
1116    for &axis in batch.iter().chain(contracting) {
1117        let Some(slot) = bound.get_mut(axis) else {
1118            return Err(SemanticAdTransformError::UnsupportedMetadata {
1119                role,
1120                message: format!("dot_general axis {axis} is out of bounds for rank {rank}"),
1121            });
1122        };
1123        *slot = true;
1124    }
1125    Ok((0..rank).filter(|axis| !bound[*axis]).collect())
1126}
1127
1128fn dot_general_transpose_plan_for_lhs(
1129    config: &DotGeneralConfig,
1130    lhs_rank: usize,
1131    rhs_rank: usize,
1132    lhs_free: &[usize],
1133    rhs_free: &[usize],
1134) -> Result<(DotGeneralConfig, Vec<usize>), SemanticAdTransformError> {
1135    let batch_count = config.lhs_batch_dims.len();
1136    let output_rank = lhs_free.len() + rhs_free.len() + batch_count;
1137    let rhs_free_positions = (lhs_free.len()..lhs_free.len() + rhs_free.len()).collect();
1138    let rhs_contracting_order = dot_general_free_dims(
1139        rhs_rank,
1140        rhs_free,
1141        &config.rhs_batch_dims,
1142        SemanticTransformRole::Vjp,
1143    )?;
1144    let mut result_order = lhs_free.to_vec();
1145    for rhs_axis in rhs_contracting_order {
1146        let Some(pair) = config
1147            .rhs_contracting_dims
1148            .iter()
1149            .position(|&axis| axis == rhs_axis)
1150        else {
1151            return Err(dot_general_transpose_metadata_error(format!(
1152                "rhs contracting axis {rhs_axis} has no lhs pair"
1153            )));
1154        };
1155        result_order.push(config.lhs_contracting_dims[pair]);
1156    }
1157    result_order.extend(config.lhs_batch_dims.iter().copied());
1158    Ok((
1159        DotGeneralConfig {
1160            lhs_contracting_dims: rhs_free_positions,
1161            rhs_contracting_dims: rhs_free.into(),
1162            lhs_batch_dims: (lhs_free.len() + rhs_free.len()..output_rank).collect(),
1163            rhs_batch_dims: config.rhs_batch_dims.clone(),
1164        },
1165        permutation_to_original_order(lhs_rank, &result_order)?,
1166    ))
1167}
1168
1169fn dot_general_transpose_plan_for_rhs(
1170    config: &DotGeneralConfig,
1171    lhs_rank: usize,
1172    rhs_rank: usize,
1173    lhs_free: &[usize],
1174    rhs_free: &[usize],
1175) -> Result<(DotGeneralConfig, Vec<usize>), SemanticAdTransformError> {
1176    let batch_count = config.lhs_batch_dims.len();
1177    let lhs_contracting_order = dot_general_free_dims(
1178        lhs_rank,
1179        lhs_free,
1180        &config.lhs_batch_dims,
1181        SemanticTransformRole::Vjp,
1182    )?;
1183    let mut result_order = Vec::with_capacity(rhs_rank);
1184    for lhs_axis in lhs_contracting_order {
1185        let Some(pair) = config
1186            .lhs_contracting_dims
1187            .iter()
1188            .position(|&axis| axis == lhs_axis)
1189        else {
1190            return Err(dot_general_transpose_metadata_error(format!(
1191                "lhs contracting axis {lhs_axis} has no rhs pair"
1192            )));
1193        };
1194        result_order.push(config.rhs_contracting_dims[pair]);
1195    }
1196    result_order.extend(rhs_free.iter().copied());
1197    result_order.extend(config.rhs_batch_dims.iter().copied());
1198    let output_rank = lhs_free.len() + rhs_free.len() + batch_count;
1199    Ok((
1200        DotGeneralConfig {
1201            lhs_contracting_dims: lhs_free.into(),
1202            rhs_contracting_dims: (0..lhs_free.len()).collect(),
1203            lhs_batch_dims: config.lhs_batch_dims.clone(),
1204            rhs_batch_dims: (lhs_free.len() + rhs_free.len()..output_rank).collect(),
1205        },
1206        permutation_to_original_order(rhs_rank, &result_order)?,
1207    ))
1208}
1209
1210fn permutation_to_original_order(
1211    rank: usize,
1212    result_order: &[usize],
1213) -> Result<Vec<usize>, SemanticAdTransformError> {
1214    let mut permutation = vec![0; rank];
1215    for (result_axis, &original_axis) in result_order.iter().enumerate() {
1216        let Some(slot) = permutation.get_mut(original_axis) else {
1217            return Err(dot_general_transpose_metadata_error(format!(
1218                "dot_general transpose axis {original_axis} is out of bounds for rank {rank}"
1219            )));
1220        };
1221        *slot = result_axis;
1222    }
1223    Ok(permutation)
1224}
1225
1226fn transpose_if_needed(
1227    builder: &mut SemanticProgramBuilder,
1228    value: ProgramValue,
1229    permutation: &[usize],
1230) -> Result<ProgramValue, ProgramBuildError> {
1231    if permutation
1232        .iter()
1233        .enumerate()
1234        .all(|(axis, &mapped)| axis == mapped)
1235    {
1236        Ok(value)
1237    } else {
1238        Ok(builder.add_op(
1239            CoreSemanticOp::Transpose {
1240                perm: permutation.to_vec(),
1241            },
1242            &[value],
1243        )?[0])
1244    }
1245}
1246
1247fn dot_general_transpose_metadata_error(message: String) -> SemanticAdTransformError {
1248    SemanticAdTransformError::UnsupportedMetadata {
1249        role: SemanticTransformRole::Vjp,
1250        message,
1251    }
1252}
1253
1254fn primary_cotangent(
1255    builder: &mut SemanticProgramBuilder,
1256    cotangent: AdValue,
1257    active_inputs: &[bool],
1258    primal_inputs: &[ProgramValue],
1259    normalize: bool,
1260) -> Result<Vec<AdValue>, SemanticAdTransformError> {
1261    let mut result = vec![AdValue::Absent; primal_inputs.len()];
1262    if active_inputs.first().copied().unwrap_or(false) {
1263        result[0] = if normalize {
1264            normalize_ad_value(builder, cotangent, true, primal_inputs[0])?
1265        } else {
1266            cotangent
1267        };
1268    }
1269    Ok(result)
1270}
1271
1272fn inverse_permutation(perm: &[usize]) -> Vec<usize> {
1273    let mut inverse = vec![0; perm.len()];
1274    for (axis, mapped) in perm.iter().copied().enumerate() {
1275        inverse[mapped] = axis;
1276    }
1277    inverse
1278}
1279
1280fn reshape_ad_value_to_input(
1281    builder: &mut SemanticProgramBuilder,
1282    value: AdValue,
1283    primal_input: ProgramValue,
1284) -> Result<AdValue, SemanticAdTransformError> {
1285    let shape = value_shape_plan(
1286        builder,
1287        primal_input,
1288        SemanticTransformRole::Vjp,
1289        "reshape input",
1290    )?;
1291    let reshaped = reshape_ad_value_to_shape(builder, value, primal_input, &shape)?;
1292    truncate_ad_value_to_dynamic_axes(builder, reshaped, primal_input, &shape.dynamic_axes)
1293}
1294
1295fn transpose_broadcast(
1296    builder: &mut SemanticProgramBuilder,
1297    value: AdValue,
1298    primal_input: ProgramValue,
1299    dims: &[usize],
1300) -> Result<AdValue, SemanticAdTransformError> {
1301    let AdValue::Value(mut value) = value else {
1302        return Ok(AdValue::Absent);
1303    };
1304    let input_shape = value_shape_plan(
1305        builder,
1306        primal_input,
1307        SemanticTransformRole::Vjp,
1308        "broadcast input",
1309    )?;
1310    let output_shape = value_shape_plan(
1311        builder,
1312        value,
1313        SemanticTransformRole::Vjp,
1314        "broadcast cotangent",
1315    )?;
1316    let input_rank = input_shape.shape.len();
1317    let output_rank = output_shape.shape.len();
1318    let metadata_error = |message| SemanticAdTransformError::UnsupportedMetadata {
1319        role: SemanticTransformRole::Vjp,
1320        message,
1321    };
1322    if dims.len() != input_rank {
1323        return Err(metadata_error(format!(
1324            "broadcast dims length {} does not match input rank {input_rank}",
1325            dims.len()
1326        )));
1327    }
1328    let mut seen = HashSet::with_capacity(dims.len());
1329    for (input_axis, &output_axis) in dims.iter().enumerate() {
1330        if output_axis >= output_rank {
1331            return Err(metadata_error(format!(
1332                "broadcast dims[{input_axis}] = {output_axis} is out of bounds for output rank {output_rank}"
1333            )));
1334        }
1335        if !seen.insert(output_axis) {
1336            return Err(metadata_error(format!(
1337                "broadcast dims[{input_axis}] = {output_axis} duplicates an earlier output axis"
1338            )));
1339        }
1340    }
1341    // INVARIANT: these repeated membership and position scans are bounded by tensor rank;
1342    // replace them with axis maps only if unusually high-rank tensors make this measurable.
1343    let mut reduce_axes: Vec<_> = (0..output_rank)
1344        .filter(|axis| !dims.contains(axis))
1345        .collect();
1346    reduce_axes.extend(
1347        dims.iter()
1348            .copied()
1349            .enumerate()
1350            .filter_map(|(input_axis, output_axis)| {
1351                (matches!(
1352                    input_shape.shape[input_axis],
1353                    tenferro_ops::dim_expr::DimExpr::Const(1)
1354                ) && input_shape.shape[input_axis] != output_shape.shape[output_axis])
1355                    .then_some(output_axis)
1356            }),
1357    );
1358    reduce_axes.sort_unstable();
1359    reduce_axes.dedup();
1360    if !reduce_axes.is_empty() {
1361        value = builder.add_op(
1362            CoreSemanticOp::ReduceSum {
1363                axes: reduce_axes.clone(),
1364            },
1365            &[value],
1366        )?[0];
1367    }
1368
1369    let remaining_output_axes: Vec<_> = (0..output_rank)
1370        .filter(|axis| !reduce_axes.contains(axis))
1371        .collect();
1372    let perm: Vec<_> = dims
1373        .iter()
1374        .copied()
1375        .filter(|axis| !reduce_axes.contains(axis))
1376        .map(|axis| {
1377            remaining_output_axes
1378                .iter()
1379                .position(|candidate| *candidate == axis)
1380                .ok_or_else(|| {
1381                    metadata_error(format!(
1382                        "broadcast output axis {axis} did not survive cotangent reduction"
1383                    ))
1384                })
1385        })
1386        .collect::<Result<_, _>>()?;
1387    if perm.iter().copied().ne(0..perm.len()) {
1388        value = builder.add_op(CoreSemanticOp::Transpose { perm }, &[value])?[0];
1389    }
1390    if builder.value_metadata(value)?.shape() != builder.value_metadata(primal_input)?.shape() {
1391        value = reshape_value_to_shape(builder, value, primal_input, &input_shape)?;
1392    }
1393    value =
1394        truncate_value_to_dynamic_axes(builder, value, primal_input, &input_shape.dynamic_axes)?;
1395    Ok(AdValue::Value(value))
1396}
1397
1398fn exact_value_shape(
1399    builder: &SemanticProgramBuilder,
1400    value: ProgramValue,
1401    role: SemanticTransformRole,
1402    field: &'static str,
1403) -> Result<Vec<tenferro_ops::dim_expr::DimExpr>, SemanticAdTransformError> {
1404    exact_shape(builder.value_metadata(value)?.shape(), role, field)
1405}
1406
1407fn is_differentiable_dtype(dtype: DType) -> bool {
1408    matches!(dtype, DType::F32 | DType::F64 | DType::C32 | DType::C64)
1409}
1410
1411fn is_complex_dtype(dtype: DType) -> bool {
1412    matches!(dtype, DType::C32 | DType::C64)
1413}
1414
1415fn abs_output_dtype(dtype: DType) -> DType {
1416    match dtype {
1417        DType::C32 => DType::F32,
1418        DType::C64 => DType::F64,
1419        other => other,
1420    }
1421}
1422
1423fn add_ad_values(
1424    builder: &mut SemanticProgramBuilder,
1425    lhs: AdValue,
1426    rhs: AdValue,
1427) -> Result<AdValue, ProgramBuildError> {
1428    match (lhs, rhs) {
1429        (AdValue::Absent, value) | (value, AdValue::Absent) => Ok(value),
1430        (AdValue::Value(lhs), AdValue::Value(rhs)) => Ok(AdValue::Value(
1431            builder.add_op(CoreSemanticOp::Add, &[lhs, rhs])?[0],
1432        )),
1433    }
1434}
1435
1436fn sub_ad_values(
1437    builder: &mut SemanticProgramBuilder,
1438    lhs: AdValue,
1439    rhs: AdValue,
1440) -> Result<AdValue, ProgramBuildError> {
1441    match (lhs, rhs) {
1442        (AdValue::Absent, AdValue::Absent) => Ok(AdValue::Absent),
1443        (value, AdValue::Absent) => Ok(value),
1444        (AdValue::Absent, AdValue::Value(rhs)) => Ok(AdValue::Value(
1445            builder.add_op(CoreSemanticOp::Neg, &[rhs])?[0],
1446        )),
1447        (AdValue::Value(lhs), AdValue::Value(rhs)) => Ok(AdValue::Value(
1448            builder.add_op(CoreSemanticOp::Sub, &[lhs, rhs])?[0],
1449        )),
1450    }
1451}
1452
1453fn unary_ad_value(
1454    builder: &mut SemanticProgramBuilder,
1455    op: CoreSemanticOp,
1456    value: AdValue,
1457) -> Result<AdValue, ProgramBuildError> {
1458    match value {
1459        AdValue::Absent => Ok(AdValue::Absent),
1460        AdValue::Value(value) => Ok(AdValue::Value(builder.add_op(op, &[value])?[0])),
1461    }
1462}
1463
1464fn multiply_ad_value(
1465    builder: &mut SemanticProgramBuilder,
1466    value: AdValue,
1467    coefficient: ProgramValue,
1468) -> Result<AdValue, ProgramBuildError> {
1469    match value {
1470        AdValue::Absent => Ok(AdValue::Absent),
1471        AdValue::Value(value) => Ok(AdValue::Value(
1472            builder.add_op(CoreSemanticOp::Mul, &[value, coefficient])?[0],
1473        )),
1474    }
1475}
1476
1477fn divide_ad_value(
1478    builder: &mut SemanticProgramBuilder,
1479    value: AdValue,
1480    denominator: ProgramValue,
1481) -> Result<AdValue, ProgramBuildError> {
1482    match value {
1483        AdValue::Absent => Ok(AdValue::Absent),
1484        AdValue::Value(value) => Ok(AdValue::Value(
1485            builder.add_op(CoreSemanticOp::Div, &[value, denominator])?[0],
1486        )),
1487    }
1488}
1489
1490fn convert_ad_value(
1491    builder: &mut SemanticProgramBuilder,
1492    value: AdValue,
1493    from: DType,
1494    to: DType,
1495) -> Result<AdValue, ProgramBuildError> {
1496    if from == to {
1497        return Ok(value);
1498    }
1499    unary_ad_value(builder, CoreSemanticOp::Convert { from, to }, value)
1500}
1501
1502fn select_ad_values(
1503    builder: &mut SemanticProgramBuilder,
1504    condition: ProgramValue,
1505    on_true: AdValue,
1506    on_false: AdValue,
1507) -> Result<AdValue, SemanticAdTransformError> {
1508    match (on_true, on_false) {
1509        (AdValue::Absent, AdValue::Absent) => Ok(AdValue::Absent),
1510        (AdValue::Value(on_true), AdValue::Value(on_false)) => Ok(AdValue::Value(
1511            builder.add_op(CoreSemanticOp::Select, &[condition, on_true, on_false])?[0],
1512        )),
1513        (AdValue::Value(on_true), AdValue::Absent) => {
1514            let zero = zero_constant_like(builder, on_true, SemanticTransformRole::Jvp)?;
1515            Ok(AdValue::Value(
1516                builder.add_op(CoreSemanticOp::Select, &[condition, on_true, zero])?[0],
1517            ))
1518        }
1519        (AdValue::Absent, AdValue::Value(on_false)) => {
1520            let zero = zero_constant_like(builder, on_false, SemanticTransformRole::Jvp)?;
1521            Ok(AdValue::Value(
1522                builder.add_op(CoreSemanticOp::Select, &[condition, zero, on_false])?[0],
1523            ))
1524        }
1525    }
1526}
1527
1528fn split_select_cotangent(
1529    builder: &mut SemanticProgramBuilder,
1530    condition: ProgramValue,
1531    cotangent: AdValue,
1532    true_active: bool,
1533    false_active: bool,
1534) -> Result<(AdValue, AdValue), SemanticAdTransformError> {
1535    if !true_active && !false_active {
1536        return Ok((AdValue::Absent, AdValue::Absent));
1537    }
1538    let AdValue::Value(cotangent) = cotangent else {
1539        return Ok((AdValue::Absent, AdValue::Absent));
1540    };
1541    let zero = zero_constant_like(builder, cotangent, SemanticTransformRole::Vjp)?;
1542    let on_true = if true_active {
1543        AdValue::Value(builder.add_op(CoreSemanticOp::Select, &[condition, cotangent, zero])?[0])
1544    } else {
1545        AdValue::Absent
1546    };
1547    let on_false = if false_active {
1548        AdValue::Value(builder.add_op(CoreSemanticOp::Select, &[condition, zero, cotangent])?[0])
1549    } else {
1550        AdValue::Absent
1551    };
1552    Ok((on_true, on_false))
1553}
1554
1555fn linearize_extrema(
1556    builder: &mut SemanticProgramBuilder,
1557    op: &CoreSemanticOp,
1558    primal_inputs: &[ProgramValue],
1559    tangent_inputs: &[AdValue],
1560) -> Result<AdValue, SemanticAdTransformError> {
1561    let output = builder.add_op(op.clone(), primal_inputs)?[0];
1562    let lhs_eq_output = builder.add_op(
1563        CoreSemanticOp::Compare(CompareDir::Eq),
1564        &[primal_inputs[0], output],
1565    )?[0];
1566    let rhs_eq_output = builder.add_op(
1567        CoreSemanticOp::Compare(CompareDir::Eq),
1568        &[primal_inputs[1], output],
1569    )?[0];
1570    let lhs = balanced_extrema_contribution(
1571        builder,
1572        tangent_inputs[0],
1573        lhs_eq_output,
1574        rhs_eq_output,
1575        SemanticTransformRole::Jvp,
1576    )?;
1577    let rhs = balanced_extrema_contribution(
1578        builder,
1579        tangent_inputs[1],
1580        rhs_eq_output,
1581        lhs_eq_output,
1582        SemanticTransformRole::Jvp,
1583    )?;
1584    Ok(add_ad_values(builder, lhs, rhs)?)
1585}
1586
1587fn extrema_vjp(
1588    builder: &mut SemanticProgramBuilder,
1589    op: &CoreSemanticOp,
1590    primal_inputs: &[ProgramValue],
1591    cotangent: AdValue,
1592    active_inputs: &[bool],
1593) -> Result<Vec<AdValue>, SemanticAdTransformError> {
1594    let output = builder.add_op(op.clone(), primal_inputs)?[0];
1595    let lhs_eq_output = builder.add_op(
1596        CoreSemanticOp::Compare(CompareDir::Eq),
1597        &[primal_inputs[0], output],
1598    )?[0];
1599    let rhs_eq_output = builder.add_op(
1600        CoreSemanticOp::Compare(CompareDir::Eq),
1601        &[primal_inputs[1], output],
1602    )?[0];
1603    let lhs = balanced_extrema_contribution(
1604        builder,
1605        cotangent,
1606        lhs_eq_output,
1607        rhs_eq_output,
1608        SemanticTransformRole::Vjp,
1609    )?;
1610    let rhs = balanced_extrema_contribution(
1611        builder,
1612        cotangent,
1613        rhs_eq_output,
1614        lhs_eq_output,
1615        SemanticTransformRole::Vjp,
1616    )?;
1617    Ok(vec![
1618        normalize_ad_value(builder, lhs, active_inputs[0], primal_inputs[0])?,
1619        normalize_ad_value(builder, rhs, active_inputs[1], primal_inputs[1])?,
1620    ])
1621}
1622
1623fn balanced_extrema_contribution(
1624    builder: &mut SemanticProgramBuilder,
1625    active: AdValue,
1626    self_eq_output: ProgramValue,
1627    other_eq_output: ProgramValue,
1628    role: SemanticTransformRole,
1629) -> Result<AdValue, SemanticAdTransformError> {
1630    let AdValue::Value(active) = active else {
1631        return Ok(AdValue::Absent);
1632    };
1633    let zero = zero_constant_like(builder, active, role)?;
1634    let selected = builder.add_op(CoreSemanticOp::Select, &[self_eq_output, active, zero])?[0];
1635    let one = one_like(builder, active, role)?;
1636    let two = builder.add_op(CoreSemanticOp::Add, &[one, one])?[0];
1637    let half = builder.add_op(CoreSemanticOp::Div, &[selected, two])?[0];
1638    Ok(AdValue::Value(
1639        builder.add_op(CoreSemanticOp::Select, &[other_eq_output, half, selected])?[0],
1640    ))
1641}
1642
1643fn linearize_clamp(
1644    builder: &mut SemanticProgramBuilder,
1645    primal_inputs: &[ProgramValue],
1646    tangent_inputs: &[AdValue],
1647) -> Result<AdValue, SemanticAdTransformError> {
1648    let masks = clamp_masks(builder, primal_inputs)?;
1649    let input = mask_ad_value(
1650        builder,
1651        tangent_inputs[0],
1652        &[masks[0], masks[1]],
1653        SemanticTransformRole::Jvp,
1654    )?;
1655    let lower = mask_ad_value(
1656        builder,
1657        tangent_inputs[1],
1658        &[masks[2], masks[3]],
1659        SemanticTransformRole::Jvp,
1660    )?;
1661    let upper = mask_ad_value(
1662        builder,
1663        tangent_inputs[2],
1664        &[masks[4]],
1665        SemanticTransformRole::Jvp,
1666    )?;
1667    let input_and_lower = add_ad_values(builder, input, lower)?;
1668    Ok(add_ad_values(builder, input_and_lower, upper)?)
1669}
1670
1671fn clamp_vjp(
1672    builder: &mut SemanticProgramBuilder,
1673    primal_inputs: &[ProgramValue],
1674    cotangent: AdValue,
1675    active_inputs: &[bool],
1676) -> Result<Vec<AdValue>, SemanticAdTransformError> {
1677    let masks = clamp_masks(builder, primal_inputs)?;
1678    let input = mask_ad_value(
1679        builder,
1680        cotangent,
1681        &[masks[0], masks[1]],
1682        SemanticTransformRole::Vjp,
1683    )?;
1684    let lower = mask_ad_value(
1685        builder,
1686        cotangent,
1687        &[masks[2], masks[3]],
1688        SemanticTransformRole::Vjp,
1689    )?;
1690    let upper = mask_ad_value(builder, cotangent, &[masks[4]], SemanticTransformRole::Vjp)?;
1691    Ok(vec![
1692        normalize_ad_value(builder, input, active_inputs[0], primal_inputs[0])?,
1693        normalize_ad_value(builder, lower, active_inputs[1], primal_inputs[1])?,
1694        normalize_ad_value(builder, upper, active_inputs[2], primal_inputs[2])?,
1695    ])
1696}
1697
1698fn clamp_masks(
1699    builder: &mut SemanticProgramBuilder,
1700    primal_inputs: &[ProgramValue],
1701) -> Result<[ProgramValue; 5], ProgramBuildError> {
1702    let input = primal_inputs[0];
1703    let lower = primal_inputs[1];
1704    let upper = primal_inputs[2];
1705    let input_gt_lower =
1706        builder.add_op(CoreSemanticOp::Compare(CompareDir::Gt), &[input, lower])?[0];
1707    let input_lt_upper =
1708        builder.add_op(CoreSemanticOp::Compare(CompareDir::Lt), &[input, upper])?[0];
1709    let lower_gt_input =
1710        builder.add_op(CoreSemanticOp::Compare(CompareDir::Gt), &[lower, input])?[0];
1711    let lower_lt_upper =
1712        builder.add_op(CoreSemanticOp::Compare(CompareDir::Lt), &[lower, upper])?[0];
1713    let max_input_lower = builder.add_op(CoreSemanticOp::Maximum, &[input, lower])?[0];
1714    let upper_lt_max_input_lower = builder.add_op(
1715        CoreSemanticOp::Compare(CompareDir::Lt),
1716        &[upper, max_input_lower],
1717    )?[0];
1718    Ok([
1719        input_gt_lower,
1720        input_lt_upper,
1721        lower_gt_input,
1722        lower_lt_upper,
1723        upper_lt_max_input_lower,
1724    ])
1725}
1726
1727fn mask_ad_value(
1728    builder: &mut SemanticProgramBuilder,
1729    active: AdValue,
1730    conditions: &[ProgramValue],
1731    role: SemanticTransformRole,
1732) -> Result<AdValue, SemanticAdTransformError> {
1733    let AdValue::Value(active) = active else {
1734        return Ok(AdValue::Absent);
1735    };
1736    let zero = zero_constant_like(builder, active, role)?;
1737    let mut value = active;
1738    for condition in conditions {
1739        value = builder.add_op(CoreSemanticOp::Select, &[*condition, value, zero])?[0];
1740    }
1741    Ok(AdValue::Value(value))
1742}
1743
1744fn linearize_analytic_unary(
1745    builder: &mut SemanticProgramBuilder,
1746    op: &CoreSemanticOp,
1747    primal_input: ProgramValue,
1748    tangent: AdValue,
1749) -> Result<AdValue, SemanticAdTransformError> {
1750    if matches!(tangent, AdValue::Absent) {
1751        return Ok(AdValue::Absent);
1752    }
1753    let coefficient =
1754        analytic_unary_coefficient(builder, op, primal_input, SemanticTransformRole::Jvp)?;
1755    Ok(multiply_ad_value(builder, tangent, coefficient)?)
1756}
1757
1758fn linearize_sign(
1759    builder: &mut SemanticProgramBuilder,
1760    primal_input: ProgramValue,
1761    tangent: AdValue,
1762) -> Result<AdValue, SemanticAdTransformError> {
1763    let AdValue::Value(tangent_value) = tangent else {
1764        return Ok(AdValue::Absent);
1765    };
1766    let input_dtype = builder.value_metadata(primal_input)?.dtype();
1767    if !is_complex_dtype(input_dtype) {
1768        return Ok(AdValue::Absent);
1769    }
1770
1771    let zero = zero_constant_like(builder, primal_input, SemanticTransformRole::Jvp)?;
1772    let zero_mask = builder.add_op(
1773        CoreSemanticOp::Compare(CompareDir::Eq),
1774        &[primal_input, zero],
1775    )?[0];
1776    let sign = builder.add_op(CoreSemanticOp::Sign, &[primal_input])?[0];
1777    let abs = builder.add_op(CoreSemanticOp::Abs, &[primal_input])?[0];
1778    let output_dtype = abs_output_dtype(input_dtype);
1779    let complex_abs = builder.add_op(
1780        CoreSemanticOp::Convert {
1781            from: output_dtype,
1782            to: input_dtype,
1783        },
1784        &[abs],
1785    )?[0];
1786    let one = one_like(builder, complex_abs, SemanticTransformRole::Jvp)?;
1787    let safe_abs = builder.add_op(CoreSemanticOp::Select, &[zero_mask, one, complex_abs])?[0];
1788    let safe_sign = builder.add_op(CoreSemanticOp::Select, &[zero_mask, zero, sign])?[0];
1789    let conj_sign = builder.add_op(CoreSemanticOp::Conj, &[safe_sign])?[0];
1790
1791    let abs_tangent_complex = multiply_ad_value(builder, AdValue::Value(tangent_value), conj_sign)?;
1792    let abs_tangent = convert_ad_value(builder, abs_tangent_complex, input_dtype, output_dtype)?;
1793    let abs_tangent = convert_ad_value(builder, abs_tangent, output_dtype, input_dtype)?;
1794    let tangent_over_abs = divide_ad_value(builder, AdValue::Value(tangent_value), safe_abs)?;
1795    let sign_times_abs_tangent = multiply_ad_value(builder, abs_tangent, safe_sign)?;
1796    let correction = divide_ad_value(builder, sign_times_abs_tangent, safe_abs)?;
1797    let derivative = sub_ad_values(builder, tangent_over_abs, correction)?;
1798    let zero_derivative = zero_constant_like(builder, tangent_value, SemanticTransformRole::Jvp)?;
1799    select_ad_values(
1800        builder,
1801        zero_mask,
1802        AdValue::Value(zero_derivative),
1803        derivative,
1804    )
1805}
1806
1807fn analytic_unary_coefficient(
1808    builder: &mut SemanticProgramBuilder,
1809    op: &CoreSemanticOp,
1810    primal_input: ProgramValue,
1811    role: SemanticTransformRole,
1812) -> Result<ProgramValue, SemanticAdTransformError> {
1813    let coefficient = match op {
1814        CoreSemanticOp::Exp | CoreSemanticOp::Expm1 => {
1815            builder.add_op(CoreSemanticOp::Exp, &[primal_input])?[0]
1816        }
1817        CoreSemanticOp::Log => {
1818            let one = one_like(builder, primal_input, role)?;
1819            builder.add_op(CoreSemanticOp::Div, &[one, primal_input])?[0]
1820        }
1821        CoreSemanticOp::Sin => builder.add_op(CoreSemanticOp::Cos, &[primal_input])?[0],
1822        CoreSemanticOp::Cos => {
1823            let sin = builder.add_op(CoreSemanticOp::Sin, &[primal_input])?[0];
1824            builder.add_op(CoreSemanticOp::Neg, &[sin])?[0]
1825        }
1826        CoreSemanticOp::Tanh => {
1827            let tanh = builder.add_op(CoreSemanticOp::Tanh, &[primal_input])?[0];
1828            let square = builder.add_op(CoreSemanticOp::Mul, &[tanh, tanh])?[0];
1829            let one = one_like(builder, primal_input, role)?;
1830            builder.add_op(CoreSemanticOp::Sub, &[one, square])?[0]
1831        }
1832        CoreSemanticOp::Sqrt => {
1833            let sqrt = builder.add_op(CoreSemanticOp::Sqrt, &[primal_input])?[0];
1834            let twice = builder.add_op(CoreSemanticOp::Add, &[sqrt, sqrt])?[0];
1835            let one = one_like(builder, primal_input, role)?;
1836            builder.add_op(CoreSemanticOp::Div, &[one, twice])?[0]
1837        }
1838        CoreSemanticOp::Rsqrt => {
1839            let rsqrt = builder.add_op(CoreSemanticOp::Rsqrt, &[primal_input])?[0];
1840            let negated = builder.add_op(CoreSemanticOp::Neg, &[rsqrt])?[0];
1841            let twice = builder.add_op(CoreSemanticOp::Add, &[primal_input, primal_input])?[0];
1842            builder.add_op(CoreSemanticOp::Div, &[negated, twice])?[0]
1843        }
1844        CoreSemanticOp::Log1p => {
1845            let one = one_like(builder, primal_input, role)?;
1846            let denominator = builder.add_op(CoreSemanticOp::Add, &[primal_input, one])?[0];
1847            builder.add_op(CoreSemanticOp::Div, &[one, denominator])?[0]
1848        }
1849        CoreSemanticOp::Erf => {
1850            // d erf(x) = 2/sqrt(pi) * exp(-x^2)
1851            let square = builder.add_op(CoreSemanticOp::Mul, &[primal_input, primal_input])?[0];
1852            let negated = builder.add_op(CoreSemanticOp::Neg, &[square])?[0];
1853            let gaussian = builder.add_op(CoreSemanticOp::Exp, &[negated])?[0];
1854            let scale = real_float_constant_like(
1855                builder,
1856                primal_input,
1857                std::f64::consts::FRAC_2_SQRT_PI,
1858                role,
1859            )?;
1860            builder.add_op(CoreSemanticOp::Mul, &[scale, gaussian])?[0]
1861        }
1862        _ => return Err(unsupported_core(role, op)),
1863    };
1864    Ok(coefficient)
1865}
1866
1867fn one_like(
1868    builder: &mut SemanticProgramBuilder,
1869    anchor: ProgramValue,
1870    role: SemanticTransformRole,
1871) -> Result<ProgramValue, SemanticAdTransformError> {
1872    let metadata = builder.value_metadata(anchor)?.clone();
1873    let dtype = metadata.dtype();
1874    let bytes = match dtype {
1875        DType::F32 => 1.0_f32.to_le_bytes().to_vec(),
1876        DType::F64 => 1.0_f64.to_le_bytes().to_vec(),
1877        DType::C32 => {
1878            let mut bytes = 1.0_f32.to_le_bytes().to_vec();
1879            bytes.extend_from_slice(&0.0_f32.to_le_bytes());
1880            bytes
1881        }
1882        DType::C64 => {
1883            let mut bytes = 1.0_f64.to_le_bytes().to_vec();
1884            bytes.extend_from_slice(&0.0_f64.to_le_bytes());
1885            bytes
1886        }
1887        _ => {
1888            return Err(SemanticAdTransformError::UnsupportedMetadata {
1889                role,
1890                message: format!("cannot construct a differentiable one for {dtype:?}"),
1891            });
1892        }
1893    };
1894    let scalar = builder.add_op(CoreSemanticOp::Constant { dtype, bytes }, &[])?[0];
1895    broadcast_scalar_like(builder, scalar, anchor, &metadata, role, "one-like anchor")
1896}
1897
1898/// Broadcast a rank-0 `scalar` to `anchor`'s (possibly dynamic) shape.
1899fn broadcast_scalar_like(
1900    builder: &mut SemanticProgramBuilder,
1901    scalar: ProgramValue,
1902    anchor: ProgramValue,
1903    metadata: &ProgramValueMetadata,
1904    role: SemanticTransformRole,
1905    what: &'static str,
1906) -> Result<ProgramValue, SemanticAdTransformError> {
1907    if metadata.shape().is_empty() {
1908        Ok(scalar)
1909    } else {
1910        let shape = shape_plan(metadata.shape(), role, what)?;
1911        let value = broadcast_value_in_dim_to_shape(builder, scalar, anchor, &shape, Vec::new())?;
1912        Ok(truncate_value_to_dynamic_axes(
1913            builder,
1914            value,
1915            anchor,
1916            &shape.dynamic_axes,
1917        )?)
1918    }
1919}
1920
1921/// Build the real constant `value` shaped like a real floating `anchor`.
1922fn real_float_constant_like(
1923    builder: &mut SemanticProgramBuilder,
1924    anchor: ProgramValue,
1925    value: f64,
1926    role: SemanticTransformRole,
1927) -> Result<ProgramValue, SemanticAdTransformError> {
1928    let metadata = builder.value_metadata(anchor)?.clone();
1929    let dtype = metadata.dtype();
1930    let bytes = match dtype {
1931        DType::F32 => (value as f32).to_le_bytes().to_vec(),
1932        DType::F64 => value.to_le_bytes().to_vec(),
1933        _ => {
1934            return Err(SemanticAdTransformError::UnsupportedMetadata {
1935                role,
1936                message: format!("cannot construct a real floating constant for {dtype:?}"),
1937            });
1938        }
1939    };
1940    let scalar = builder.add_op(CoreSemanticOp::Constant { dtype, bytes }, &[])?[0];
1941    broadcast_scalar_like(
1942        builder,
1943        scalar,
1944        anchor,
1945        &metadata,
1946        role,
1947        "real-constant anchor",
1948    )
1949}
1950
1951/// Build a true dtype-aware zero shaped like `anchor`.
1952///
1953/// A zero synthesized as `x - x` or `x + (-x)` evaluates to `NaN` when `x` is
1954/// non-finite, so emission sites that need an additive identity for masked-out
1955/// values must materialize a `Constant` zero instead of deriving it from the
1956/// active value.
1957fn zero_constant_like(
1958    builder: &mut SemanticProgramBuilder,
1959    anchor: ProgramValue,
1960    role: SemanticTransformRole,
1961) -> Result<ProgramValue, SemanticAdTransformError> {
1962    let metadata = builder.value_metadata(anchor)?.clone();
1963    let dtype = metadata.dtype();
1964    let bytes = match dtype {
1965        DType::F32 => 0.0_f32.to_le_bytes().to_vec(),
1966        DType::F64 => 0.0_f64.to_le_bytes().to_vec(),
1967        DType::I32 => 0_i32.to_le_bytes().to_vec(),
1968        DType::I64 => 0_i64.to_le_bytes().to_vec(),
1969        DType::Bool => vec![0],
1970        DType::C32 => {
1971            let mut bytes = 0.0_f32.to_le_bytes().to_vec();
1972            bytes.extend_from_slice(&0.0_f32.to_le_bytes());
1973            bytes
1974        }
1975        DType::C64 => {
1976            let mut bytes = 0.0_f64.to_le_bytes().to_vec();
1977            bytes.extend_from_slice(&0.0_f64.to_le_bytes());
1978            bytes
1979        }
1980        // INVARIANT: the semantic AD catalog is closed to the preset scalars,
1981        // so an externally defined dtype cannot reach a constant emission.
1982        DType::External(_) => unreachable!("the AD catalog is closed to the presets"),
1983    };
1984    let scalar = builder.add_op(CoreSemanticOp::Constant { dtype, bytes }, &[])?[0];
1985    if metadata.shape().is_empty() {
1986        return Ok(scalar);
1987    }
1988    let shape = shape_plan(metadata.shape(), role, "zero-like anchor")?;
1989    let zero = broadcast_value_in_dim_to_shape(builder, scalar, anchor, &shape, Vec::new())?;
1990    Ok(truncate_value_to_dynamic_axes(
1991        builder,
1992        zero,
1993        anchor,
1994        &shape.dynamic_axes,
1995    )?)
1996}
1997
1998fn active_cotangent(
1999    builder: &mut SemanticProgramBuilder,
2000    cotangent: AdValue,
2001    active: bool,
2002    primal_input: ProgramValue,
2003) -> Result<AdValue, SemanticAdTransformError> {
2004    normalize_ad_value(builder, cotangent, active, primal_input)
2005}
2006
2007fn normalize_ad_value(
2008    builder: &mut SemanticProgramBuilder,
2009    value: AdValue,
2010    active: bool,
2011    primal_input: ProgramValue,
2012) -> Result<AdValue, SemanticAdTransformError> {
2013    if !active {
2014        return Ok(AdValue::Absent);
2015    }
2016    let AdValue::Value(mut value) = value else {
2017        return Ok(AdValue::Absent);
2018    };
2019    let target_metadata = builder.value_metadata(primal_input)?.clone();
2020    let value_metadata = builder.value_metadata(value)?.clone();
2021    let target_shape = shape_plan(
2022        target_metadata.shape(),
2023        SemanticTransformRole::Vjp,
2024        "primal input",
2025    )?;
2026    let value_shape = shape_plan(
2027        value_metadata.shape(),
2028        SemanticTransformRole::Vjp,
2029        "cotangent",
2030    )?;
2031    if value_shape.shape.len() < target_shape.shape.len() {
2032        return Err(SemanticAdTransformError::UnsupportedMetadata {
2033            role: SemanticTransformRole::Vjp,
2034            message: "cotangent rank is smaller than its primal-input rank".into(),
2035        });
2036    }
2037    let leading = value_shape.shape.len() - target_shape.shape.len();
2038    let mut axes: Vec<_> = (0..leading).collect();
2039    axes.extend(
2040        target_shape
2041            .shape
2042            .iter()
2043            .zip(value_shape.shape.iter().skip(leading))
2044            .enumerate()
2045            .filter_map(|(axis, (target, actual))| {
2046                (matches!(target, tenferro_ops::dim_expr::DimExpr::Const(1)) && target != actual)
2047                    .then_some(axis + leading)
2048            }),
2049    );
2050    if !axes.is_empty() {
2051        value = builder.add_op(CoreSemanticOp::ReduceSum { axes }, &[value])?[0];
2052    }
2053    if builder.value_metadata(value)?.shape() != target_metadata.shape() {
2054        value = reshape_value_to_shape(builder, value, primal_input, &target_shape)?;
2055    }
2056    value =
2057        truncate_value_to_dynamic_axes(builder, value, primal_input, &target_shape.dynamic_axes)?;
2058    let value_dtype = builder.value_metadata(value)?.dtype();
2059    if value_dtype != target_metadata.dtype() {
2060        value = builder.add_op(
2061            CoreSemanticOp::Convert {
2062                from: value_dtype,
2063                to: target_metadata.dtype(),
2064            },
2065            &[value],
2066        )?[0];
2067    }
2068    Ok(AdValue::Value(value))
2069}
2070
2071fn exact_shape(
2072    shape: &[tenferro_ops::ShapeExtent<tenferro_ops::dim_expr::DimExpr>],
2073    role: SemanticTransformRole,
2074    field: &'static str,
2075) -> Result<Vec<tenferro_ops::dim_expr::DimExpr>, SemanticAdTransformError> {
2076    shape
2077        .iter()
2078        .map(|extent| {
2079            extent.as_exact().cloned().ok_or_else(|| {
2080                SemanticAdTransformError::UnsupportedMetadata {
2081                    role,
2082                    message: format!("{field} has a bounded or unknown extent"),
2083                }
2084            })
2085        })
2086        .collect()
2087}
2088
2089fn value_shape_plan(
2090    builder: &SemanticProgramBuilder,
2091    value: ProgramValue,
2092    role: SemanticTransformRole,
2093    field: &'static str,
2094) -> Result<ValueShapePlan, SemanticAdTransformError> {
2095    shape_plan(builder.value_metadata(value)?.shape(), role, field)
2096}
2097
2098fn shape_plan(
2099    shape: &[ShapeExtent<DimExpr>],
2100    role: SemanticTransformRole,
2101    field: &'static str,
2102) -> Result<ValueShapePlan, SemanticAdTransformError> {
2103    let mut planned_shape = Vec::with_capacity(shape.len());
2104    let mut dynamic_axes = Vec::new();
2105    for (axis, extent) in shape.iter().enumerate() {
2106        match extent {
2107            ShapeExtent::Exact(expression) => planned_shape.push(expression.clone()),
2108            ShapeExtent::UpperBound(expression) => {
2109                planned_shape.push(expression.clone());
2110                dynamic_axes.push(axis);
2111            }
2112            ShapeExtent::Unknown => {
2113                return Err(SemanticAdTransformError::UnsupportedMetadata {
2114                    role,
2115                    message: format!("{field} has an unknown extent without an upper bound"),
2116                });
2117            }
2118        }
2119    }
2120    Ok(ValueShapePlan {
2121        shape: planned_shape,
2122        dynamic_axes,
2123    })
2124}
2125
2126fn reshape_ad_value_to_shape(
2127    builder: &mut SemanticProgramBuilder,
2128    value: AdValue,
2129    shape_source: ProgramValue,
2130    shape: &ValueShapePlan,
2131) -> Result<AdValue, ProgramBuildError> {
2132    match value {
2133        AdValue::Absent => Ok(AdValue::Absent),
2134        AdValue::Value(value) => Ok(AdValue::Value(reshape_value_to_shape(
2135            builder,
2136            value,
2137            shape_source,
2138            shape,
2139        )?)),
2140    }
2141}
2142
2143fn reshape_value_to_shape(
2144    builder: &mut SemanticProgramBuilder,
2145    value: ProgramValue,
2146    shape_source: ProgramValue,
2147    shape: &ValueShapePlan,
2148) -> Result<ProgramValue, ProgramBuildError> {
2149    let mut inputs = vec![value];
2150    let to_shape = payload_shape_for_shape_source(shape, &mut inputs, shape_source);
2151    Ok(builder.add_op(CoreSemanticOp::Reshape { to_shape }, &inputs)?[0])
2152}
2153
2154fn broadcast_ad_value_in_dim_to_shape(
2155    builder: &mut SemanticProgramBuilder,
2156    value: AdValue,
2157    shape_source: ProgramValue,
2158    shape: &ValueShapePlan,
2159    dims: Vec<usize>,
2160) -> Result<AdValue, ProgramBuildError> {
2161    match value {
2162        AdValue::Absent => Ok(AdValue::Absent),
2163        AdValue::Value(value) => Ok(AdValue::Value(broadcast_value_in_dim_to_shape(
2164            builder,
2165            value,
2166            shape_source,
2167            shape,
2168            dims,
2169        )?)),
2170    }
2171}
2172
2173fn broadcast_value_in_dim_to_shape(
2174    builder: &mut SemanticProgramBuilder,
2175    value: ProgramValue,
2176    shape_source: ProgramValue,
2177    shape: &ValueShapePlan,
2178    dims: Vec<usize>,
2179) -> Result<ProgramValue, ProgramBuildError> {
2180    let mut inputs = vec![value];
2181    let shape = payload_shape_for_shape_source(shape, &mut inputs, shape_source);
2182    Ok(builder.add_op(CoreSemanticOp::BroadcastInDim { shape, dims }, &inputs)?[0])
2183}
2184
2185fn payload_shape_for_shape_source(
2186    shape: &ValueShapePlan,
2187    inputs: &mut Vec<ProgramValue>,
2188    shape_source: ProgramValue,
2189) -> Vec<DimExpr> {
2190    if DimExpr::max_input_idx_all(&shape.shape).is_none() {
2191        return shape.shape.clone();
2192    }
2193    let input_idx = inputs
2194        .iter()
2195        .position(|&input| input == shape_source)
2196        .unwrap_or_else(|| {
2197            let input_idx = inputs.len();
2198            inputs.push(shape_source);
2199            input_idx
2200        });
2201    DimExpr::input_shape(input_idx, shape.shape.len())
2202}
2203
2204fn truncate_ad_value_to_dynamic_axes(
2205    builder: &mut SemanticProgramBuilder,
2206    value: AdValue,
2207    shape_source: ProgramValue,
2208    dynamic_axes: &[usize],
2209) -> Result<AdValue, SemanticAdTransformError> {
2210    let AdValue::Value(value) = value else {
2211        return Ok(AdValue::Absent);
2212    };
2213    Ok(AdValue::Value(truncate_value_to_dynamic_axes(
2214        builder,
2215        value,
2216        shape_source,
2217        dynamic_axes,
2218    )?))
2219}
2220
2221fn truncate_value_to_dynamic_axes(
2222    builder: &mut SemanticProgramBuilder,
2223    mut value: ProgramValue,
2224    shape_source: ProgramValue,
2225    dynamic_axes: &[usize],
2226) -> Result<ProgramValue, ProgramBuildError> {
2227    for &axis in dynamic_axes {
2228        let size = builder.add_op(CoreSemanticOp::ShapeOf { axis }, &[shape_source])?[0];
2229        value = builder.add_op(CoreSemanticOp::DynamicTruncate { axis }, &[value, size])?[0];
2230    }
2231    Ok(value)
2232}
2233
2234fn conjugate_if_complex(
2235    builder: &mut SemanticProgramBuilder,
2236    value: ProgramValue,
2237) -> Result<ProgramValue, ProgramBuildError> {
2238    if matches!(
2239        builder.value_metadata(value)?.dtype(),
2240        DType::C32 | DType::C64
2241    ) {
2242        Ok(builder.add_op(CoreSemanticOp::Conj, &[value])?[0])
2243    } else {
2244        Ok(value)
2245    }
2246}
2247
2248fn finish_derivative(
2249    builder: SemanticProgramBuilder,
2250    derivative_input_indices: Vec<Option<usize>>,
2251    values: Vec<AdValue>,
2252) -> Result<SemanticAdProgram, SemanticAdTransformError> {
2253    let mut outputs = Vec::new();
2254    let derivative_output_indices = values
2255        .into_iter()
2256        .map(|value| match value {
2257            AdValue::Absent => None,
2258            AdValue::Value(value) => {
2259                let index = outputs.len();
2260                outputs.push(value);
2261                Some(index)
2262            }
2263        })
2264        .collect();
2265    let frozen = builder.finish(&outputs)?;
2266    let frozen = prune_dead_derivative_operations(frozen)?;
2267    let frozen = cancel_double_neg_derivative_operations(frozen)?;
2268    Ok(SemanticAdProgram {
2269        frozen,
2270        derivative_input_indices: derivative_input_indices.into_boxed_slice(),
2271        derivative_output_indices,
2272    })
2273}
2274
2275fn prune_dead_derivative_operations(
2276    frozen: FrozenProgram,
2277) -> Result<FrozenProgram, SemanticAdTransformError> {
2278    let mut roots = frozen.program.inputs().to_vec();
2279    let output_offset = roots.len();
2280    roots.extend_from_slice(frozen.program.outputs());
2281
2282    let mut builder = SemanticProgramBuilder::new();
2283    let imported = builder.import(ProgramImport {
2284        program: frozen.program.as_ref(),
2285        bindings: &frozen.bindings,
2286        roots: &roots,
2287    })?;
2288    let outputs = imported.roots()[output_offset..].to_vec();
2289    Ok(builder.finish(&outputs)?)
2290}
2291
2292fn cancel_double_neg_derivative_operations(
2293    frozen: FrozenProgram,
2294) -> Result<FrozenProgram, SemanticAdTransformError> {
2295    let operations = frozen.program.operations().collect::<Vec<_>>();
2296    if operations.iter().any(|operation| {
2297        !operation.effects().is_empty()
2298            || !operation.shape_guards().is_empty()
2299            || !matches!(operation.op(), SemanticOpRef::Core(_))
2300    }) {
2301        return Ok(frozen);
2302    }
2303
2304    let mut builder = SemanticProgramBuilder::new();
2305    let imported = builder.import(ProgramImport {
2306        program: frozen.program.as_ref(),
2307        bindings: &frozen.bindings,
2308        roots: frozen.program.inputs(),
2309    })?;
2310    let mut values = frozen
2311        .program
2312        .inputs()
2313        .iter()
2314        .copied()
2315        .zip(imported.roots().iter().copied())
2316        .collect::<HashMap<_, _>>();
2317    let mut neg_inputs = HashMap::<ProgramValue, ProgramValue>::new();
2318    let mut changed = false;
2319
2320    for operation in operations {
2321        let inputs = operation
2322            .inputs()
2323            .iter()
2324            .copied()
2325            .map(|value| {
2326                values.get(&value).copied().ok_or_else(|| {
2327                    SemanticAdTransformError::UnsupportedMetadata {
2328                        role: SemanticTransformRole::Jvp,
2329                        message: "derivative simplifier saw an unmapped value".into(),
2330                    }
2331                })
2332            })
2333            .collect::<Result<Vec<_>, _>>()?;
2334        let SemanticOpRef::Core(op) = operation.op() else {
2335            unreachable!("non-core operations returned above");
2336        };
2337
2338        if matches!(op, CoreSemanticOp::Neg) {
2339            let input = inputs[0];
2340            if let Some(inner) = neg_inputs.get(&input).copied() {
2341                values.insert(operation.outputs()[0], inner);
2342                changed = true;
2343                continue;
2344            }
2345            let output = builder.add_op(CoreSemanticOp::Neg, &[input])?[0];
2346            neg_inputs.insert(output, input);
2347            values.insert(operation.outputs()[0], output);
2348            continue;
2349        }
2350
2351        let outputs = builder.add_op(op.clone(), &inputs)?;
2352        for (source, output) in operation
2353            .outputs()
2354            .iter()
2355            .copied()
2356            .zip(outputs.iter().copied())
2357        {
2358            values.insert(source, output);
2359        }
2360    }
2361
2362    if !changed {
2363        return Ok(frozen);
2364    }
2365
2366    let outputs = frozen
2367        .program
2368        .outputs()
2369        .iter()
2370        .copied()
2371        .map(|value| {
2372            values.get(&value).copied().ok_or_else(|| {
2373                SemanticAdTransformError::UnsupportedMetadata {
2374                    role: SemanticTransformRole::Jvp,
2375                    message: "derivative simplifier saw an unmapped output".into(),
2376                }
2377            })
2378        })
2379        .collect::<Result<Vec<_>, _>>()?;
2380    prune_dead_derivative_operations(builder.finish(&outputs)?)
2381}
2382
2383fn validate_activity(
2384    role: SemanticTransformRole,
2385    field: &'static str,
2386    expected: usize,
2387    actual: usize,
2388) -> Result<(), SemanticAdTransformError> {
2389    if expected == actual {
2390        Ok(())
2391    } else {
2392        Err(SemanticAdTransformError::ActivityArity {
2393            role,
2394            field,
2395            expected,
2396            actual,
2397        })
2398    }
2399}
2400
2401fn unsupported_core(role: SemanticTransformRole, op: &CoreSemanticOp) -> SemanticAdTransformError {
2402    SemanticAdTransformError::UnsupportedCore {
2403        role,
2404        op: format!("{op:?}"),
2405    }
2406}