use std::collections::HashMap;
use std::sync::Arc;
use computegraph::graph::GraphBuilder;
use computegraph::{LocalValueId, OperationRole, ValueRef};
use tenferro_ops::input_key::TensorInputKey;
use tenferro_ops::std_tensor_op::StdTensorOp;
use tenferro_ops::{SymDim, TensorMeta};
use tenferro_runtime::ad_support::{
allocate_input_key, allocate_shape_tensor_id, checkpoint_tensor, compile_ad_source,
frozen_input_value, inputs_map as tensor_inputs_map, leaf_input_key,
metadata_scopes as tensor_metadata_scopes, metadata_scopes_with_new, ones_tensor,
register_scoped_graph_analysis, shape_hint as tensor_shape_hint, tensor_from_parts,
ConstraintScopeTransfer, RetainedValue, TracedTensorParts,
};
use tenferro_runtime::program::{FrozenProgram, ProgramValue, ProgramValueMetadata, SemanticOpRef};
use tenferro_runtime::{
CompiledGraph, Error, ErrorPhase, GraphCompiler, Result, Runtime, Tensor, TracedTensor,
};
use crate::semantic_extension::{SemanticAdError, SemanticExtensionRuleSet};
use crate::semantic_transform::{
semantic_jvp, semantic_vjp, SemanticAdProgram, SemanticAdTransformError,
};
use crate::transform_cache::{AdTransformCache, SemanticAdTransformCacheKey};
pub(crate) fn next_input_key() -> TensorInputKey {
tenferro_runtime::ad_support::allocate_input_key()
}
fn error_shape_hint(tensor: &TracedTensor) -> Vec<usize> {
tensor
.try_concrete_shape()
.unwrap_or_else(|| vec![0; tensor.rank])
}
pub(crate) fn grad_with_rules_and_cache(
output: &TracedTensor,
wrt: &TracedTensor,
rules: &SemanticExtensionRuleSet,
ad_transform_cache: Option<&AdTransformCache>,
) -> Result<TracedTensor> {
grad_with_optional_rules(output, wrt, rules, ad_transform_cache)
}
pub(crate) fn jvp_with_rules_and_cache(
output: &TracedTensor,
wrt: &TracedTensor,
tangent: &TracedTensor,
rules: &SemanticExtensionRuleSet,
ad_transform_cache: Option<&AdTransformCache>,
) -> Result<TracedTensor> {
let wrt_input_key = leaf_input_key(wrt)?;
jvp_optional_impl(output, wrt, tangent, rules, ad_transform_cache)?
.ok_or_else(|| Error::Internal(format!("jvp output is inactive for {:?}", wrt_input_key)))
}
pub(crate) fn grad_optional_with_rules_and_cache(
output: &TracedTensor,
wrt: &TracedTensor,
rules: &SemanticExtensionRuleSet,
ad_transform_cache: Option<&AdTransformCache>,
) -> Result<Option<TracedTensor>> {
if output.rank != 0 {
return Err(Error::NonScalarGrad {
shape: error_shape_hint(output),
});
}
let ones = ones_tensor(output.dtype, vec![])?;
let seed = TracedTensor::from_tensor_concrete_shape(ones)?;
vjp_optional_impl(output, wrt, &seed, rules, "grad", ad_transform_cache)
}
pub(crate) fn jvp_optional_with_rules_and_cache(
output: &TracedTensor,
wrt: &TracedTensor,
tangent: &TracedTensor,
rules: &SemanticExtensionRuleSet,
ad_transform_cache: Option<&AdTransformCache>,
) -> Result<Option<TracedTensor>> {
jvp_optional_impl(output, wrt, tangent, rules, ad_transform_cache)
}
pub(crate) fn vjp_with_rules_and_cache(
output: &TracedTensor,
wrt: &TracedTensor,
cotangent: &TracedTensor,
rules: &SemanticExtensionRuleSet,
ad_transform_cache: Option<&AdTransformCache>,
) -> Result<TracedTensor> {
let wrt_input_key = leaf_input_key(wrt)?;
vjp_optional_impl(output, wrt, cotangent, rules, "vjp", ad_transform_cache)?
.ok_or_else(|| Error::Internal(format!("vjp output is inactive for {:?}", wrt_input_key)))
}
pub(crate) fn vjp_optional_with_rules_and_cache(
output: &TracedTensor,
wrt: &TracedTensor,
cotangent: &TracedTensor,
rules: &SemanticExtensionRuleSet,
ad_transform_cache: Option<&AdTransformCache>,
) -> Result<Option<TracedTensor>> {
vjp_optional_impl(output, wrt, cotangent, rules, "vjp", ad_transform_cache)
}
fn grad_with_optional_rules(
output: &TracedTensor,
wrt: &TracedTensor,
rules: &SemanticExtensionRuleSet,
ad_transform_cache: Option<&AdTransformCache>,
) -> Result<TracedTensor> {
if output.rank != 0 {
return Err(Error::NonScalarGrad {
shape: error_shape_hint(output),
});
}
let ones = ones_tensor(output.dtype, vec![])?;
let seed = TracedTensor::from_tensor_concrete_shape(ones)?;
let wrt_input_key = leaf_input_key(wrt)?;
vjp_optional_impl(output, wrt, &seed, rules, "grad", ad_transform_cache)?
.ok_or_else(|| Error::Internal(format!("grad output is inactive for {:?}", wrt_input_key)))
}
fn single_runtime_output(mut outputs: Vec<Tensor>, op: &'static str) -> Result<Tensor> {
let actual = outputs.len();
if actual != 1 {
return Err(Error::runtime_state(
op,
ErrorPhase::Execution,
format!("expected one runtime output, got {actual}"),
));
}
outputs.pop().ok_or_else(|| {
Error::runtime_state(
op,
ErrorPhase::Execution,
"runtime returned no output after successful output-count validation",
)
})
}
pub trait TracedTensorAdExt {
fn grad(&self, wrt: &TracedTensor) -> Result<TracedTensor>;
fn grad_optional(&self, wrt: &TracedTensor) -> Result<Option<TracedTensor>>;
fn checkpoint(&mut self, compiler: &mut GraphCompiler, runtime: &Runtime) -> Result<()>;
fn jvp(&self, wrt: &TracedTensor, tangent: &TracedTensor) -> Result<TracedTensor>;
fn jvp_optional(
&self,
wrt: &TracedTensor,
tangent: &TracedTensor,
) -> Result<Option<TracedTensor>>;
fn vjp(&self, wrt: &TracedTensor, cotangent: &TracedTensor) -> Result<TracedTensor>;
fn vjp_optional(
&self,
wrt: &TracedTensor,
cotangent: &TracedTensor,
) -> Result<Option<TracedTensor>>;
}
impl TracedTensorAdExt for TracedTensor {
fn grad(&self, wrt: &TracedTensor) -> Result<TracedTensor> {
let rules = SemanticExtensionRuleSet::default();
grad_with_optional_rules(self, wrt, &rules, None)
}
fn grad_optional(&self, wrt: &TracedTensor) -> Result<Option<TracedTensor>> {
if self.rank != 0 {
return Err(Error::NonScalarGrad {
shape: error_shape_hint(self),
});
}
let ones = ones_tensor(self.dtype, vec![])?;
let seed = TracedTensor::from_tensor_concrete_shape(ones)?;
let rules = SemanticExtensionRuleSet::default();
vjp_optional_impl(self, wrt, &seed, &rules, "grad", None)
}
fn checkpoint(&mut self, compiler: &mut GraphCompiler, runtime: &Runtime) -> Result<()> {
let data = if let Some(data) = self.attached_value() {
Arc::clone(data)
} else {
let program = compiler.compile(self)?;
Arc::new(RetainedValue::from_tensor(single_runtime_output(
runtime.run_compiled(&program, &[])?,
"TracedTensorAdExt::checkpoint",
)?))
};
checkpoint_tensor(self, data)?;
Ok(())
}
fn jvp(&self, wrt: &TracedTensor, tangent: &TracedTensor) -> Result<TracedTensor> {
let wrt_input_key = leaf_input_key(wrt)?;
self.jvp_optional(wrt, tangent)?.ok_or_else(|| {
Error::Internal(format!("jvp output is inactive for {:?}", wrt_input_key))
})
}
fn jvp_optional(
&self,
wrt: &TracedTensor,
tangent: &TracedTensor,
) -> Result<Option<TracedTensor>> {
let rules = SemanticExtensionRuleSet::default();
jvp_optional_impl(self, wrt, tangent, &rules, None)
}
fn vjp(&self, wrt: &TracedTensor, cotangent: &TracedTensor) -> Result<TracedTensor> {
let wrt_input_key = leaf_input_key(wrt)?;
self.vjp_optional(wrt, cotangent)?.ok_or_else(|| {
Error::Internal(format!("vjp output is inactive for {:?}", wrt_input_key))
})
}
fn vjp_optional(
&self,
wrt: &TracedTensor,
cotangent: &TracedTensor,
) -> Result<Option<TracedTensor>> {
let rules = SemanticExtensionRuleSet::default();
vjp_optional_impl(self, wrt, cotangent, &rules, "vjp", None)
}
}
fn jvp_optional_impl(
output: &TracedTensor,
wrt: &TracedTensor,
tangent: &TracedTensor,
rules: &SemanticExtensionRuleSet,
ad_transform_cache: Option<&AdTransformCache>,
) -> Result<Option<TracedTensor>> {
let wrt_input_key = leaf_input_key(wrt)?;
let tangent_data = tangent.attached_value().cloned().ok_or_else(|| {
Error::invalid_argument(
"jvp",
ErrorPhase::GraphBuild,
"tangent",
"jvp tangent must have concrete tensor data",
)
})?;
let mut compiler = GraphCompiler::new();
let source = compile_ad_source(&mut compiler, output)?;
let Some(wrt_input_index) = source.input_key_index(&wrt_input_key) else {
return Ok(None);
};
let mut active_inputs = vec![false; source.input_count()];
active_inputs[wrt_input_index] = true;
let derivative = semantic_jvp_with_cache(
source.frozen_program(),
&active_inputs,
rules,
ad_transform_cache,
)?;
let Some(seed_input_index) = derivative
.derivative_input_indices()
.get(wrt_input_index)
.copied()
.flatten()
else {
return Ok(None);
};
let Some(derivative_output_index) = derivative
.derivative_output_indices()
.first()
.copied()
.flatten()
else {
return Ok(None);
};
derivative_tensor_from_program(
&source,
&derivative,
derivative_output_index,
&[(seed_input_index, tangent_data)],
[output, wrt, tangent],
tensor_shape_hint(output),
"jvp",
)
.map(Some)
}
fn vjp_optional_impl(
output: &TracedTensor,
wrt: &TracedTensor,
cotangent: &TracedTensor,
rules: &SemanticExtensionRuleSet,
transform: &'static str,
ad_transform_cache: Option<&AdTransformCache>,
) -> Result<Option<TracedTensor>> {
let wrt_input_key = leaf_input_key(wrt)?;
let cotangent_data = cotangent.attached_value().cloned().ok_or_else(|| {
Error::invalid_argument(
transform,
ErrorPhase::GraphBuild,
"cotangent",
"vjp cotangent must have concrete tensor data",
)
})?;
let mut compiler = GraphCompiler::new();
let source = compile_ad_source(&mut compiler, output)?;
let Some(wrt_input_index) = source.input_key_index(&wrt_input_key) else {
return Ok(None);
};
let mut active_inputs = vec![false; source.input_count()];
active_inputs[wrt_input_index] = true;
let active_outputs = vec![true; source.output_count()];
let derivative = semantic_vjp_with_cache(
source.frozen_program(),
&active_inputs,
&active_outputs,
rules,
ad_transform_cache,
)?;
let Some(seed_input_index) = derivative
.derivative_input_indices()
.first()
.copied()
.flatten()
else {
return Ok(None);
};
let Some(derivative_output_index) = derivative
.derivative_output_indices()
.get(wrt_input_index)
.copied()
.flatten()
else {
return Ok(None);
};
derivative_tensor_from_program(
&source,
&derivative,
derivative_output_index,
&[(seed_input_index, cotangent_data)],
[output, wrt, cotangent],
tensor_shape_hint(wrt),
transform,
)
.map(Some)
}
fn semantic_jvp_with_cache(
source: &FrozenProgram,
active_inputs: &[bool],
rules: &SemanticExtensionRuleSet,
ad_transform_cache: Option<&AdTransformCache>,
) -> Result<SemanticAdProgram> {
let key = SemanticAdTransformCacheKey::jvp(source, active_inputs);
if let Some(cache) = ad_transform_cache {
if let Some(cached) = cache.get_semantic(&key, source)? {
return cached
.as_ref()
.with_input_prefix_bindings_from(source)
.map_err(|source| {
Error::runtime_state_source(
"semantic traced jvp cache",
ErrorPhase::GraphBuild,
source,
)
});
}
}
let derivative =
semantic_jvp(source, active_inputs, rules).map_err(semantic_transform_error("jvp"))?;
if let Some(cache) = ad_transform_cache {
cache.put_semantic(key, source, Arc::new(derivative.clone()))?;
}
Ok(derivative)
}
fn semantic_vjp_with_cache(
source: &FrozenProgram,
active_inputs: &[bool],
active_outputs: &[bool],
rules: &SemanticExtensionRuleSet,
ad_transform_cache: Option<&AdTransformCache>,
) -> Result<SemanticAdProgram> {
let key = SemanticAdTransformCacheKey::vjp(source, active_inputs, active_outputs);
if let Some(cache) = ad_transform_cache {
if let Some(cached) = cache.get_semantic(&key, source)? {
return cached
.as_ref()
.with_input_prefix_bindings_from(source)
.map_err(|source| {
Error::runtime_state_source(
"semantic traced vjp cache",
ErrorPhase::GraphBuild,
source,
)
});
}
}
let derivative = semantic_vjp(source, active_inputs, active_outputs, rules)
.map_err(semantic_transform_error("vjp"))?;
if let Some(cache) = ad_transform_cache {
cache.put_semantic(key, source, Arc::new(derivative.clone()))?;
}
Ok(derivative)
}
fn semantic_transform_error(
transform: &'static str,
) -> impl FnOnce(SemanticAdTransformError) -> Error {
move |source| {
semantic_transform_validation_error(transform, &source).unwrap_or_else(|| {
Error::runtime_state_source(transform, ErrorPhase::GraphBuild, source)
})
}
}
fn semantic_transform_validation_error(
transform: &'static str,
source: &SemanticAdTransformError,
) -> Option<Error> {
if let SemanticAdTransformError::Extension(
SemanticAdError::Unsupported { family_id, .. }
| SemanticAdError::MissingRule { family_id, .. },
) = source
{
return Some(Error::UnsupportedAdRule {
transform,
op: (*family_id).to_owned(),
});
}
let SemanticAdTransformError::Extension(SemanticAdError::Rule { source, .. }) = source else {
return None;
};
let tenferro_ops::ad::ADRuleError::InvalidInput { op, message, .. } =
source.downcast_ref::<tenferro_ops::ad::ADRuleError>()?
else {
return None;
};
Some(Error::invalid_argument(
transform,
ErrorPhase::GraphBuild,
"semantic_ad_rule",
format!("{op}: {message}"),
))
}
fn derivative_tensor_from_program(
source: &CompiledGraph,
derivative: &SemanticAdProgram,
derivative_output_index: usize,
seed_tensors: &[(usize, Arc<RetainedValue>)],
inherited_tensors: [&TracedTensor; 3],
fallback_shape_hint: Option<Vec<SymDim>>,
transform: &'static str,
) -> Result<TracedTensor> {
derivative_trace_from_frozen_program(
source,
derivative.frozen(),
derivative_output_index,
seed_tensors,
&inherited_tensors,
fallback_shape_hint,
transform,
)
}
pub(crate) fn derivative_trace_from_frozen_program(
source: &CompiledGraph,
frozen: &FrozenProgram,
derivative_output_index: usize,
seed_tensors: &[(usize, Arc<RetainedValue>)],
inherited_tensors: &[&TracedTensor],
fallback_shape_hint: Option<Vec<SymDim>>,
transform: &'static str,
) -> Result<TracedTensor> {
let input_shapes = symbolic_input_shapes(frozen)?;
let input_shape_refs: Vec<_> = input_shapes.iter().map(Vec::as_slice).collect();
let input_metas = frozen
.program
.inputs()
.iter()
.copied()
.map(|value| tensor_meta_for_value(frozen, value, &input_shape_refs, transform))
.collect::<Result<Vec<_>>>()?;
let output_value = *frozen
.program
.outputs()
.get(derivative_output_index)
.ok_or_else(|| {
Error::runtime_state(
transform,
ErrorPhase::GraphBuild,
format!(
"derivative output index {derivative_output_index} is outside {} outputs",
frozen.program.outputs().len()
),
)
})?;
let output_meta = tensor_meta_for_value(frozen, output_value, &input_shape_refs, transform)?;
let mut builder = GraphBuilder::<StdTensorOp>::new();
let mut value_map = HashMap::<ProgramValue, LocalValueId>::new();
let mut input_keys = Vec::with_capacity(frozen.program.inputs().len());
for (input_index, input) in frozen.program.inputs().iter().copied().enumerate() {
let key = if input_index < source.input_keys().len() {
source.input_keys()[input_index].clone()
} else {
allocate_input_key()
};
let local = builder.add_input(key.clone());
value_map.insert(input, local);
input_keys.push(key);
}
for operation in frozen.program.operations() {
let inputs = operation
.inputs()
.iter()
.copied()
.map(|value| {
value_map
.get(&value)
.copied()
.map(ValueRef::Local)
.ok_or_else(|| missing_program_value(transform, "operation input"))
})
.collect::<Result<Vec<_>>>()?;
let op = match operation.op() {
SemanticOpRef::Core(op) => StdTensorOp::from(op),
SemanticOpRef::Extension(op) => StdTensorOp::Extension(op.clone_arc()),
_ => {
return Err(Error::runtime_state(
transform,
ErrorPhase::GraphBuild,
"unsupported semantic operation variant in derivative graph",
));
}
};
let outputs = builder.add_operation(op, inputs, OperationRole::Primary);
if outputs.len() != operation.outputs().len() {
return Err(Error::runtime_state(
transform,
ErrorPhase::GraphBuild,
format!(
"semantic operation expected {} outputs, graph builder produced {}",
operation.outputs().len(),
outputs.len()
),
));
}
for (value, local) in operation.outputs().iter().copied().zip(outputs) {
value_map.insert(value, local);
}
}
let graph_outputs = frozen
.program
.outputs()
.iter()
.copied()
.map(|value| {
value_map
.get(&value)
.copied()
.ok_or_else(|| missing_program_value(transform, "program output"))
})
.collect::<Result<Vec<_>>>()?;
let val = *graph_outputs.get(derivative_output_index).ok_or_else(|| {
Error::runtime_state(
transform,
ErrorPhase::GraphBuild,
"derivative output index missing after graph conversion",
)
})?;
builder.set_outputs(graph_outputs);
let graph = Arc::new(builder.build());
let Some(primary_tensor) = inherited_tensors.first() else {
return Err(Error::runtime_state(
transform,
ErrorPhase::GraphBuild,
"derivative trace construction requires inherited source tensors",
));
};
let mut inputs_map = (*tensor_inputs_map(primary_tensor)).clone();
for (input_index, key) in input_keys.iter().enumerate() {
if let Some(tensor) = frozen_input_value(frozen, input_index) {
inputs_map.insert(key.clone(), tensor);
}
}
for (seed_input_index, tensor) in seed_tensors {
let meta = input_metas.get(*seed_input_index).ok_or_else(|| {
Error::runtime_state(
transform,
ErrorPhase::GraphBuild,
format!("seed input index {seed_input_index} is outside derivative inputs"),
)
})?;
validate_seed_tensor(transform, *seed_input_index, tensor.as_ref(), meta)?;
let key = input_keys.get(*seed_input_index).ok_or_else(|| {
Error::runtime_state(
transform,
ErrorPhase::GraphBuild,
format!("seed input key {seed_input_index} is outside derivative inputs"),
)
})?;
inputs_map.insert(key.clone(), Arc::clone(tensor));
}
let source_input_count = source.input_keys().len();
let graph_input_metadata = graph
.inputs()
.iter()
.copied()
.zip(input_metas.iter().cloned())
.enumerate()
.filter_map(|(input_index, (input, meta))| {
if input_index < source_input_count {
None
} else {
Some((graph.values()[input].key.clone(), meta))
}
});
let analysis = register_scoped_graph_analysis(graph.as_ref(), graph_input_metadata)?;
let inherited_constraint_scopes = inherited_tensors
.iter()
.map(|tensor| ConstraintScopeTransfer::from_tensor(tensor))
.collect::<Vec<_>>();
Ok(tensor_from_parts(TracedTensorParts {
rank: output_meta.rank(),
dtype: output_meta.dtype,
graph,
val,
data: None,
shape_hint: output_meta.exact_shape().or(fallback_shape_hint),
inputs_map: Arc::new(inputs_map),
extra_roots: Vec::new(),
checkpoint_chain: None,
metadata_scopes: metadata_scopes_with_new(
analysis.metadata,
inherited_tensors
.iter()
.map(|tensor| tensor_metadata_scopes(tensor)),
),
constraint_scope_transfer: ConstraintScopeTransfer::with_new(
analysis.constraints,
inherited_constraint_scopes.iter(),
),
}))
}
fn missing_program_value(transform: &'static str, role: &'static str) -> Error {
Error::runtime_state(
transform,
ErrorPhase::GraphBuild,
format!("semantic derivative graph references missing {role}"),
)
}
fn symbolic_input_shapes(frozen: &FrozenProgram) -> Result<Vec<Vec<SymDim>>> {
frozen
.program
.inputs()
.iter()
.copied()
.map(|value| {
let meta = frozen.program.value_metadata(value).map_err(|source| {
Error::runtime_state_source(
"semantic traced AD input metadata",
ErrorPhase::GraphBuild,
source,
)
})?;
let tensor_id = allocate_shape_tensor_id();
Ok((0..meta.shape().len())
.map(|axis| SymDim::tensor_axis(tensor_id, axis))
.collect())
})
.collect()
}
fn tensor_meta_for_value(
frozen: &FrozenProgram,
value: ProgramValue,
input_shapes: &[&[SymDim]],
transform: &'static str,
) -> Result<TensorMeta> {
let meta = frozen
.program
.value_metadata(value)
.map_err(|source| Error::runtime_state_source(transform, ErrorPhase::GraphBuild, source))?;
Ok(program_metadata_to_tensor_meta(meta, input_shapes))
}
fn program_metadata_to_tensor_meta(
metadata: &ProgramValueMetadata,
input_shapes: &[&[SymDim]],
) -> TensorMeta {
let extents = metadata
.shape()
.iter()
.cloned()
.map(|extent| extent.map(|dim| SymDim::from_dim_expr(&dim, input_shapes)))
.collect();
TensorMeta::with_extents(metadata.dtype(), extents)
}
fn validate_seed_tensor(
transform: &'static str,
input_index: usize,
tensor: &RetainedValue,
expected: &TensorMeta,
) -> Result<()> {
let actual_dtype = tensor.dtype();
if actual_dtype != expected.dtype {
return Err(Error::invalid_argument(
transform,
ErrorPhase::GraphBuild,
"seed",
format!(
"seed input {input_index} dtype mismatch: expected {:?}, got {:?}",
expected.dtype, actual_dtype
),
));
}
let actual_shape = tensor.shape();
if actual_shape.len() != expected.rank() {
return Err(Error::invalid_argument(
transform,
ErrorPhase::GraphBuild,
"seed",
format!(
"seed input {input_index} rank mismatch: expected {}, got {}",
expected.rank(),
actual_shape.len()
),
));
}
if let Some(expected_shape) = expected
.exact_shape()
.filter(|shape| shape.iter().all(|dim| dim.constant_value().is_some()))
.map(|shape| {
shape
.into_iter()
.map(|dim| dim.constant_value().expect("filtered constant shape"))
.collect::<Vec<_>>()
})
{
if expected_shape != actual_shape {
return Err(Error::invalid_argument(
transform,
ErrorPhase::GraphBuild,
"seed",
format!(
"seed input {input_index} shape mismatch: expected {:?}, got {:?}",
expected_shape, actual_shape
),
));
}
}
Ok(())
}
#[cfg(test)]
mod semantic_transform_error_tests {
use super::*;
use crate::semantic_extension::SemanticAdRuleRole;
#[test]
fn unsupported_semantic_rule_maps_to_public_transform_error() {
let source = SemanticAdTransformError::Extension(SemanticAdError::Unsupported {
family_id: "tenferro-tests.unsupported.v1",
role: SemanticAdRuleRole::LinearTranspose,
message: "unsupported test payload".into(),
});
let error = semantic_transform_validation_error("vjp", &source)
.expect("semantic rejection must map to a public unsupported-rule error");
assert!(matches!(
error,
Error::UnsupportedAdRule { transform: "vjp", ref op }
if op == "tenferro-tests.unsupported.v1"
));
}
#[test]
fn missing_semantic_rule_maps_to_public_transform_error() {
let source = SemanticAdTransformError::Extension(SemanticAdError::MissingRule {
family_id: "tenferro-tests.missing.v1",
role: SemanticAdRuleRole::Linearize,
});
let error = semantic_transform_validation_error("jvp", &source)
.expect("missing semantic rule must map to a public unsupported-rule error");
assert!(matches!(
error,
Error::UnsupportedAdRule { transform: "jvp", ref op }
if op == "tenferro-tests.missing.v1"
));
}
}