1mod 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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
26pub enum SemanticTransformRole {
27 Jvp,
29 Vjp,
31}
32
33#[derive(Debug, thiserror::Error)]
35pub enum SemanticAdTransformError {
36 #[error("semantic {role:?} {field} expects {expected} entries, got {actual}")]
38 ActivityArity {
39 role: SemanticTransformRole,
41 field: &'static str,
43 expected: usize,
45 actual: usize,
47 },
48 #[error("semantic {role:?} does not support active core operation {op}")]
50 UnsupportedCore {
51 role: SemanticTransformRole,
53 op: String,
55 },
56 #[error("semantic {role:?} does not support this semantic operation variant")]
58 UnsupportedOperationVariant {
59 role: SemanticTransformRole,
61 },
62 #[error("semantic {role:?} does not support derivative metadata: {message}")]
64 UnsupportedMetadata {
65 role: SemanticTransformRole,
67 message: String,
69 },
70 #[error("semantic AD source-program query failed: {0}")]
72 Query(#[from] ProgramQueryError),
73 #[error("semantic AD program construction failed: {0}")]
75 Build(#[from] ProgramBuildError),
76 #[error("semantic extension AD failed: {0}")]
78 Extension(#[from] SemanticAdError),
79 #[error("semantic AD program finalization failed: {0}")]
81 Finish(#[from] ProgramFinishError),
82 #[error("semantic AD transform cache failed: {0}")]
84 Cache(#[source] tenferro_runtime::Error),
85}
86
87#[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 pub const fn frozen(&self) -> &FrozenProgram {
107 &self.frozen
108 }
109
110 pub fn derivative_input_indices(&self) -> &[Option<usize>] {
112 &self.derivative_input_indices
113 }
114
115 pub fn derivative_output_indices(&self) -> &[Option<usize>] {
117 &self.derivative_output_indices
118 }
119
120 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
137pub 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
236pub 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
258pub(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 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
639pub(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 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 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
1898fn 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
1921fn 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
1951fn 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 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}