1use std::collections::{HashMap, HashSet};
5use std::sync::Arc;
6
7use computegraph::traits::GraphOperation;
8use computegraph::types::{LocalValueId, OperationRole, ValueKey, ValueRef};
9use tenferro_ad::semantic_extension::{
10 AdValue, ResidualSpec, SemanticAdError, SemanticAdRuleRole, SemanticExtensionRegistryError,
11 SemanticExtensionRuleSet, SemanticLinearTransposeRequest, SemanticLinearTransposeRule,
12 SemanticLinearizeRequest, SemanticLinearizeResult, SemanticLinearizeRule,
13};
14use tenferro_ops::ad::PrimitiveRuleBuilder;
15use tenferro_ops::ad::PrimitiveTransposeInput;
16use tenferro_ops::dim_expr::DimExpr;
17use tenferro_ops::input_key::TensorInputKey;
18use tenferro_ops::shape_extent::ShapeExtent;
19use tenferro_ops::std_tensor_op::StdTensorOp;
20use tenferro_ops::{ShapeGuardContext, SymDim, TensorMeta};
21use tenferro_runtime::program::{
22 CoreSemanticOp, ProgramValue, ProgramValueMetadata, SemanticProgramBuilder,
23};
24
25use super::LinalgAdRule;
26use crate::extension::{LinalgExtensionOp, LinalgOp};
27use crate::LINALG_EXTENSION_FAMILY_ID;
28
29pub fn semantic_ad_rules() -> Result<SemanticExtensionRuleSet, SemanticExtensionRegistryError> {
53 SemanticExtensionRuleSet::new()
54 .with_linearize(Arc::new(LinalgAdRule))?
55 .with_linear_transpose(Arc::new(LinalgAdRule))
56}
57
58impl SemanticLinearizeRule for LinalgAdRule {
59 fn family_id(&self) -> &'static str {
60 LINALG_EXTENSION_FAMILY_ID
61 }
62
63 fn linearize(
64 &self,
65 request: SemanticLinearizeRequest<'_>,
66 builder: &mut SemanticProgramBuilder,
67 ) -> Result<SemanticLinearizeResult, SemanticAdError> {
68 let op = semantic_linalg_op(request.op(), SemanticAdRuleRole::Linearize)?;
69 if matches!(op.op(), LinalgOp::HouseholderQrThinQ { .. }) {
70 return Err(SemanticAdError::Unsupported {
71 family_id: LINALG_EXTENSION_FAMILY_ID,
72 role: SemanticAdRuleRole::Linearize,
73 message: "internal thin-Q residual is not differentiable".into(),
74 });
75 }
76 if matches!(op.op(), LinalgOp::LuFactor | LinalgOp::SvdFull) {
77 return Ok(SemanticLinearizeResult::new(
86 std::iter::repeat_n(AdValue::Absent, request.primal_outputs().len()),
87 [],
88 ));
89 }
90 let legacy = LegacyInvocation::new(
91 request.primal_inputs(),
92 request.primal_outputs(),
93 request.active_outputs(),
94 builder,
95 )?;
96 let seed_values: Vec<_> = request
97 .tangent_inputs()
98 .iter()
99 .copied()
100 .map(AdValue::value)
101 .collect();
102 let tangent_inputs: Vec<_> = seed_values
103 .iter()
104 .enumerate()
105 .map(|(index, value)| value.map(|_| index))
106 .collect();
107 let mut emitted = SemanticRuleBuilder::with_seeds(
108 &seed_values,
109 &legacy.external_values,
110 &legacy.shape_sources,
111 builder,
112 SemanticAdRuleRole::Linearize,
113 );
114 let tangent_outputs = LinalgAdRule
115 .linearize(
116 op,
117 &mut emitted,
118 &legacy.input_keys,
119 &legacy.output_keys,
120 &tangent_inputs,
121 &mut legacy.context.clone(),
122 )
123 .map_err(|error| legacy_error(SemanticAdRuleRole::Linearize, error))?;
124 let locals = emitted.finish()?;
125 Ok(SemanticLinearizeResult::new(
126 tangent_outputs.into_iter().map(|value| {
127 value
128 .and_then(|local| locals.get(local).copied().flatten())
129 .map_or(AdValue::Absent, AdValue::Value)
130 }),
131 [],
132 ))
133 }
134}
135
136impl SemanticLinearTransposeRule for LinalgAdRule {
137 fn family_id(&self) -> &'static str {
138 LINALG_EXTENSION_FAMILY_ID
139 }
140
141 fn residual_mask(&self) -> ResidualSpec {
142 ResidualSpec::all_inputs().with_all_outputs()
150 }
151
152 fn linear_transpose(
153 &self,
154 request: SemanticLinearTransposeRequest<'_>,
155 builder: &mut SemanticProgramBuilder,
156 ) -> Result<Box<[AdValue]>, SemanticAdError> {
157 let primal_inputs = (0..request.primal_input_count())
158 .map(|index| request.primal_input_value(index))
159 .collect::<Result<Vec<_>, _>>()?;
160 let primal_outputs = (0..request.primal_output_count())
161 .map(|index| request.primal_output_value(index))
162 .collect::<Result<Vec<_>, _>>()?;
163 let op = semantic_linalg_op(request.op(), SemanticAdRuleRole::LinearTranspose)?;
164 match op.op() {
165 LinalgOp::TriangularSolve {
166 left_side,
167 lower,
168 transpose_a,
169 unit_diagonal,
170 } => semantic_triangular_solve_transpose(
171 &primal_inputs,
172 &primal_outputs,
173 request.cotangent_outputs(),
174 request.active_inputs(),
175 request.residual_mask(),
176 builder,
177 left_side,
178 lower,
179 transpose_a,
180 unit_diagonal,
181 ),
182 LinalgOp::LuSolvePrepared {
183 transpose_a,
184 conjugate_a,
185 } => {
186 let active_inputs = lu_solve_prepared_transpose_active_inputs(
187 request.active_inputs(),
188 SemanticAdRuleRole::LinearTranspose,
189 )?;
190 let mut result = vec![AdValue::Absent; 4];
191 let (a_cotangent, b_cotangent) = semantic_prepared_solve_transpose(
192 builder,
193 SemanticPreparedSolve {
194 op: "lu_solve_prepared",
195 a: primal_inputs[0],
196 packed_lu: primal_inputs[1],
197 pivots: primal_inputs[2],
198 solution: primal_outputs.first().copied(),
199 },
200 request.cotangent_outputs(),
201 (active_inputs[0], active_inputs[3]),
202 transpose_a,
203 conjugate_a,
204 )?;
205 result[0] = a_cotangent;
206 result[3] = b_cotangent;
207 Ok(result.into_boxed_slice())
208 }
209 LinalgOp::LuFactorSolve => {
213 let active_inputs = request.active_inputs();
214 let (Some(&a_active), Some(&b_active), Some(&packed_lu), Some(&pivots)) = (
215 active_inputs.first(),
216 active_inputs.get(1),
217 primal_outputs.get(1),
218 primal_outputs.get(2),
219 ) else {
220 return Err(semantic_internal(
221 SemanticAdRuleRole::LinearTranspose,
222 "lu_factor_solve transpose expected inputs (a, b) and outputs (x, lu, pivots)",
223 ));
224 };
225 let (a_cotangent, b_cotangent) = semantic_prepared_solve_transpose(
226 builder,
227 SemanticPreparedSolve {
228 op: "lu_factor_solve",
229 a: primal_inputs[0],
230 packed_lu,
231 pivots,
232 solution: primal_outputs.first().copied(),
233 },
234 request.cotangent_outputs(),
235 (a_active, b_active),
236 false,
237 false,
238 )?;
239 Ok(vec![a_cotangent, b_cotangent].into_boxed_slice())
240 }
241 LinalgOp::FullPivLuSolve { .. } => semantic_custom_transpose(
242 request.op(),
243 &primal_inputs,
244 &primal_outputs,
245 request.cotangent_outputs(),
246 request.active_inputs(),
247 builder,
248 SemanticAdRuleRole::LinearTranspose,
249 ),
250 LinalgOp::Solve => semantic_custom_transpose(
251 request.op(),
252 &primal_inputs,
253 &primal_outputs,
254 request.cotangent_outputs(),
255 request.active_inputs(),
256 builder,
257 SemanticAdRuleRole::LinearTranspose,
258 ),
259 LinalgOp::LuFactor | LinalgOp::SvdFull | LinalgOp::HouseholderQrThinQ { .. } => {
260 Err(SemanticAdError::Unsupported {
261 family_id: LINALG_EXTENSION_FAMILY_ID,
262 role: SemanticAdRuleRole::LinearTranspose,
263 message: format!("semantic linear transpose is unsupported for {:?}", op.op()),
264 })
265 }
266 _ => semantic_linearized_transpose(
267 request.op(),
268 &primal_inputs,
269 &primal_outputs,
270 request.cotangent_outputs(),
271 request.active_inputs(),
272 builder,
273 ),
274 }
275 }
276}
277
278fn lu_solve_prepared_transpose_active_inputs(
279 active_inputs: &[bool],
280 role: SemanticAdRuleRole,
281) -> Result<[bool; 4], SemanticAdError> {
282 let active_inputs: [bool; 4] = active_inputs.try_into().map_err(|_| {
283 semantic_internal(
284 role,
285 format!(
286 "lu_solve_prepared semantic transpose expected 4 active inputs, got {}",
287 active_inputs.len()
288 ),
289 )
290 })?;
291 Ok([active_inputs[0], false, false, active_inputs[3]])
295}
296
297struct SemanticPreparedSolve {
299 op: &'static str,
300 a: ProgramValue,
301 packed_lu: ProgramValue,
302 pivots: ProgramValue,
303 solution: Option<ProgramValue>,
304}
305
306fn semantic_prepared_solve_transpose(
312 builder: &mut SemanticProgramBuilder,
313 primal: SemanticPreparedSolve,
314 cotangent_outputs: &[AdValue],
315 (a_active, b_active): (bool, bool),
316 transpose_a: bool,
317 conjugate_a: bool,
318) -> Result<(AdValue, AdValue), SemanticAdError> {
319 let Some(ct) = cotangent_outputs.first().copied().and_then(AdValue::value) else {
320 return Ok((AdValue::Absent, AdValue::Absent));
321 };
322 if !a_active && !b_active {
323 return Ok((AdValue::Absent, AdValue::Absent));
324 }
325 let rhs_cotangent = builder.add_extension(
326 Arc::new(LinalgExtensionOp::new(LinalgOp::LuSolvePrepared {
327 transpose_a: !transpose_a,
328 conjugate_a: !conjugate_a,
329 })),
330 &[primal.a, primal.packed_lu, primal.pivots, ct],
331 )?[0];
332 let a_cotangent = if a_active {
333 let solution = primal.solution.ok_or_else(|| {
334 semantic_internal(
335 SemanticAdRuleRole::LinearTranspose,
336 format!("{} transpose requires its primal solution", primal.op),
337 )
338 })?;
339 let rank = builder.value_metadata(primal.a)?.shape().len();
340 let matrix_cotangent = semantic_solve_matrix_cotangent(
341 builder,
342 rhs_cotangent,
343 solution,
344 true,
345 transpose_a,
346 rank,
347 )?;
348 AdValue::Value(if conjugate_a {
349 conjugate_if_complex(builder, matrix_cotangent)?
350 } else {
351 matrix_cotangent
352 })
353 } else {
354 AdValue::Absent
355 };
356 let b_cotangent = if b_active {
357 AdValue::Value(rhs_cotangent)
358 } else {
359 AdValue::Absent
360 };
361 Ok((a_cotangent, b_cotangent))
362}
363
364#[allow(clippy::too_many_arguments)]
365fn semantic_triangular_solve_transpose(
366 primal_inputs: &[ProgramValue],
367 primal_outputs: &[ProgramValue],
368 cotangent_outputs: &[AdValue],
369 active_inputs: &[bool],
370 residual_mask: ResidualSpec,
371 builder: &mut SemanticProgramBuilder,
372 left_side: bool,
373 lower: bool,
374 transpose_a: bool,
375 unit_diagonal: bool,
376) -> Result<Box<[AdValue]>, SemanticAdError> {
377 let role = SemanticAdRuleRole::LinearTranspose;
378 if primal_inputs.len() != 2
379 || primal_outputs.len() != 1
380 || cotangent_outputs.len() != 1
381 || active_inputs.len() != 2
382 {
383 return Err(semantic_internal(
384 role,
385 "triangular_solve semantic transpose received malformed arity",
386 ));
387 }
388 let Some(ct) = cotangent_outputs.first().copied().and_then(AdValue::value) else {
389 return Ok(vec![AdValue::Absent; 2].into_boxed_slice());
390 };
391
392 let mut result = vec![AdValue::Absent; 2];
393 if !active_inputs[0] && !active_inputs[1] {
394 return Ok(result.into_boxed_slice());
395 }
396
397 let matrix_rank = builder.value_metadata(primal_inputs[0])?.shape().len();
398 let rhs_rank = builder.value_metadata(primal_inputs[1])?.shape().len();
399 if matrix_rank < 2 || rhs_rank < 2 {
400 return Err(semantic_internal(
401 role,
402 "triangular_solve semantic transpose expects matrix operands",
403 ));
404 }
405 if matrix_rank != rhs_rank {
406 return Err(semantic_internal(
407 role,
408 "triangular_solve semantic transpose expects equal-rank operands",
409 ));
410 }
411
412 let conjugated_a = conjugate_if_complex(builder, primal_inputs[0])?;
413 debug_assert!(
414 residual_mask.declares_input(0),
415 "linalg triangular_solve transpose read primal input 0 as a tensor operand but the \
416 residual mask does not declare it; declare it in the linalg rule's residual mask"
417 );
418 let rhs_cotangent = builder.add_extension(
419 Arc::new(LinalgExtensionOp::new(LinalgOp::TriangularSolve {
420 left_side,
421 lower,
422 transpose_a: !transpose_a,
423 unit_diagonal,
424 })),
425 &[conjugated_a, ct],
426 )?[0];
427
428 if active_inputs[1] {
429 result[1] = AdValue::Value(rhs_cotangent);
430 }
431 if active_inputs[0] {
432 debug_assert!(
433 residual_mask.declares_output(0),
434 "linalg triangular_solve transpose read primal output 0 as a tensor operand but the \
435 residual mask does not declare it; declare it in the linalg rule's residual mask"
436 );
437 let matrix_cotangent = semantic_solve_matrix_cotangent(
438 builder,
439 rhs_cotangent,
440 primal_outputs[0],
441 left_side,
442 transpose_a,
443 matrix_rank,
444 )?;
445 let k = if unit_diagonal {
446 if lower {
447 -1
448 } else {
449 1
450 }
451 } else {
452 0
453 };
454 let projected = if lower {
455 builder.add_op(CoreSemanticOp::Tril { k }, &[matrix_cotangent])?[0]
456 } else {
457 builder.add_op(CoreSemanticOp::Triu { k }, &[matrix_cotangent])?[0]
458 };
459 result[0] = AdValue::Value(projected);
460 }
461
462 Ok(result.into_boxed_slice())
463}
464
465fn semantic_linearized_transpose(
466 op: &dyn tenferro_ad::extension::ExtensionOp,
467 primal_inputs: &[ProgramValue],
468 primal_outputs: &[ProgramValue],
469 cotangent_outputs: &[AdValue],
470 active_inputs: &[bool],
471 builder: &mut SemanticProgramBuilder,
472) -> Result<Box<[AdValue]>, SemanticAdError> {
473 let legacy = LegacyInvocation::new(
474 primal_inputs,
475 primal_outputs,
476 &cotangent_outputs
477 .iter()
478 .map(|value| matches!(value, AdValue::Value(_)))
479 .collect::<Vec<_>>(),
480 builder,
481 )?;
482 let tangent_inputs: Vec<_> = active_inputs
483 .iter()
484 .copied()
485 .enumerate()
486 .map(|(index, active)| active.then_some(index))
487 .collect();
488 let mut fragment = SemanticLinearFragmentBuilder::with_seed_count(primal_inputs.len());
489 let tangent_outputs = LinalgAdRule
490 .linearize(
491 op,
492 &mut fragment,
493 &legacy.input_keys,
494 &legacy.output_keys,
495 &tangent_inputs,
496 &mut legacy.context.clone(),
497 )
498 .map_err(|error| legacy_error(SemanticAdRuleRole::LinearTranspose, error))?;
499 fragment.transpose_linear_fragment(
500 &tangent_outputs,
501 cotangent_outputs,
502 active_inputs,
503 &legacy.external_values,
504 &legacy.shape_sources,
505 builder,
506 )
507}
508
509fn semantic_custom_transpose(
510 op: &dyn tenferro_ad::extension::ExtensionOp,
511 primal_inputs: &[ProgramValue],
512 primal_outputs: &[ProgramValue],
513 cotangent_outputs: &[AdValue],
514 active_inputs: &[bool],
515 builder: &mut SemanticProgramBuilder,
516 role: SemanticAdRuleRole,
517) -> Result<Box<[AdValue]>, SemanticAdError> {
518 let legacy = LegacyInvocation::new(
519 primal_inputs,
520 primal_outputs,
521 &vec![true; primal_outputs.len()],
522 builder,
523 )?;
524 let seed_values: Vec<_> = cotangent_outputs
525 .iter()
526 .copied()
527 .map(AdValue::value)
528 .collect();
529 let cotangents: Vec<_> = seed_values
530 .iter()
531 .enumerate()
532 .map(|(index, value)| value.map(|_| index))
533 .collect();
534 let transpose_inputs: Vec<_> = legacy
535 .input_keys
536 .iter()
537 .cloned()
538 .map(PrimitiveTransposeInput::Residual)
539 .collect();
540 let mut emitted = SemanticRuleBuilder::with_seeds(
541 &seed_values,
542 &legacy.external_values,
543 &legacy.shape_sources,
544 builder,
545 role,
546 );
547 let cotangent_inputs = LinalgAdRule
548 .linear_transpose(
549 op,
550 &mut emitted,
551 &cotangents,
552 &transpose_inputs,
553 active_inputs,
554 &mut legacy.context.clone(),
555 )
556 .map_err(|error| legacy_error(role, error))?;
557 let locals = emitted.finish()?;
558 Ok(cotangent_inputs
559 .into_iter()
560 .map(|value| {
561 value
562 .and_then(|local| locals.get(local).copied().flatten())
563 .map_or(AdValue::Absent, AdValue::Value)
564 })
565 .collect())
566}
567
568struct LegacyInvocation {
569 context: ShapeGuardContext,
570 input_keys: Vec<ValueKey<StdTensorOp>>,
571 output_keys: Vec<ValueKey<StdTensorOp>>,
572 external_values: HashMap<ValueKey<StdTensorOp>, ProgramValue>,
573 shape_sources: Vec<ProgramValue>,
574}
575
576impl LegacyInvocation {
577 fn new(
578 primal_inputs: &[ProgramValue],
579 primal_outputs: &[ProgramValue],
580 active_outputs: &[bool],
581 builder: &SemanticProgramBuilder,
582 ) -> Result<Self, SemanticAdError> {
583 let values: Vec<_> = primal_inputs
584 .iter()
585 .chain(primal_outputs)
586 .copied()
587 .collect();
588 let metadata: Vec<_> = values
589 .iter()
590 .copied()
591 .map(|value| builder.value_metadata(value).cloned())
592 .collect::<Result<_, _>>()?;
593 let symbolic_inputs = synthetic_input_shapes(&metadata);
594 let symbolic_input_refs: Vec<_> = symbolic_inputs.iter().map(Vec::as_slice).collect();
595 let mut context = ShapeGuardContext::default();
596 let mut external_values = HashMap::new();
597 let keys: Vec<_> = values
598 .iter()
599 .copied()
600 .enumerate()
601 .map(|(index, value)| {
602 let key = ValueKey::Input(TensorInputKey::User {
603 id: u64::try_from(index + 1).expect("small semantic AD invocation"),
604 });
605 context.insert_metadata(
606 key.clone(),
607 legacy_metadata(&metadata[index], &symbolic_input_refs),
608 );
609 external_values.insert(key.clone(), value);
610 key
611 })
612 .collect();
613 let input_count = primal_inputs.len();
614 let input_keys = keys[..input_count].to_vec();
615 let output_keys = keys[input_count..].to_vec();
616 let active_values: HashSet<_> = output_keys
617 .iter()
618 .zip(active_outputs)
619 .filter(|(_, active)| **active)
620 .map(|(key, _)| key.clone())
621 .collect();
622 context = context.with_linearize_active_values(Arc::new(active_values));
623 Ok(Self {
624 context,
625 input_keys,
626 output_keys,
627 external_values,
628 shape_sources: values,
629 })
630 }
631}
632
633fn legacy_metadata(metadata: &ProgramValueMetadata, input_shapes: &[&[SymDim]]) -> TensorMeta {
634 let extents = metadata
635 .shape()
636 .iter()
637 .cloned()
638 .map(|extent| extent.map(|dim| SymDim::from_dim_expr(&dim, input_shapes)))
639 .collect();
640 TensorMeta::with_extents(metadata.dtype(), extents)
641}
642
643fn synthetic_input_shapes(metadata: &[ProgramValueMetadata]) -> Vec<Vec<SymDim>> {
644 let mut ranks = Vec::<usize>::new();
645 for expression in metadata
646 .iter()
647 .flat_map(ProgramValueMetadata::shape)
648 .filter_map(ShapeExtent::bound_expr)
649 {
650 collect_input_ranks(expression, &mut ranks);
651 }
652 ranks
653 .into_iter()
654 .enumerate()
655 .map(|(input, rank)| {
656 (0..rank)
657 .map(|axis| {
658 SymDim::tensor_axis(
659 u64::try_from(input + 1).expect("small semantic input index"),
660 axis,
661 )
662 })
663 .collect()
664 })
665 .collect()
666}
667
668fn collect_input_ranks(expression: &DimExpr, ranks: &mut Vec<usize>) {
669 match expression {
670 DimExpr::Const(_) => {}
671 DimExpr::InputDim { input_idx, axis } => {
672 if ranks.len() <= *input_idx {
673 ranks.resize(*input_idx + 1, 0);
674 }
675 ranks[*input_idx] = ranks[*input_idx].max(*axis + 1);
676 }
677 DimExpr::Add(lhs, rhs)
678 | DimExpr::Sub(lhs, rhs)
679 | DimExpr::Mul(lhs, rhs)
680 | DimExpr::FloorDiv(lhs, rhs)
681 | DimExpr::Min(lhs, rhs)
682 | DimExpr::Max(lhs, rhs) => {
683 collect_input_ranks(lhs, ranks);
684 collect_input_ranks(rhs, ranks);
685 }
686 }
687}
688
689#[derive(Clone, Debug)]
690enum SemanticLinearFragmentOp {
691 Core(CoreSemanticOp),
692 Extension(Arc<dyn tenferro_ad::extension::ExtensionOp>),
693 Unsupported(String),
694}
695
696#[derive(Clone)]
697enum SemanticLinearFragmentInput {
698 External(ValueKey<StdTensorOp>),
699 Local(LocalValueId),
700}
701
702impl From<ValueRef<StdTensorOp>> for SemanticLinearFragmentInput {
703 fn from(value: ValueRef<StdTensorOp>) -> Self {
704 match value {
705 ValueRef::External(key) => Self::External(key),
706 ValueRef::Local(local) => Self::Local(local),
707 }
708 }
709}
710
711struct SemanticLinearFragmentOperation {
712 operation: SemanticLinearFragmentOp,
713 inputs: Vec<SemanticLinearFragmentInput>,
714 role: OperationRole,
715 outputs: Vec<LocalValueId>,
716}
717
718fn semantic_linear_fragment_op(operation: &StdTensorOp) -> SemanticLinearFragmentOp {
719 match operation {
720 StdTensorOp::Extension(extension) => {
721 SemanticLinearFragmentOp::Extension(Arc::clone(extension))
722 }
723 core => CoreSemanticOp::try_from(core).map_or_else(
724 |_| SemanticLinearFragmentOp::Unsupported(format!("{core:?}")),
725 SemanticLinearFragmentOp::Core,
726 ),
727 }
728}
729
730struct SemanticRuleBuilder<'a, 'builder> {
731 next_local: usize,
732 locals: Vec<Option<ProgramValue>>,
733 external_values: &'a HashMap<ValueKey<StdTensorOp>, ProgramValue>,
734 shape_sources: &'a [ProgramValue],
735 builder: &'builder mut SemanticProgramBuilder,
736 role: SemanticAdRuleRole,
737 error: Option<SemanticAdError>,
738}
739
740impl<'a, 'builder> SemanticRuleBuilder<'a, 'builder> {
741 fn with_seeds(
742 seeds: &[Option<ProgramValue>],
743 external_values: &'a HashMap<ValueKey<StdTensorOp>, ProgramValue>,
744 shape_sources: &'a [ProgramValue],
745 builder: &'builder mut SemanticProgramBuilder,
746 role: SemanticAdRuleRole,
747 ) -> Self {
748 Self {
749 next_local: seeds.len(),
750 locals: seeds.to_vec(),
751 external_values,
752 shape_sources,
753 builder,
754 role,
755 error: None,
756 }
757 }
758
759 fn finish(self) -> Result<Vec<Option<ProgramValue>>, SemanticAdError> {
760 if let Some(error) = self.error {
761 Err(error)
762 } else {
763 Ok(self.locals)
764 }
765 }
766}
767
768impl PrimitiveRuleBuilder for SemanticRuleBuilder<'_, '_> {
769 fn add_operation(
770 &mut self,
771 operation: StdTensorOp,
772 inputs: Vec<ValueRef<StdTensorOp>>,
773 _role: OperationRole,
774 ) -> Vec<LocalValueId> {
775 let output_count = GraphOperation::output_count(&operation);
776 let outputs: Vec<_> = (self.next_local..self.next_local + output_count).collect();
777 self.next_local += output_count;
778 self.locals.resize(self.next_local, None);
779 if self.error.is_none() {
780 let fragment_op = semantic_linear_fragment_op(&operation);
781 let fragment_inputs: Vec<_> = inputs.into_iter().map(Into::into).collect();
782 let emitted = resolve_semantic_linear_fragment_inputs(
783 &fragment_inputs,
784 self.external_values,
785 &self.locals,
786 self.role,
787 )
788 .and_then(|resolved| {
789 emit_semantic_linear_fragment_operation(
790 &fragment_op,
791 &resolved,
792 self.shape_sources,
793 self.builder,
794 self.role,
795 )
796 });
797 match emitted {
798 Ok(values) => {
799 if values.len() != outputs.len() {
800 self.error = Some(semantic_internal(
801 self.role,
802 format!(
803 "semantic linalg AD operation emitted {} outputs for {} slots",
804 values.len(),
805 outputs.len()
806 ),
807 ));
808 } else {
809 for (local, value) in outputs.iter().copied().zip(values.iter().copied()) {
810 self.locals[local] = Some(value);
811 }
812 }
813 }
814 Err(error) => {
815 self.error = Some(error);
816 }
817 }
818 }
819 outputs
820 }
821}
822
823struct SemanticLinearFragmentBuilder {
829 seed_count: usize,
830 next_local: usize,
831 operations: Vec<SemanticLinearFragmentOperation>,
832}
833
834impl SemanticLinearFragmentBuilder {
835 fn with_seed_count(seed_count: usize) -> Self {
836 Self {
837 seed_count,
838 next_local: seed_count,
839 operations: Vec::new(),
840 }
841 }
842
843 fn transpose_linear_fragment(
844 &self,
845 tangent_outputs: &[Option<LocalValueId>],
846 cotangent_outputs: &[AdValue],
847 active_inputs: &[bool],
848 external_values: &HashMap<ValueKey<StdTensorOp>, ProgramValue>,
849 shape_sources: &[ProgramValue],
850 builder: &mut SemanticProgramBuilder,
851 ) -> Result<Box<[AdValue]>, SemanticAdError> {
852 let role = SemanticAdRuleRole::LinearTranspose;
853 let fixed_locals =
854 self.emit_fixed_primal_ops(external_values, shape_sources, builder, role)?;
855 let mut cotangents = HashMap::<LocalValueId, ProgramValue>::new();
856 for (tangent, cotangent) in tangent_outputs
857 .iter()
858 .copied()
859 .zip(cotangent_outputs.iter().copied())
860 {
861 if let (Some(tangent), AdValue::Value(cotangent)) = (tangent, cotangent) {
862 accumulate_local_cotangent(builder, &mut cotangents, tangent, cotangent)?;
863 }
864 }
865 for operation in self.operations.iter().rev() {
866 let Some(active_mask) = linear_active_mask(&operation.role) else {
867 continue;
868 };
869 if !active_mask.iter().any(|active| *active) {
870 continue;
871 }
872 let output_cotangents: Vec<_> = operation
873 .outputs
874 .iter()
875 .map(|output| cotangents.remove(output))
876 .collect();
877 if output_cotangents.iter().all(Option::is_none) {
878 continue;
879 }
880 let context = SemanticLinearFragmentTransposeContext {
881 fragment: self,
882 external_values,
883 fixed_locals: &fixed_locals,
884 shape_sources,
885 role,
886 };
887 let input_cotangents = transpose_semantic_linear_fragment_operation(
888 operation,
889 &output_cotangents,
890 active_mask,
891 &context,
892 builder,
893 )?;
894 for ((input, active), cotangent) in operation
895 .inputs
896 .iter()
897 .zip(active_mask)
898 .zip(input_cotangents)
899 {
900 if !active {
901 continue;
902 }
903 let (SemanticLinearFragmentInput::Local(input), Some(cotangent)) =
904 (input, cotangent)
905 else {
906 return Err(semantic_internal(
907 role,
908 "linear linalg fragment has a non-local active input",
909 ));
910 };
911 accumulate_local_cotangent(builder, &mut cotangents, *input, cotangent)?;
912 }
913 }
914 Ok(active_inputs
915 .iter()
916 .copied()
917 .enumerate()
918 .map(|(input, active)| {
919 if active {
920 cotangents
921 .remove(&input)
922 .map_or(AdValue::Absent, AdValue::Value)
923 } else {
924 AdValue::Absent
925 }
926 })
927 .collect())
928 }
929
930 fn emit_fixed_primal_ops(
931 &self,
932 external_values: &HashMap<ValueKey<StdTensorOp>, ProgramValue>,
933 shape_sources: &[ProgramValue],
934 builder: &mut SemanticProgramBuilder,
935 role: SemanticAdRuleRole,
936 ) -> Result<Vec<Option<ProgramValue>>, SemanticAdError> {
937 let mut locals = vec![None; self.next_local];
938 for operation in &self.operations {
939 if linear_active_mask(&operation.role)
940 .is_some_and(|mask| mask.iter().any(|active| *active))
941 {
942 continue;
943 }
944 let inputs = resolve_semantic_linear_fragment_inputs(
945 &operation.inputs,
946 external_values,
947 &locals,
948 role,
949 )?;
950 let outputs = emit_semantic_linear_fragment_operation(
951 &operation.operation,
952 &inputs,
953 shape_sources,
954 builder,
955 role,
956 )?;
957 for (local, value) in operation
958 .outputs
959 .iter()
960 .copied()
961 .zip(outputs.iter().copied())
962 {
963 locals[local] = Some(value);
964 }
965 }
966 Ok(locals)
967 }
968}
969
970impl PrimitiveRuleBuilder for SemanticLinearFragmentBuilder {
971 fn add_operation(
972 &mut self,
973 operation: StdTensorOp,
974 inputs: Vec<ValueRef<StdTensorOp>>,
975 role: OperationRole,
976 ) -> Vec<LocalValueId> {
977 let output_count = GraphOperation::output_count(&operation);
978 let outputs: Vec<_> = (self.next_local..self.next_local + output_count).collect();
979 self.next_local += output_count;
980 self.operations.push(SemanticLinearFragmentOperation {
981 operation: semantic_linear_fragment_op(&operation),
982 inputs: inputs.into_iter().map(Into::into).collect(),
983 role,
984 outputs: outputs.clone(),
985 });
986 outputs
987 }
988}
989
990fn linear_active_mask(role: &OperationRole) -> Option<&[bool]> {
991 match role {
992 OperationRole::Primary => None,
993 OperationRole::Linearized { active_mask } => Some(active_mask),
994 }
995}
996
997struct SemanticLinearFragmentTransposeContext<'a> {
998 fragment: &'a SemanticLinearFragmentBuilder,
999 external_values: &'a HashMap<ValueKey<StdTensorOp>, ProgramValue>,
1000 fixed_locals: &'a [Option<ProgramValue>],
1001 shape_sources: &'a [ProgramValue],
1002 role: SemanticAdRuleRole,
1003}
1004
1005fn resolve_semantic_linear_fragment_inputs(
1006 inputs: &[SemanticLinearFragmentInput],
1007 external_values: &HashMap<ValueKey<StdTensorOp>, ProgramValue>,
1008 locals: &[Option<ProgramValue>],
1009 role: SemanticAdRuleRole,
1010) -> Result<Vec<ProgramValue>, SemanticAdError> {
1011 inputs
1012 .iter()
1013 .map(|input| match input {
1014 SemanticLinearFragmentInput::External(key) => external_values.get(key).copied(),
1015 SemanticLinearFragmentInput::Local(local) => locals.get(*local).copied().flatten(),
1016 })
1017 .collect::<Option<_>>()
1018 .ok_or_else(|| {
1019 semantic_internal(
1020 role,
1021 "semantic linalg linear fragment references an unavailable fixed value",
1022 )
1023 })
1024}
1025
1026fn emit_semantic_linear_fragment_operation(
1027 operation: &SemanticLinearFragmentOp,
1028 inputs: &[ProgramValue],
1029 shape_sources: &[ProgramValue],
1030 builder: &mut SemanticProgramBuilder,
1031 role: SemanticAdRuleRole,
1032) -> Result<Box<[ProgramValue]>, SemanticAdError> {
1033 match operation {
1034 SemanticLinearFragmentOp::Extension(extension) => {
1035 Ok(builder.add_extension(Arc::clone(extension), inputs)?)
1036 }
1037 SemanticLinearFragmentOp::Core(core) => {
1038 let fragment_core = core.clone();
1039 let (core, inputs) =
1040 localize_shape_expressions(core.clone(), inputs, shape_sources, builder, role)
1041 .map_err(|error| match error {
1042 SemanticAdError::Invariant {
1043 family_id,
1044 role,
1045 message,
1046 } => SemanticAdError::Invariant {
1047 family_id,
1048 role,
1049 message: format!(
1050 "{message}; linear fragment operation {fragment_core:?}"
1051 ),
1052 },
1053 other => other,
1054 })?;
1055 Ok(builder.add_op(core, &inputs)?)
1056 }
1057 SemanticLinearFragmentOp::Unsupported(operation) => Err(semantic_internal(
1058 role,
1059 format!("linalg AD emitted a non-semantic standard operation {operation}"),
1060 )),
1061 }
1062}
1063
1064fn localize_shape_expressions(
1065 operation: CoreSemanticOp,
1066 data_inputs: &[ProgramValue],
1067 shape_sources: &[ProgramValue],
1068 builder: &SemanticProgramBuilder,
1069 role: SemanticAdRuleRole,
1070) -> Result<(CoreSemanticOp, Vec<ProgramValue>), SemanticAdError> {
1071 let mut inputs = data_inputs.to_vec();
1072 let operation = match operation {
1073 CoreSemanticOp::Reshape { to_shape } => CoreSemanticOp::Reshape {
1074 to_shape: localize_dims(
1075 &to_shape,
1076 data_inputs,
1077 shape_sources,
1078 1,
1079 &mut inputs,
1080 builder,
1081 role,
1082 )?,
1083 },
1084 CoreSemanticOp::BroadcastInDim { shape, dims } => CoreSemanticOp::BroadcastInDim {
1085 shape: localize_dims(
1086 &shape,
1087 data_inputs,
1088 shape_sources,
1089 1,
1090 &mut inputs,
1091 builder,
1092 role,
1093 )?,
1094 dims,
1095 },
1096 CoreSemanticOp::GatherDynamicSliceSizes {
1097 offset_dims,
1098 collapsed_slice_dims,
1099 start_index_map,
1100 index_vector_dim,
1101 slice_sizes,
1102 } => CoreSemanticOp::GatherDynamicSliceSizes {
1103 offset_dims,
1104 collapsed_slice_dims,
1105 start_index_map,
1106 index_vector_dim,
1107 slice_sizes: localize_dims(
1108 &slice_sizes,
1109 data_inputs,
1110 shape_sources,
1111 2,
1112 &mut inputs,
1113 builder,
1114 role,
1115 )?,
1116 },
1117 other => other,
1118 };
1119 Ok((operation, inputs))
1120}
1121
1122fn localize_dims(
1123 dims: &[DimExpr],
1124 data_inputs: &[ProgramValue],
1125 shape_sources: &[ProgramValue],
1126 fixed_data_arity: usize,
1127 operation_inputs: &mut Vec<ProgramValue>,
1128 builder: &SemanticProgramBuilder,
1129 role: SemanticAdRuleRole,
1130) -> Result<Vec<DimExpr>, SemanticAdError> {
1131 dims.iter()
1132 .map(|dim| {
1133 localize_dim(
1134 dim,
1135 data_inputs,
1136 shape_sources,
1137 fixed_data_arity,
1138 operation_inputs,
1139 builder,
1140 role,
1141 )
1142 })
1143 .collect()
1144}
1145
1146fn localize_dim(
1147 dim: &DimExpr,
1148 data_inputs: &[ProgramValue],
1149 shape_sources: &[ProgramValue],
1150 fixed_data_arity: usize,
1151 operation_inputs: &mut Vec<ProgramValue>,
1152 builder: &SemanticProgramBuilder,
1153 role: SemanticAdRuleRole,
1154) -> Result<DimExpr, SemanticAdError> {
1155 let binary = |lhs: &DimExpr,
1156 rhs: &DimExpr,
1157 constructor: fn(Box<DimExpr>, Box<DimExpr>) -> DimExpr,
1158 operation_inputs: &mut Vec<ProgramValue>|
1159 -> Result<DimExpr, SemanticAdError> {
1160 Ok(constructor(
1161 Box::new(localize_dim(
1162 lhs,
1163 data_inputs,
1164 shape_sources,
1165 fixed_data_arity,
1166 operation_inputs,
1167 builder,
1168 role,
1169 )?),
1170 Box::new(localize_dim(
1171 rhs,
1172 data_inputs,
1173 shape_sources,
1174 fixed_data_arity,
1175 operation_inputs,
1176 builder,
1177 role,
1178 )?),
1179 ))
1180 };
1181 match dim {
1182 DimExpr::Const(value) => Ok(DimExpr::Const(*value)),
1183 DimExpr::InputDim { input_idx, axis } => {
1184 let source = if *input_idx >= fixed_data_arity {
1191 data_inputs
1192 .get(*input_idx)
1193 .copied()
1194 .or_else(|| shape_sources.get(*input_idx).copied())
1195 } else {
1196 shape_sources.get(*input_idx).copied()
1197 }
1198 .ok_or_else(|| {
1199 semantic_internal(
1200 role,
1201 format!(
1202 "linalg AD symbolic shape input {input_idx} is out of bounds for {} operation inputs and {} primal shape sources",
1203 data_inputs.len(),
1204 shape_sources.len()
1205 ),
1206 )
1207 })?;
1208 let rank = builder.value_metadata(source)?.shape().len();
1209 if *axis >= rank {
1210 return Err(semantic_internal(
1211 role,
1212 format!(
1213 "linalg AD symbolic shape axis {axis} is out of bounds for source rank {rank}"
1214 ),
1215 ));
1216 }
1217 let input_idx = operation_inputs
1218 .iter()
1219 .position(|value| *value == source)
1220 .unwrap_or_else(|| {
1221 operation_inputs.push(source);
1222 operation_inputs.len() - 1
1223 });
1224 debug_assert!(
1225 input_idx < data_inputs.len() + shape_sources.len(),
1226 "localized shape source must be an operation input"
1227 );
1228 Ok(DimExpr::InputDim {
1229 input_idx,
1230 axis: *axis,
1231 })
1232 }
1233 DimExpr::Add(lhs, rhs) => binary(lhs, rhs, DimExpr::Add, operation_inputs),
1234 DimExpr::Sub(lhs, rhs) => binary(lhs, rhs, DimExpr::Sub, operation_inputs),
1235 DimExpr::Mul(lhs, rhs) => binary(lhs, rhs, DimExpr::Mul, operation_inputs),
1236 DimExpr::FloorDiv(lhs, rhs) => binary(lhs, rhs, DimExpr::FloorDiv, operation_inputs),
1237 DimExpr::Min(lhs, rhs) => binary(lhs, rhs, DimExpr::Min, operation_inputs),
1238 DimExpr::Max(lhs, rhs) => binary(lhs, rhs, DimExpr::Max, operation_inputs),
1239 }
1240}
1241
1242fn transpose_semantic_linear_fragment_operation(
1243 operation: &SemanticLinearFragmentOperation,
1244 cotangent_outputs: &[Option<ProgramValue>],
1245 active_mask: &[bool],
1246 context: &SemanticLinearFragmentTransposeContext<'_>,
1247 builder: &mut SemanticProgramBuilder,
1248) -> Result<Vec<Option<ProgramValue>>, SemanticAdError> {
1249 let Some(cotangent) = cotangent_outputs.first().copied().flatten() else {
1250 return Ok(vec![None; operation.inputs.len()]);
1251 };
1252 let fixed = |index: usize| {
1253 if active_mask.get(index).copied().unwrap_or(false) {
1254 None
1255 } else {
1256 match operation.inputs.get(index) {
1257 Some(SemanticLinearFragmentInput::External(key)) => {
1258 context.external_values.get(key).copied()
1259 }
1260 Some(SemanticLinearFragmentInput::Local(local)) => {
1261 context.fixed_locals.get(*local).copied().flatten()
1262 }
1263 None => None,
1264 }
1265 }
1266 };
1267 let unary = |value| Ok(vec![Some(value)]);
1268 match &operation.operation {
1269 SemanticLinearFragmentOp::Core(CoreSemanticOp::Add) => Ok(active_mask
1270 .iter()
1271 .map(|active| active.then_some(cotangent))
1272 .collect()),
1273 SemanticLinearFragmentOp::Core(CoreSemanticOp::Sub) => {
1274 let rhs = builder.add_op(CoreSemanticOp::Neg, &[cotangent])?[0];
1275 Ok(vec![
1276 active_mask[0].then_some(cotangent),
1277 active_mask[1].then_some(rhs),
1278 ])
1279 }
1280 SemanticLinearFragmentOp::Core(CoreSemanticOp::Neg) => {
1281 let value = builder.add_op(CoreSemanticOp::Neg, &[cotangent])?[0];
1282 unary(value)
1283 }
1284 SemanticLinearFragmentOp::Core(CoreSemanticOp::Conj) => {
1285 let value = builder.add_op(CoreSemanticOp::Conj, &[cotangent])?[0];
1286 unary(value)
1287 }
1288 SemanticLinearFragmentOp::Core(CoreSemanticOp::Mul) => {
1289 transpose_mul(cotangent, active_mask, &fixed, builder, context.role)
1290 }
1291 SemanticLinearFragmentOp::Core(CoreSemanticOp::Div) => {
1292 transpose_div(cotangent, active_mask, &fixed, builder, context.role)
1293 }
1294 SemanticLinearFragmentOp::Core(CoreSemanticOp::DotGeneral { config }) => {
1295 transpose_matrix_dot(
1296 cotangent,
1297 config,
1298 active_mask,
1299 &fixed,
1300 builder,
1301 context.role,
1302 )
1303 }
1304 SemanticLinearFragmentOp::Core(CoreSemanticOp::ReduceSum { axes }) => {
1305 transpose_reduce_sum(context, operation, cotangent, axes, active_mask, builder)
1306 }
1307 SemanticLinearFragmentOp::Core(CoreSemanticOp::Transpose { perm }) => {
1308 let mut inverse = vec![0; perm.len()];
1309 for (output_axis, input_axis) in perm.iter().copied().enumerate() {
1310 inverse[input_axis] = output_axis;
1311 }
1312 let value =
1313 builder.add_op(CoreSemanticOp::Transpose { perm: inverse }, &[cotangent])?[0];
1314 unary(value)
1315 }
1316 SemanticLinearFragmentOp::Core(CoreSemanticOp::Convert { from, to }) => {
1317 let value = builder.add_op(
1318 CoreSemanticOp::Convert {
1319 from: *to,
1320 to: *from,
1321 },
1322 &[cotangent],
1323 )?[0];
1324 unary(value)
1325 }
1326 SemanticLinearFragmentOp::Core(CoreSemanticOp::ExtractDiag { axis_a, axis_b }) => {
1327 let value = builder.add_op(
1328 CoreSemanticOp::EmbedDiag {
1329 axis_a: *axis_a,
1330 axis_b: *axis_b,
1331 },
1332 &[cotangent],
1333 )?[0];
1334 unary(value)
1335 }
1336 SemanticLinearFragmentOp::Core(CoreSemanticOp::EmbedDiag { axis_a, axis_b }) => {
1337 let value = builder.add_op(
1338 CoreSemanticOp::ExtractDiag {
1339 axis_a: *axis_a,
1340 axis_b: *axis_b,
1341 },
1342 &[cotangent],
1343 )?[0];
1344 unary(value)
1345 }
1346 SemanticLinearFragmentOp::Core(CoreSemanticOp::Tril { k }) => {
1347 let value = builder.add_op(CoreSemanticOp::Tril { k: *k }, &[cotangent])?[0];
1348 unary(value)
1349 }
1350 SemanticLinearFragmentOp::Core(CoreSemanticOp::Triu { k }) => {
1351 let value = builder.add_op(CoreSemanticOp::Triu { k: *k }, &[cotangent])?[0];
1352 unary(value)
1353 }
1354 SemanticLinearFragmentOp::Extension(extension) => transpose_linalg_extension(
1355 extension.as_ref(),
1356 operation,
1357 cotangent,
1358 active_mask,
1359 context.external_values,
1360 context.fixed_locals,
1361 builder,
1362 ),
1363 SemanticLinearFragmentOp::Unsupported(operation) => Err(semantic_internal(
1364 context.role,
1365 format!("unsupported linear linalg fragment operation {operation}"),
1366 )),
1367 other => Err(semantic_internal(
1368 context.role,
1369 format!("unsupported linear linalg fragment operation {other:?}"),
1370 )),
1371 }
1372}
1373
1374fn transpose_reduce_sum(
1375 context: &SemanticLinearFragmentTransposeContext<'_>,
1376 operation: &SemanticLinearFragmentOperation,
1377 cotangent: ProgramValue,
1378 axes: &[usize],
1379 active_mask: &[bool],
1380 builder: &mut SemanticProgramBuilder,
1381) -> Result<Vec<Option<ProgramValue>>, SemanticAdError> {
1382 if operation.inputs.len() != 1 || active_mask.len() != 1 {
1383 return Err(semantic_internal(
1384 context.role,
1385 "linear reduce_sum fragment has malformed arity",
1386 ));
1387 }
1388 if !active_mask[0] {
1389 return Ok(vec![None]);
1390 }
1391 let mut cache = HashMap::new();
1392 let input_shape = semantic_linear_fragment_value_shape(
1393 context.fragment,
1394 &operation.inputs[0],
1395 context.external_values,
1396 context.shape_sources,
1397 builder,
1398 context.role,
1399 &mut cache,
1400 )?
1401 .ok_or_else(|| {
1402 semantic_internal(
1403 context.role,
1404 "linear reduce_sum fragment is missing input shape metadata",
1405 )
1406 })?;
1407 if axes.iter().any(|axis| *axis >= input_shape.len()) {
1408 return Err(semantic_internal(
1409 context.role,
1410 format!(
1411 "linear reduce_sum axis is out of bounds for input rank {}",
1412 input_shape.len()
1413 ),
1414 ));
1415 }
1416 let dims: Vec<_> = (0..input_shape.len())
1417 .filter(|axis| !axes.contains(axis))
1418 .collect();
1419 let mut inputs = vec![cotangent];
1420 let shape = localize_dims(
1421 &input_shape,
1422 &[cotangent],
1423 context.shape_sources,
1424 1,
1425 &mut inputs,
1426 builder,
1427 context.role,
1428 )?;
1429 Ok(vec![Some(
1430 builder.add_op(CoreSemanticOp::BroadcastInDim { shape, dims }, &inputs)?[0],
1431 )])
1432}
1433
1434fn semantic_linear_fragment_value_shape(
1435 fragment: &SemanticLinearFragmentBuilder,
1436 value: &SemanticLinearFragmentInput,
1437 external_values: &HashMap<ValueKey<StdTensorOp>, ProgramValue>,
1438 shape_sources: &[ProgramValue],
1439 builder: &SemanticProgramBuilder,
1440 role: SemanticAdRuleRole,
1441 cache: &mut HashMap<LocalValueId, Option<Vec<DimExpr>>>,
1442) -> Result<Option<Vec<DimExpr>>, SemanticAdError> {
1443 match value {
1444 SemanticLinearFragmentInput::External(key) => {
1445 let source = external_values.get(key).copied().ok_or_else(|| {
1446 semantic_internal(
1447 role,
1448 "semantic linalg linear-fragment shape references missing external value",
1449 )
1450 })?;
1451 source_shape(source, shape_sources, builder, role).map(Some)
1452 }
1453 SemanticLinearFragmentInput::Local(local) => semantic_linear_fragment_local_shape(
1454 fragment,
1455 *local,
1456 external_values,
1457 shape_sources,
1458 builder,
1459 role,
1460 cache,
1461 ),
1462 }
1463}
1464
1465fn semantic_linear_fragment_local_shape(
1466 fragment: &SemanticLinearFragmentBuilder,
1467 local: LocalValueId,
1468 external_values: &HashMap<ValueKey<StdTensorOp>, ProgramValue>,
1469 shape_sources: &[ProgramValue],
1470 builder: &SemanticProgramBuilder,
1471 role: SemanticAdRuleRole,
1472 cache: &mut HashMap<LocalValueId, Option<Vec<DimExpr>>>,
1473) -> Result<Option<Vec<DimExpr>>, SemanticAdError> {
1474 if let Some(cached) = cache.get(&local) {
1475 return Ok(cached.clone());
1476 }
1477 let shape = if local < fragment.seed_count {
1478 let source = shape_sources.get(local).copied().ok_or_else(|| {
1479 semantic_internal(
1480 role,
1481 format!("semantic linalg linear-fragment seed local {local} has no shape source"),
1482 )
1483 })?;
1484 Some(source_shape(source, shape_sources, builder, role)?)
1485 } else {
1486 let (operation, output_index) = fragment
1487 .operations
1488 .iter()
1489 .find_map(|operation| {
1490 operation
1491 .outputs
1492 .iter()
1493 .position(|output| *output == local)
1494 .map(|index| (operation, index))
1495 })
1496 .ok_or_else(|| {
1497 semantic_internal(
1498 role,
1499 format!(
1500 "semantic linalg linear-fragment local {local} has no producing operation"
1501 ),
1502 )
1503 })?;
1504 semantic_linear_fragment_operation_output_shape(
1505 fragment,
1506 operation,
1507 output_index,
1508 external_values,
1509 shape_sources,
1510 builder,
1511 role,
1512 cache,
1513 )?
1514 };
1515 cache.insert(local, shape.clone());
1516 Ok(shape)
1517}
1518
1519fn source_shape(
1520 source: ProgramValue,
1521 shape_sources: &[ProgramValue],
1522 builder: &SemanticProgramBuilder,
1523 role: SemanticAdRuleRole,
1524) -> Result<Vec<DimExpr>, SemanticAdError> {
1525 let index = shape_sources
1526 .iter()
1527 .position(|candidate| *candidate == source)
1528 .ok_or_else(|| {
1529 semantic_internal(role, "shape source is not part of the linalg invocation")
1530 })?;
1531 let rank = builder.value_metadata(source)?.shape().len();
1532 Ok(DimExpr::input_shape(index, rank))
1533}
1534
1535#[allow(clippy::too_many_arguments)]
1536fn semantic_linear_fragment_operation_output_shape(
1537 fragment: &SemanticLinearFragmentBuilder,
1538 operation: &SemanticLinearFragmentOperation,
1539 output_index: usize,
1540 external_values: &HashMap<ValueKey<StdTensorOp>, ProgramValue>,
1541 shape_sources: &[ProgramValue],
1542 builder: &SemanticProgramBuilder,
1543 role: SemanticAdRuleRole,
1544 cache: &mut HashMap<LocalValueId, Option<Vec<DimExpr>>>,
1545) -> Result<Option<Vec<DimExpr>>, SemanticAdError> {
1546 let input_shape = |input_index: usize,
1547 cache: &mut HashMap<LocalValueId, Option<Vec<DimExpr>>>|
1548 -> Result<Option<Vec<DimExpr>>, SemanticAdError> {
1549 let input = operation.inputs.get(input_index).ok_or_else(|| {
1550 semantic_internal(
1551 role,
1552 "semantic linalg linear-fragment shape requested missing operation input",
1553 )
1554 })?;
1555 semantic_linear_fragment_value_shape(
1556 fragment,
1557 input,
1558 external_values,
1559 shape_sources,
1560 builder,
1561 role,
1562 cache,
1563 )
1564 };
1565 match &operation.operation {
1566 SemanticLinearFragmentOp::Extension(extension) => {
1567 let linalg = semantic_linalg_op(extension.as_ref(), role)?;
1568 match linalg.op() {
1569 LinalgOp::LuFactor => match output_index {
1570 0 => input_shape(0, cache),
1571 1 => Ok(input_shape(0, cache)?.map(|shape| {
1572 let (rows, cols, batch) =
1573 semantic_linear_fragment_matrix_shape_parts(&shape);
1574 let mut pivots_shape =
1575 vec![DimExpr::Min(Box::new(rows.clone()), Box::new(cols.clone()))];
1576 pivots_shape.extend_from_slice(batch);
1577 pivots_shape
1578 })),
1579 2 => Ok(input_shape(0, cache)?.map(|shape| shape[2..].to_vec())),
1580 _ => Ok(None),
1581 },
1582 LinalgOp::LuSolvePrepared { .. } => input_shape(3, cache),
1583 LinalgOp::Eigh { .. } => match output_index {
1589 0 => Ok(input_shape(0, cache)?.map(|shape| {
1590 shape
1591 .into_iter()
1592 .enumerate()
1593 .filter_map(|(axis, dim)| (axis != 1).then_some(dim))
1594 .collect()
1595 })),
1596 1 => input_shape(0, cache),
1597 _ => Ok(None),
1598 },
1599 _ => Ok(None),
1600 }
1601 }
1602 SemanticLinearFragmentOp::Core(CoreSemanticOp::ExtractDiag { axis_a, axis_b }) => {
1603 Ok(input_shape(0, cache)?
1604 .map(|shape| extract_diag_shape(&shape, *axis_a, *axis_b))
1605 .transpose()?)
1606 }
1607 SemanticLinearFragmentOp::Core(CoreSemanticOp::ReduceSum { axes }) => {
1608 Ok(input_shape(0, cache)?.map(|shape| {
1609 shape
1610 .into_iter()
1611 .enumerate()
1612 .filter_map(|(axis, dim)| (!axes.contains(&axis)).then_some(dim))
1613 .collect()
1614 }))
1615 }
1616 SemanticLinearFragmentOp::Core(
1617 CoreSemanticOp::Convert { .. }
1618 | CoreSemanticOp::Neg
1619 | CoreSemanticOp::Conj
1620 | CoreSemanticOp::Mul
1624 | CoreSemanticOp::Tril { .. }
1625 | CoreSemanticOp::Triu { .. },
1626 ) => input_shape(0, cache),
1627 SemanticLinearFragmentOp::Core(CoreSemanticOp::Transpose { perm }) => {
1628 Ok(input_shape(0, cache)?
1629 .map(|shape| perm.iter().map(|axis| shape[*axis].clone()).collect()))
1630 }
1631 _ => Ok(None),
1632 }
1633}
1634
1635fn semantic_linear_fragment_matrix_shape_parts(
1636 shape: &[DimExpr],
1637) -> (&DimExpr, &DimExpr, &[DimExpr]) {
1638 (&shape[0], &shape[1], &shape[2..])
1639}
1640
1641fn extract_diag_shape(
1642 shape: &[DimExpr],
1643 axis_a: usize,
1644 axis_b: usize,
1645) -> Result<Vec<DimExpr>, SemanticAdError> {
1646 if axis_a >= shape.len() || axis_b >= shape.len() || axis_a == axis_b {
1647 return Err(semantic_internal(
1648 SemanticAdRuleRole::LinearTranspose,
1649 "extract_diag shape derivation received invalid axes",
1650 ));
1651 }
1652 let diagonal = DimExpr::Min(
1653 Box::new(shape[axis_a].clone()),
1654 Box::new(shape[axis_b].clone()),
1655 );
1656 let mut output = Vec::with_capacity(shape.len() - 1);
1657 for (axis, dim) in shape.iter().enumerate() {
1658 if axis == axis_b {
1659 continue;
1660 }
1661 if axis == axis_a {
1662 output.push(diagonal.clone());
1663 } else {
1664 output.push(dim.clone());
1665 }
1666 }
1667 Ok(output)
1668}
1669
1670fn transpose_mul(
1671 cotangent: ProgramValue,
1672 active_mask: &[bool],
1673 fixed: &impl Fn(usize) -> Option<ProgramValue>,
1674 builder: &mut SemanticProgramBuilder,
1675 role: SemanticAdRuleRole,
1676) -> Result<Vec<Option<ProgramValue>>, SemanticAdError> {
1677 let mut result = vec![None; 2];
1678 for input in 0..2 {
1679 if !active_mask[input] {
1680 continue;
1681 }
1682 let coefficient = fixed(1 - input).ok_or_else(|| {
1683 semantic_internal(role, "linear multiply is missing its fixed coefficient")
1684 })?;
1685 let coefficient = conjugate_if_complex(builder, coefficient)?;
1686 result[input] = Some(builder.add_op(CoreSemanticOp::Mul, &[cotangent, coefficient])?[0]);
1687 }
1688 Ok(result)
1689}
1690
1691fn transpose_div(
1692 cotangent: ProgramValue,
1693 active_mask: &[bool],
1694 fixed: &impl Fn(usize) -> Option<ProgramValue>,
1695 builder: &mut SemanticProgramBuilder,
1696 role: SemanticAdRuleRole,
1697) -> Result<Vec<Option<ProgramValue>>, SemanticAdError> {
1698 let mut result = vec![None; 2];
1699 if active_mask[0] {
1700 let denominator = fixed(1).ok_or_else(|| {
1701 semantic_internal(role, "linear divide is missing its fixed denominator")
1702 })?;
1703 let denominator = conjugate_if_complex(builder, denominator)?;
1704 result[0] = Some(builder.add_op(CoreSemanticOp::Div, &[cotangent, denominator])?[0]);
1705 }
1706 if active_mask[1] {
1707 let numerator = fixed(0).ok_or_else(|| {
1708 semantic_internal(role, "linear divide is missing its fixed numerator")
1709 })?;
1710 let denominator = fixed(1).ok_or_else(|| {
1711 semantic_internal(role, "linear divide is missing its fixed denominator")
1712 })?;
1713 let square = builder.add_op(CoreSemanticOp::Mul, &[denominator, denominator])?[0];
1714 let coefficient = builder.add_op(CoreSemanticOp::Div, &[numerator, square])?[0];
1715 let coefficient = conjugate_if_complex(builder, coefficient)?;
1716 let value = builder.add_op(CoreSemanticOp::Mul, &[cotangent, coefficient])?[0];
1717 result[1] = Some(builder.add_op(CoreSemanticOp::Neg, &[value])?[0]);
1718 }
1719 Ok(result)
1720}
1721
1722fn transpose_matrix_dot(
1723 cotangent: ProgramValue,
1724 config: &tenferro_tensor::DotGeneralConfig,
1725 active_mask: &[bool],
1726 fixed: &impl Fn(usize) -> Option<ProgramValue>,
1727 builder: &mut SemanticProgramBuilder,
1728 role: SemanticAdRuleRole,
1729) -> Result<Vec<Option<ProgramValue>>, SemanticAdError> {
1730 let rank = 2 + config.lhs_batch_dims.len();
1731 let expected_batch: Vec<_> = (2..rank).collect();
1732 if config.lhs_contracting_dims.as_slice() != [1]
1733 || config.rhs_contracting_dims.as_slice() != [0]
1734 || config.lhs_batch_dims.as_slice() != expected_batch.as_slice()
1735 || config.rhs_batch_dims.as_slice() != expected_batch.as_slice()
1736 {
1737 return Err(semantic_internal(
1738 role,
1739 "linalg AD emitted an unsupported dot-general configuration",
1740 ));
1741 }
1742 let mut result = vec![None; 2];
1743 if active_mask[0] {
1744 let rhs = fixed(1).ok_or_else(|| {
1745 semantic_internal(role, "linear matrix product is missing its fixed rhs")
1746 })?;
1747 let rhs_h = matrix_adjoint(builder, rhs, rank)?;
1748 result[0] = Some(
1749 builder.add_op(
1750 CoreSemanticOp::DotGeneral {
1751 config: config.clone(),
1752 },
1753 &[cotangent, rhs_h],
1754 )?[0],
1755 );
1756 }
1757 if active_mask[1] {
1758 let lhs = fixed(0).ok_or_else(|| {
1759 semantic_internal(role, "linear matrix product is missing its fixed lhs")
1760 })?;
1761 let lhs_h = matrix_adjoint(builder, lhs, rank)?;
1762 result[1] = Some(
1763 builder.add_op(
1764 CoreSemanticOp::DotGeneral {
1765 config: config.clone(),
1766 },
1767 &[lhs_h, cotangent],
1768 )?[0],
1769 );
1770 }
1771 Ok(result)
1772}
1773
1774fn semantic_solve_matrix_cotangent(
1775 builder: &mut SemanticProgramBuilder,
1776 rhs_cotangent: ProgramValue,
1777 solution: ProgramValue,
1778 left_side: bool,
1779 transpose_a: bool,
1780 rank: usize,
1781) -> Result<ProgramValue, SemanticAdError> {
1782 let negative_rhs_cotangent = builder.add_op(CoreSemanticOp::Neg, &[rhs_cotangent])?[0];
1783 let solution_h = matrix_adjoint(builder, solution, rank)?;
1784 let config = semantic_matrix_multiply_config(rank)?;
1785 let matrix_cotangent = if left_side {
1786 builder.add_op(
1787 CoreSemanticOp::DotGeneral {
1788 config: config.clone(),
1789 },
1790 &[negative_rhs_cotangent, solution_h],
1791 )?[0]
1792 } else {
1793 builder.add_op(
1794 CoreSemanticOp::DotGeneral { config },
1795 &[solution_h, negative_rhs_cotangent],
1796 )?[0]
1797 };
1798 if transpose_a {
1799 semantic_matrix_transpose(builder, matrix_cotangent, rank)
1800 } else {
1801 Ok(matrix_cotangent)
1802 }
1803}
1804
1805fn semantic_matrix_multiply_config(
1806 rank: usize,
1807) -> Result<tenferro_tensor::DotGeneralConfig, SemanticAdError> {
1808 if rank < 2 {
1809 return Err(semantic_internal(
1810 SemanticAdRuleRole::LinearTranspose,
1811 "matrix multiply semantic helper expects rank >= 2",
1812 ));
1813 }
1814 let batch_dims: Vec<usize> = (2..rank).collect();
1815 Ok(tenferro_tensor::DotGeneralConfig {
1816 lhs_contracting_dims: [1].as_slice().into(),
1817 rhs_contracting_dims: [0].as_slice().into(),
1818 lhs_batch_dims: batch_dims.clone().into(),
1819 rhs_batch_dims: batch_dims.into(),
1820 })
1821}
1822
1823fn semantic_matrix_transpose(
1824 builder: &mut SemanticProgramBuilder,
1825 value: ProgramValue,
1826 rank: usize,
1827) -> Result<ProgramValue, SemanticAdError> {
1828 if rank < 2 {
1829 return Err(semantic_internal(
1830 SemanticAdRuleRole::LinearTranspose,
1831 "matrix transpose semantic helper expects rank >= 2",
1832 ));
1833 }
1834 let mut perm: Vec<_> = (0..rank).collect();
1835 perm.swap(0, 1);
1836 Ok(builder.add_op(CoreSemanticOp::Transpose { perm }, &[value])?[0])
1837}
1838
1839fn matrix_adjoint(
1840 builder: &mut SemanticProgramBuilder,
1841 value: ProgramValue,
1842 rank: usize,
1843) -> Result<ProgramValue, SemanticAdError> {
1844 let value = conjugate_if_complex(builder, value)?;
1845 let mut perm: Vec<_> = (0..rank).collect();
1846 perm.swap(0, 1);
1847 Ok(builder.add_op(CoreSemanticOp::Transpose { perm }, &[value])?[0])
1848}
1849
1850fn conjugate_if_complex(
1851 builder: &mut SemanticProgramBuilder,
1852 value: ProgramValue,
1853) -> Result<ProgramValue, SemanticAdError> {
1854 if matches!(
1855 builder.value_metadata(value)?.dtype(),
1856 tenferro_tensor::DType::C32 | tenferro_tensor::DType::C64
1857 ) {
1858 Ok(builder.add_op(CoreSemanticOp::Conj, &[value])?[0])
1859 } else {
1860 Ok(value)
1861 }
1862}
1863
1864fn transpose_linalg_extension(
1865 extension: &dyn tenferro_ad::extension::ExtensionOp,
1866 operation: &SemanticLinearFragmentOperation,
1867 cotangent: ProgramValue,
1868 active_mask: &[bool],
1869 external_values: &HashMap<ValueKey<StdTensorOp>, ProgramValue>,
1870 fixed_locals: &[Option<ProgramValue>],
1871 builder: &mut SemanticProgramBuilder,
1872) -> Result<Vec<Option<ProgramValue>>, SemanticAdError> {
1873 let role = SemanticAdRuleRole::LinearTranspose;
1874 let linalg = semantic_linalg_op(extension, role)?;
1875 if !matches!(
1876 linalg.op(),
1877 LinalgOp::TriangularSolve { .. }
1878 | LinalgOp::LuSolvePrepared { .. }
1879 | LinalgOp::FullPivLuSolve { .. }
1880 | LinalgOp::Solve
1881 | LinalgOp::HouseholderQrAppendTangent
1882 ) {
1883 return Err(semantic_internal(
1884 role,
1885 format!(
1886 "linear linalg fragment contains unsupported extension {:?}",
1887 linalg.op()
1888 ),
1889 ));
1890 }
1891 let mut context = ShapeGuardContext::default();
1892 let mut fixed_values = HashMap::new();
1893 let mut keys = Vec::with_capacity(operation.inputs.len());
1894 let mut shape_sources = Vec::with_capacity(operation.inputs.len());
1895 for (index, (input, active)) in operation.inputs.iter().zip(active_mask).enumerate() {
1896 let key = ValueKey::Input(TensorInputKey::User {
1897 id: 10_000 + u64::try_from(index).expect("small linalg extension arity"),
1898 });
1899 let value = if *active {
1900 cotangent
1901 } else {
1902 match input {
1903 SemanticLinearFragmentInput::External(key) => external_values.get(key).copied(),
1904 SemanticLinearFragmentInput::Local(local) => {
1905 fixed_locals.get(*local).copied().flatten()
1906 }
1907 }
1908 .ok_or_else(|| {
1909 semantic_internal(role, "linear solve fragment is missing a fixed operand")
1910 })?
1911 };
1912 let metadata = builder.value_metadata(value)?.clone();
1913 let symbolic_shapes = synthetic_input_shapes(std::slice::from_ref(&metadata));
1914 let symbolic_shape_refs: Vec<_> = symbolic_shapes.iter().map(Vec::as_slice).collect();
1915 context.insert_metadata(
1916 key.clone(),
1917 legacy_metadata(&metadata, &symbolic_shape_refs),
1918 );
1919 if !active {
1920 fixed_values.insert(key.clone(), value);
1921 }
1922 shape_sources.push(value);
1923 keys.push(key);
1924 }
1925 let inputs: Vec<_> = keys
1926 .iter()
1927 .cloned()
1928 .map(PrimitiveTransposeInput::Residual)
1929 .collect();
1930 let mut emitted = SemanticRuleBuilder::with_seeds(
1931 &[Some(cotangent)],
1932 &fixed_values,
1933 &shape_sources,
1934 builder,
1935 role,
1936 );
1937 let outputs = LinalgAdRule
1938 .linear_transpose(
1939 extension,
1940 &mut emitted,
1941 &[Some(0)],
1942 &inputs,
1943 active_mask,
1944 &mut context,
1945 )
1946 .map_err(|error| legacy_error(role, error))?;
1947 let locals = emitted.finish()?;
1948 Ok(outputs
1949 .into_iter()
1950 .map(|output| output.and_then(|local| locals.get(local).copied().flatten()))
1951 .collect())
1952}
1953
1954fn accumulate_local_cotangent(
1955 builder: &mut SemanticProgramBuilder,
1956 cotangents: &mut HashMap<LocalValueId, ProgramValue>,
1957 local: LocalValueId,
1958 cotangent: ProgramValue,
1959) -> Result<(), SemanticAdError> {
1960 if let Some(existing) = cotangents.get_mut(&local) {
1961 *existing = builder.add_op(CoreSemanticOp::Add, &[*existing, cotangent])?[0];
1962 } else {
1963 cotangents.insert(local, cotangent);
1964 }
1965 Ok(())
1966}
1967
1968fn semantic_linalg_op(
1969 op: &dyn tenferro_ad::extension::ExtensionOp,
1970 role: SemanticAdRuleRole,
1971) -> Result<&LinalgExtensionOp, SemanticAdError> {
1972 op.as_any()
1973 .downcast_ref::<LinalgExtensionOp>()
1974 .ok_or_else(|| SemanticAdError::Unsupported {
1975 family_id: LINALG_EXTENSION_FAMILY_ID,
1976 role,
1977 message: "linalg semantic AD received an incompatible payload".into(),
1978 })
1979}
1980
1981fn legacy_error(role: SemanticAdRuleRole, error: tenferro_ops::ad::ADRuleError) -> SemanticAdError {
1982 SemanticAdError::Rule {
1983 family_id: LINALG_EXTENSION_FAMILY_ID,
1984 role,
1985 source: Box::new(error),
1986 }
1987}
1988
1989fn semantic_internal(role: SemanticAdRuleRole, message: impl Into<String>) -> SemanticAdError {
1990 SemanticAdError::Invariant {
1991 family_id: LINALG_EXTENSION_FAMILY_ID,
1992 role,
1993 message: message.into(),
1994 }
1995}
1996
1997#[cfg(test)]
1998mod tests {
1999 use super::*;
2000 use tenferro_runtime::program::ProgramInputSpec;
2001 use tenferro_tensor::DType;
2002
2003 #[test]
2004 fn recorded_broadcast_prefers_primal_shape_source_over_rank_compatible_data_input() {
2005 let mut builder = SemanticProgramBuilder::new();
2006 let _row_anchor = builder
2007 .input(ProgramInputSpec::new(DType::F64, [DimExpr::Const(3)]))
2008 .unwrap();
2009 let _col_anchor = builder
2010 .input(ProgramInputSpec::new(DType::F64, [DimExpr::Const(2)]))
2011 .unwrap();
2012 let matrix = builder
2013 .input(ProgramInputSpec::new(
2014 DType::F64,
2015 [
2016 DimExpr::InputDim {
2017 input_idx: 0,
2018 axis: 0,
2019 },
2020 DimExpr::InputDim {
2021 input_idx: 1,
2022 axis: 0,
2023 },
2024 ],
2025 ))
2026 .unwrap();
2027 let vector = builder
2028 .input(ProgramInputSpec::new(
2029 DType::F64,
2030 [DimExpr::Min(
2031 Box::new(DimExpr::InputDim {
2032 input_idx: 0,
2033 axis: 0,
2034 }),
2035 Box::new(DimExpr::InputDim {
2036 input_idx: 1,
2037 axis: 0,
2038 }),
2039 )],
2040 ))
2041 .unwrap();
2042
2043 let (operation, inputs) = localize_shape_expressions(
2044 CoreSemanticOp::BroadcastInDim {
2045 shape: vec![
2046 DimExpr::InputDim {
2047 input_idx: 0,
2048 axis: 0,
2049 },
2050 DimExpr::InputDim {
2051 input_idx: 0,
2052 axis: 1,
2053 },
2054 ],
2055 dims: vec![1],
2056 },
2057 &[vector],
2058 &[matrix],
2059 &builder,
2060 SemanticAdRuleRole::Linearize,
2061 )
2062 .unwrap();
2063
2064 assert_eq!(inputs, vec![vector, matrix]);
2065 assert_eq!(
2066 operation,
2067 CoreSemanticOp::BroadcastInDim {
2068 shape: vec![
2069 DimExpr::InputDim {
2070 input_idx: 1,
2071 axis: 0,
2072 },
2073 DimExpr::InputDim {
2074 input_idx: 1,
2075 axis: 1,
2076 },
2077 ],
2078 dims: vec![1],
2079 }
2080 );
2081 }
2082
2083 #[test]
2084 fn linalg_residual_mask_declares_all_inputs_and_outputs() {
2085 let mask = LinalgAdRule.residual_mask();
2089 assert!(mask.declares_input(0));
2090 assert!(mask.declares_input(1));
2091 assert!(mask.declares_output(0));
2092 assert!(mask.declares_output(1));
2093 assert!(mask.declares_output(2));
2094 }
2095}