use std::sync::Arc;
use crate::Rational;
pub trait DualCoeff:
Clone
+ PartialEq
+ std::ops::Add<Output = Self>
+ std::ops::Sub<Output = Self>
+ std::ops::Mul<Output = Self>
+ std::ops::Div<Output = Self>
+ std::ops::Neg<Output = Self>
+ std::ops::AddAssign
+ std::ops::MulAssign
{
fn zero() -> Self;
fn one() -> Self;
}
impl DualCoeff for Rational {
fn zero() -> Self {
Rational::new(0, 1)
}
fn one() -> Self {
Rational::new(1, 1)
}
}
#[derive(Debug, Clone)]
pub struct DualShape {
components: Vec<Vec<usize>>,
mult_table: Vec<(usize, usize, usize)>,
}
impl DualShape {
pub fn new(mut components: Vec<Vec<usize>>) -> Option<Self> {
if components.is_empty() {
return None;
}
let nvars = components.iter().map(|c| c.len()).max()?;
for c in components.iter_mut() {
c.resize(nvars, 0);
}
let zero = vec![0usize; nvars];
let pos = components.iter().position(|c| *c == zero)?;
if pos != 0 {
components.swap(0, pos);
}
let contains = |m: &[usize]| -> bool {
components
.iter()
.any(|c| c.len() == nvars && c.as_slice() == m)
};
for c in &components {
if !ancestors_present(c, &contains) {
return None;
}
}
let mult_table = build_mult_table(&components);
Some(Self {
components,
mult_table,
})
}
pub fn n_components(&self) -> usize {
self.components.len()
}
pub fn n_vars(&self) -> usize {
self.components.first().map(|c| c.len()).unwrap_or(0)
}
pub fn components(&self) -> &[Vec<usize>] {
&self.components
}
pub fn mult_table(&self) -> &[(usize, usize, usize)] {
&self.mult_table
}
pub fn index_of(&self, multi_index: &[usize]) -> Option<usize> {
self.components
.iter()
.position(|c| c.as_slice() == multi_index)
}
}
fn ancestors_present<F: Fn(&[usize]) -> bool>(m: &[usize], contains: &F) -> bool {
let nvars = m.len();
let mut counts: Vec<usize> = vec![0; nvars];
loop {
let is_self = counts.iter().zip(m.iter()).all(|(a, b)| a == b);
if !is_self && !contains(&counts) {
return false;
}
let mut i = 0;
while i < nvars {
counts[i] += 1;
if counts[i] > m[i] {
counts[i] = 0;
i += 1;
} else {
break;
}
}
if i == nvars {
break;
}
}
true
}
fn build_mult_table(components: &[Vec<usize>]) -> Vec<(usize, usize, usize)> {
let nvars = components.first().map(|c| c.len()).unwrap_or(0);
let mut table = Vec::new();
for a in 1..components.len() {
for b in a..components.len() {
let sum: Vec<usize> = (0..nvars)
.map(|i| components[a][i] + components[b][i])
.collect();
if let Some(c) = components
.iter()
.position(|comp| comp.len() == nvars && comp.as_slice() == sum.as_slice())
{
table.push((a, b, c));
}
}
}
table
}
#[derive(Debug, Clone)]
pub struct HyperDual<T: DualCoeff> {
values: Vec<T>,
shape: Arc<DualShape>,
}
impl<T: DualCoeff> HyperDual<T> {
pub fn from_values(shape: Arc<DualShape>, values: Vec<T>) -> Option<Self> {
if values.len() != shape.n_components() {
return None;
}
Some(Self { values, shape })
}
pub fn value(&self) -> &T {
&self.values[0]
}
pub fn values(&self) -> &[T] {
&self.values
}
pub fn shape(&self) -> &Arc<DualShape> {
&self.shape
}
pub fn deriv(&self, i: usize) -> Option<&T> {
let mut idx = vec![0usize; self.shape.n_vars()];
if i >= idx.len() {
return None;
}
idx[i] = 1;
self.shape.index_of(&idx).map(|pos| &self.values[pos])
}
pub fn constant(shape: &Arc<DualShape>, c: T) -> Self {
let mut values = vec![T::zero(); shape.n_components()];
values[0] = c;
Self {
values,
shape: shape.clone(),
}
}
pub fn variable(shape: &Arc<DualShape>, i: usize, c: T) -> Self {
let mut d = Self::constant(shape, c);
let mut idx = vec![0usize; shape.n_vars()];
if i < idx.len() {
idx[i] = 1;
if let Some(pos) = shape.index_of(&idx) {
d.values[pos] = T::one();
}
}
d
}
pub fn zero(shape: &Arc<DualShape>) -> Self {
Self::constant(shape, T::zero())
}
pub fn one(shape: &Arc<DualShape>) -> Self {
Self::constant(shape, T::one())
}
pub fn inv(&self) -> Option<Self> {
let v = self.value().clone();
if v == T::zero() {
return None;
}
let inv_v = T::one().clone() / v.clone();
let mut neg_ratio = vec![T::zero(); self.values.len()];
for (k, slot) in neg_ratio.iter_mut().enumerate().skip(1) {
*slot = (self.values[k].clone() / v.clone()).neg();
}
let mut s = vec![T::zero(); self.values.len()];
s[0] = T::one();
let mut current_power = neg_ratio.clone();
loop {
let mut changed = false;
for k in 1..s.len() {
if current_power[k] != T::zero() {
s[k] += current_power[k].clone();
changed = true;
}
}
if !changed {
break;
}
current_power = mul_truncated(¤t_power, &neg_ratio, &self.shape);
}
let result_values: Vec<T> = s.iter().map(|c| inv_v.clone() * c.clone()).collect();
Some(Self {
values: result_values,
shape: self.shape.clone(),
})
}
}
fn assert_same_shape<T: DualCoeff>(a: &HyperDual<T>, b: &HyperDual<T>) {
debug_assert_eq!(
a.shape.n_components(),
b.shape.n_components(),
"HyperDual shape mismatch"
);
}
impl<T: DualCoeff> std::ops::Add for HyperDual<T> {
type Output = Self;
fn add(self, rhs: Self) -> Self {
assert_same_shape(&self, &rhs);
let values: Vec<T> = self
.values
.iter()
.zip(rhs.values.iter())
.map(|(a, b)| a.clone() + b.clone())
.collect();
Self {
values,
shape: self.shape,
}
}
}
impl<T: DualCoeff> std::ops::Sub for HyperDual<T> {
type Output = Self;
fn sub(self, rhs: Self) -> Self {
assert_same_shape(&self, &rhs);
let values: Vec<T> = self
.values
.iter()
.zip(rhs.values.iter())
.map(|(a, b)| a.clone() - b.clone())
.collect();
Self {
values,
shape: self.shape,
}
}
}
impl<T: DualCoeff> std::ops::Neg for HyperDual<T> {
type Output = Self;
fn neg(self) -> Self {
let values: Vec<T> = self.values.iter().map(|a| a.clone().neg()).collect();
Self {
values,
shape: self.shape,
}
}
}
impl<T: DualCoeff> std::ops::Mul for HyperDual<T> {
type Output = Self;
fn mul(self, rhs: Self) -> Self {
assert_same_shape(&self, &rhs);
let mut result = vec![T::zero(); self.values.len()];
let sv = self.values[0].clone();
let rv = rhs.values[0].clone();
for (k, slot) in result.iter_mut().enumerate().skip(1) {
*slot = self.values[k].clone() * rv.clone();
*slot += sv.clone() * rhs.values[k].clone();
}
result[0] = sv * rv;
for &(a, b, c) in self.shape.mult_table() {
result[c] += self.values[a].clone() * rhs.values[b].clone();
if a != b {
result[c] += self.values[b].clone() * rhs.values[a].clone();
}
}
Self {
values: result,
shape: self.shape,
}
}
}
impl<T: DualCoeff> std::ops::Div for HyperDual<T> {
type Output = Self;
#[allow(clippy::suspicious_arithmetic_impl)]
fn div(self, rhs: Self) -> Self {
let inv = rhs.inv().expect("division by zero-valued HyperDual");
self * inv
}
}
fn mul_truncated<T: DualCoeff>(self_v: &[T], rhs_v: &[T], shape: &DualShape) -> Vec<T> {
let mut result = vec![T::zero(); self_v.len()];
let sv = &self_v[0];
let rv = &rhs_v[0];
result[0] = sv.clone() * rv.clone();
for k in 1..result.len() {
result[k] = self_v[k].clone() * rv.clone();
result[k] += sv.clone() * rhs_v[k].clone();
}
for &(a, b, c) in shape.mult_table() {
result[c] += self_v[a].clone() * rhs_v[b].clone();
if a != b {
result[c] += self_v[b].clone() * rhs_v[a].clone();
}
}
result
}
pub fn new_first_order<T: DualCoeff>(nvars: usize) -> Arc<DualShape> {
let mut components = vec![vec![0usize; nvars]];
for i in 0..nvars {
let mut m = vec![0usize; nvars];
m[i] = 1;
components.push(m);
}
Arc::new(DualShape::new(components).expect("first-order shape is ancestor-closed"))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Rational;
fn r(n: i64, d: i64) -> Rational {
Rational::new(n, d)
}
#[test]
fn first_order_shape_layout() {
let shape = new_first_order::<Rational>(2);
assert_eq!(shape.n_components(), 3);
assert_eq!(shape.n_vars(), 2);
assert_eq!(shape.components()[0], vec![0, 0]);
assert_eq!(shape.components()[1], vec![1, 0]);
assert_eq!(shape.components()[2], vec![0, 1]);
assert!(shape.mult_table().is_empty());
}
#[test]
fn variable_and_constant() {
let shape = new_first_order::<Rational>(2);
let x = HyperDual::variable(&shape, 0, r(3, 1));
assert_eq!(x.value(), &r(3, 1));
assert_eq!(x.deriv(0), Some(&r(1, 1)));
assert_eq!(x.deriv(1), Some(&r(0, 1)));
let c = HyperDual::constant(&shape, r(7, 1));
assert_eq!(c.deriv(0), Some(&r(0, 1)));
}
#[test]
fn product_of_two_variables() {
let shape = new_first_order::<Rational>(2);
let x = HyperDual::variable(&shape, 0, r(3, 1));
let y = HyperDual::variable(&shape, 1, r(5, 1));
let f = x * y;
assert_eq!(f.value(), &r(15, 1));
assert_eq!(f.deriv(0), Some(&r(5, 1))); assert_eq!(f.deriv(1), Some(&r(3, 1))); }
#[test]
fn sum_and_difference() {
let shape = new_first_order::<Rational>(2);
let x = HyperDual::variable(&shape, 0, r(2, 1));
let y = HyperDual::variable(&shape, 1, r(7, 1));
let s = x.clone() + y.clone();
assert_eq!(s.value(), &r(9, 1));
assert_eq!(s.deriv(0), Some(&r(1, 1)));
assert_eq!(s.deriv(1), Some(&r(1, 1)));
let d = x - y;
assert_eq!(d.value(), &r(-5, 1));
assert_eq!(d.deriv(1), Some(&r(-1, 1)));
}
#[test]
fn reciprocal_of_constant() {
let shape = new_first_order::<Rational>(1);
let c = HyperDual::constant(&shape, r(4, 1));
let inv = c.inv().unwrap();
assert_eq!(inv.value(), &r(1, 4));
assert_eq!(inv.deriv(0), Some(&r(0, 1)));
}
#[test]
fn quotient_derivatives() {
let shape = new_first_order::<Rational>(2);
let x = HyperDual::variable(&shape, 0, r(6, 1));
let y = HyperDual::variable(&shape, 1, r(3, 1));
let f = x / y;
assert_eq!(f.value(), &r(2, 1));
assert_eq!(f.deriv(0), Some(&r(1, 3)));
assert_eq!(f.deriv(1), Some(&r(-2, 3)));
}
#[test]
fn reciprocal_of_variable_gives_correct_derivative() {
let shape = new_first_order::<Rational>(1);
let x = HyperDual::variable(&shape, 0, r(5, 1));
let inv = x.inv().unwrap();
assert_eq!(inv.value(), &r(1, 5));
assert_eq!(inv.deriv(0), Some(&r(-1, 25)));
}
#[test]
fn power_of_variable() {
let shape = new_first_order::<Rational>(1);
let x = HyperDual::variable(&shape, 0, r(3, 1));
let x2 = x.clone() * x.clone();
let x3 = x2 * x;
assert_eq!(x3.value(), &r(27, 1));
assert_eq!(x3.deriv(0), Some(&r(27, 1)));
}
#[test]
fn three_variable_product_derivatives() {
let shape = new_first_order::<Rational>(3);
let x = HyperDual::variable(&shape, 0, r(2, 1));
let y = HyperDual::variable(&shape, 1, r(3, 1));
let z = HyperDual::variable(&shape, 2, r(5, 1));
let f = x * y * z;
assert_eq!(f.value(), &r(30, 1));
assert_eq!(f.deriv(0), Some(&r(15, 1))); assert_eq!(f.deriv(1), Some(&r(10, 1))); assert_eq!(f.deriv(2), Some(&r(6, 1))); }
#[test]
fn shape_rejects_non_ancestor_closed() {
let components = vec![vec![0], vec![2]];
assert!(DualShape::new(components).is_none());
}
#[test]
fn shape_accepts_second_order() {
let components = vec![vec![0], vec![1], vec![2]];
let shape = DualShape::new(components).unwrap();
assert_eq!(shape.n_components(), 3);
assert_eq!(shape.mult_table(), &[(1usize, 1usize, 2usize)]);
}
}