Skip to main content

tenferro_linalg/ad/
semantic.rs

1// Solve residual policy reference: PyTorch 8dd3b763, FunctionsManual.cpp,
2// linalg_solve_backward (saved LU/pivots for gB, saved X for gA).
3// Emission below reuses tenferro's existing prepared-solve and cotangent helpers.
4use 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
29/// Return the linalg semantic-program AD rule set.
30///
31/// # Errors
32///
33/// Returns [`SemanticExtensionRegistryError::MalformedFamilyId`] if the linalg
34/// family identifier is invalid, or
35/// [`SemanticExtensionRegistryError::DuplicateRule`] if a semantic rule role
36/// is already registered.
37///
38/// # Examples
39///
40/// ```rust
41/// let rules = tenferro_linalg::semantic_ad_rules().unwrap();
42/// assert!(rules
43///     .lookup_linearize(tenferro_linalg::LINALG_EXTENSION_FAMILY_ID)
44///     .is_some());
45/// assert!(rules
46///     .lookup_linear_transpose(tenferro_linalg::LINALG_EXTENSION_FAMILY_ID)
47///     .is_some());
48/// assert!(rules
49///     .lookup_primal_vjp(tenferro_linalg::LINALG_EXTENSION_FAMILY_ID)
50///     .is_none());
51/// ```
52pub 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            // These value-only operations can appear inside a differentiable
78            // composite (for example, `solve` uses `LuFactor` outputs as
79            // prepared-solve residuals).  Returning absent tangents lets the
80            // composite rule differentiate through the primal inputs without
81            // pretending that the factorization outputs are differentiable.
82            // A caller requesting those outputs directly still observes that
83            // no derivative output was produced, while VJP rejects their
84            // unsupported transpose below.
85            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        // The linalg family is one rule across solve/eigen/qr ops whose
143        // transposes collectively read every operand: triangular solve reads
144        // both inputs plus the solution output, and the linearize+fragment
145        // path can consume any input/output value (svd reads outputs 0-2,
146        // eigh/qr read outputs 0-1).
147        // ponytail: family-level mask; per-op masks would require splitting
148        // `LinalgAdRule` per op.
149        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            // The fused solve saves its own factors as outputs 1 and 2, so the
210            // adjoint solve reuses them instead of refactoring `a`. The factor
211            // outputs carry no cotangent (see the `LuFactor` linearize rule).
212            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    // Packed LU may be an active intermediate when `solve` lowers through
292    // factorization. Pivot and parity slots remain non-cotangent-producing
293    // residuals.
294    Ok([active_inputs[0], false, false, active_inputs[3]])
295}
296
297/// Primal values of one prepared LU solve `op(A) x = b`.
298struct SemanticPreparedSolve {
299    op: &'static str,
300    a: ProgramValue,
301    packed_lu: ProgramValue,
302    pivots: ProgramValue,
303    solution: Option<ProgramValue>,
304}
305
306/// Cotangents `(gA, gB)` of a prepared LU solve `op(A) x = b`.
307///
308/// `gB = op(A)^{-H} gX` is one prepared solve against the saved factors with
309/// both flags flipped, and `gA = -op(gB x^H)` reads the saved solution
310/// (PyTorch `linalg_solve_backward`).
311fn 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
823/// Op-local semantic linear fragment used by the manifest's
824/// `LinearizeThenTranspose` route. General linalg decomposition VJPs remain
825/// derived by transposing their emitted linearization, but the fragment stores
826/// semantic core ops or extension payloads instead of replaying legacy
827/// `StdTensorOp` graphs into the destination builder.
828struct 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            // Legacy linalg rules express shape dimensions in invocation
1185            // coordinates. Shape-aware primitive helpers append explicit
1186            // shape operands after the primitive's fixed data operands and
1187            // remap their dimensions into those local operand coordinates.
1188            // Keep those two coordinate spaces distinct; rank compatibility
1189            // cannot disambiguate them.
1190            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                // The eigenvalue-value rule re-emits `Eigh` inside the linear
1584                // fragment to reach the eigenvectors, so the fragment shape
1585                // resolver must describe both outputs: the eigenvector matrix
1586                // keeps the input shape, and the eigenvalues drop the column
1587                // axis of the matrix core.
1588                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            // Hadamard products in a linalg linear fragment keep the operand
1621            // shape; the eigenvalue path emits one before its reduction, and the
1622            // transposed reduction needs that shape to restore the operands.
1623            | 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        // The linalg family is one rule across solve/eigen/qr ops. SVD reads
2086        // outputs 0-2 and eigh/qr read outputs 0-1 as tensor residuals, so the
2087        // mask must declare every input and output (issue #1665 step 5).
2088        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}