#[allow(unused_imports)]
use crate::prelude::*;
use num_bigint::BigInt;
use num_rational::BigRational;
use num_traits::{One, Zero};
pub struct SymbolicDifferentiator {
cache: FxHashMap<(PolynomialKey, usize), Vec<BigRational>>,
stats: DifferentiationStats,
}
type PolynomialKey = Vec<String>;
#[derive(Debug, Clone, Default)]
pub struct DifferentiationStats {
pub derivatives_computed: usize,
pub cache_hits: usize,
pub gradients_computed: usize,
pub hessians_computed: usize,
pub jacobians_computed: usize,
}
#[derive(Debug, Clone)]
pub struct Gradient {
pub partials: Vec<Vec<BigRational>>,
}
#[derive(Debug, Clone)]
pub struct Hessian {
pub matrix: Vec<Vec<Vec<BigRational>>>,
}
#[derive(Debug, Clone)]
pub struct Jacobian {
pub matrix: Vec<Vec<Vec<BigRational>>>,
}
impl SymbolicDifferentiator {
pub fn new() -> Self {
Self {
cache: FxHashMap::default(),
stats: DifferentiationStats::default(),
}
}
pub fn derivative(&mut self, poly: &[BigRational]) -> Vec<BigRational> {
self.stats.derivatives_computed += 1;
if poly.len() <= 1 {
return vec![BigRational::zero()];
}
let degree = poly.len() - 1;
let mut deriv = Vec::new();
for (i, coeff) in poly.iter().enumerate().take(degree) {
let power = (degree - i) as i64;
let deriv_coeff = coeff * BigRational::from_integer(BigInt::from(power));
deriv.push(deriv_coeff);
}
if deriv.is_empty() {
vec![BigRational::zero()]
} else {
deriv
}
}
pub fn partial_derivative(
&mut self,
poly: &[BigRational],
var_index: usize,
) -> Vec<BigRational> {
let key = (self.poly_to_key(poly), var_index);
if let Some(cached) = self.cache.get(&key) {
self.stats.cache_hits += 1;
return cached.clone();
}
self.stats.derivatives_computed += 1;
let deriv = self.derivative(poly);
self.cache.insert(key, deriv.clone());
deriv
}
pub fn nth_derivative(&mut self, poly: &[BigRational], n: usize) -> Vec<BigRational> {
let mut result = poly.to_vec();
for _ in 0..n {
result = self.derivative(&result);
if result.len() == 1 && result[0].is_zero() {
break;
}
}
result
}
pub fn gradient(&mut self, poly: &[BigRational], num_vars: usize) -> Gradient {
self.stats.gradients_computed += 1;
let mut partials = Vec::new();
for var_idx in 0..num_vars {
let partial = self.partial_derivative(poly, var_idx);
partials.push(partial);
}
Gradient { partials }
}
pub fn hessian(&mut self, poly: &[BigRational], num_vars: usize) -> Hessian {
self.stats.hessians_computed += 1;
let mut matrix = Vec::new();
for i in 0..num_vars {
let mut row = Vec::new();
let first_deriv = self.partial_derivative(poly, i);
for j in 0..num_vars {
let second_deriv = self.partial_derivative(&first_deriv, j);
row.push(second_deriv);
}
matrix.push(row);
}
Hessian { matrix }
}
pub fn jacobian(&mut self, functions: &[Vec<BigRational>], num_vars: usize) -> Jacobian {
self.stats.jacobians_computed += 1;
let mut matrix = Vec::new();
for func in functions {
let mut row = Vec::new();
for var_idx in 0..num_vars {
let partial = self.partial_derivative(func, var_idx);
row.push(partial);
}
matrix.push(row);
}
Jacobian { matrix }
}
pub fn directional_derivative(
&mut self,
poly: &[BigRational],
direction: &[BigRational],
) -> BigRational {
let num_vars = direction.len();
let grad = self.gradient(poly, num_vars);
let mut result = BigRational::zero();
for (partial, dir_component) in grad.partials.iter().zip(direction.iter()) {
if let Some(constant_term) = partial.last() {
result += constant_term * dir_component;
}
}
result
}
pub fn is_harmonic(&mut self, poly: &[BigRational], num_vars: usize) -> bool {
let hessian = self.hessian(poly, num_vars);
let mut laplacian = vec![BigRational::zero()];
for i in 0..num_vars {
if i < hessian.matrix.len() && i < hessian.matrix[i].len() {
let diagonal_elem = &hessian.matrix[i][i];
laplacian = self.poly_add(&laplacian, diagonal_elem);
}
}
self.is_zero_poly(&laplacian)
}
pub fn is_critical_point(
&mut self,
poly: &[BigRational],
num_vars: usize,
point: &[BigRational],
) -> bool {
let grad = self.gradient(poly, num_vars);
for (i, partial) in grad.partials.iter().enumerate() {
if i < point.len() {
let value = self.evaluate_at_point(partial, point[i].clone());
if !value.is_zero() {
return false;
}
}
}
true
}
pub fn taylor_expansion(
&mut self,
poly: &[BigRational],
center: BigRational,
degree: usize,
) -> Vec<BigRational> {
let mut terms = Vec::new();
for k in 0..=degree {
let kth_deriv = self.nth_derivative(poly, k);
let value_at_center = self.evaluate_at_point(&kth_deriv, center.clone());
let factorial = self.factorial(k);
let term = value_at_center / BigRational::from_integer(factorial);
terms.push(term);
}
terms.reverse(); terms
}
fn poly_add(&self, p1: &[BigRational], p2: &[BigRational]) -> Vec<BigRational> {
let max_len = p1.len().max(p2.len());
let mut result = vec![BigRational::zero(); max_len];
for (i, coeff) in p1.iter().rev().enumerate() {
if i < result.len() {
result[max_len - 1 - i] = result[max_len - 1 - i].clone() + coeff;
}
}
for (i, coeff) in p2.iter().rev().enumerate() {
if i < result.len() {
result[max_len - 1 - i] = result[max_len - 1 - i].clone() + coeff;
}
}
result
}
fn is_zero_poly(&self, poly: &[BigRational]) -> bool {
poly.iter().all(|c| c.is_zero())
}
fn evaluate_at_point(&self, poly: &[BigRational], x: BigRational) -> BigRational {
if poly.is_empty() {
return BigRational::zero();
}
let mut result = poly[0].clone();
for coeff in &poly[1..] {
result = result * &x + coeff;
}
result
}
fn factorial(&self, n: usize) -> BigInt {
if n <= 1 {
BigInt::one()
} else {
let mut result = BigInt::one();
for i in 2..=n {
result *= BigInt::from(i);
}
result
}
}
fn poly_to_key(&self, poly: &[BigRational]) -> PolynomialKey {
poly.iter().map(|c| c.to_string()).collect()
}
pub fn stats(&self) -> &DifferentiationStats {
&self.stats
}
pub fn reset_stats(&mut self) {
self.stats = DifferentiationStats::default();
}
}
impl Default for SymbolicDifferentiator {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_symbolic_differentiator() {
let diff = SymbolicDifferentiator::new();
assert_eq!(diff.stats.derivatives_computed, 0);
}
#[test]
fn test_derivative_constant() {
let mut diff = SymbolicDifferentiator::new();
let poly = vec![BigRational::from_integer(BigInt::from(5))];
let deriv = diff.derivative(&poly);
assert_eq!(deriv.len(), 1);
assert!(deriv[0].is_zero());
}
#[test]
fn test_derivative_linear() {
let mut diff = SymbolicDifferentiator::new();
let poly = vec![
BigRational::from_integer(BigInt::from(3)),
BigRational::from_integer(BigInt::from(2)),
];
let deriv = diff.derivative(&poly);
assert_eq!(deriv.len(), 1);
assert_eq!(deriv[0], BigRational::from_integer(BigInt::from(3)));
}
#[test]
fn test_derivative_quadratic() {
let mut diff = SymbolicDifferentiator::new();
let poly = vec![
BigRational::one(),
BigRational::from_integer(BigInt::from(2)),
BigRational::one(),
];
let deriv = diff.derivative(&poly);
assert_eq!(deriv.len(), 2);
assert_eq!(deriv[0], BigRational::from_integer(BigInt::from(2)));
assert_eq!(deriv[1], BigRational::from_integer(BigInt::from(2)));
}
#[test]
fn test_nth_derivative() {
let mut diff = SymbolicDifferentiator::new();
let poly = vec![
BigRational::one(),
BigRational::zero(),
BigRational::zero(),
BigRational::zero(),
];
let deriv1 = diff.nth_derivative(&poly, 1);
assert_eq!(deriv1[0], BigRational::from_integer(BigInt::from(3)));
let deriv2 = diff.nth_derivative(&poly, 2);
assert_eq!(deriv2[0], BigRational::from_integer(BigInt::from(6)));
let deriv3 = diff.nth_derivative(&poly, 3);
assert_eq!(deriv3[0], BigRational::from_integer(BigInt::from(6)));
let deriv4 = diff.nth_derivative(&poly, 4);
assert!(deriv4[0].is_zero());
}
#[test]
fn test_gradient() {
let mut diff = SymbolicDifferentiator::new();
let poly = vec![
BigRational::one(),
BigRational::from_integer(BigInt::from(2)),
];
let grad = diff.gradient(&poly, 2);
assert_eq!(grad.partials.len(), 2);
}
#[test]
fn test_factorial() {
let diff = SymbolicDifferentiator::new();
assert_eq!(diff.factorial(0), BigInt::one());
assert_eq!(diff.factorial(1), BigInt::one());
assert_eq!(diff.factorial(5), BigInt::from(120));
}
#[test]
fn test_stats() {
let mut diff = SymbolicDifferentiator::new();
let poly = vec![BigRational::one(), BigRational::zero()];
diff.derivative(&poly);
assert_eq!(diff.stats().derivatives_computed, 1);
diff.reset_stats();
assert_eq!(diff.stats().derivatives_computed, 0);
}
}