use super::{ADKey, ADRuleResult, PrimitiveBuilder, PrimitiveValue};
use computegraph::{GraphOperation, LocalValueId, OperationRole, ValueKey};
pub trait Primitive: GraphOperation
where
Self::InputKey: ADKey,
{
type ADContext: Default;
fn add() -> Self
where
Self: Sized;
fn jvp_rule(
&self,
builder: &mut impl PrimitiveBuilder<Self>,
primal_inputs: &[ValueKey<Self>],
primal_outputs: &[ValueKey<Self>],
tangent_inputs: &[Option<LocalValueId>],
ctx: &mut Self::ADContext,
) -> Vec<Option<LocalValueId>>
where
Self: Sized;
fn try_jvp_rule(
&self,
builder: &mut impl PrimitiveBuilder<Self>,
primal_inputs: &[ValueKey<Self>],
primal_outputs: &[ValueKey<Self>],
tangent_inputs: &[Option<LocalValueId>],
ctx: &mut Self::ADContext,
) -> ADRuleResult<Vec<Option<LocalValueId>>>
where
Self: Sized,
{
Ok(self.jvp_rule(builder, primal_inputs, primal_outputs, tangent_inputs, ctx))
}
fn transpose_rule(
&self,
builder: &mut impl PrimitiveBuilder<Self>,
cotangent_outputs: &[Option<LocalValueId>],
inputs: &[PrimitiveValue<Self>],
role: &OperationRole,
ctx: &mut Self::ADContext,
) -> Vec<Option<LocalValueId>>
where
Self: Sized;
fn try_linear_transpose_rule(
&self,
builder: &mut impl PrimitiveBuilder<Self>,
cotangent_outputs: &[Option<LocalValueId>],
inputs: &[PrimitiveValue<Self>],
role: &OperationRole,
ctx: &mut Self::ADContext,
) -> ADRuleResult<Vec<Option<LocalValueId>>>
where
Self: Sized,
{
Ok(self.transpose_rule(builder, cotangent_outputs, inputs, role, ctx))
}
}