use crate::api::expr::Ex;
use crate::base::arena::Arena;
use crate::base::node::{ExprId, ExprNode, SymbolId};
use crate::domains::matrix::Matrix;
use num_traits::One;
use num_traits::Signed;
pub struct OdeResult {
pub solution: ExprId,
pub constants: Vec<ExprId>,
}
pub fn dsolve(
arena: &mut Arena,
expr: ExprId,
func: ExprId, var: ExprId, ) -> Option<OdeResult> {
let var_sym = match arena.node(var) {
ExprNode::Symbol(sid) => *sid,
_ => return None,
};
let func_sym = match arena.node(func) {
ExprNode::Symbol(sid) => *sid,
_ => return None,
};
if let Some(result) = try_second_order_const_coeff(arena, expr, func, var, func_sym, var_sym) {
return Some(result);
}
if let Some(result) =
try_second_order_cc_nonhomogeneous(arena, expr, func, var, func_sym, var_sym)
{
return Some(result);
}
if let Some(result) = try_euler_cauchy(arena, expr, func, var, func_sym, var_sym) {
return Some(result);
}
if let Some(result) = try_nth_order_linear_const_coeff(arena, expr, func, var, func_sym) {
return Some(result);
}
if let Some(result) = try_variation_of_parameters(arena, expr, func, var, func_sym, var_sym) {
return Some(result);
}
if let Some(result) = try_clairaut(arena, expr, func, var, func_sym, var_sym) {
return Some(result);
}
if let Some(result) = try_first_order_linear_general(arena, expr, func, var, func_sym, var_sym)
{
return Some(result);
}
if let Some(result) = try_exact_ode(arena, expr, func, var, func_sym, var_sym) {
return Some(result);
}
if let Some(result) = try_integrating_factor_ode(arena, expr, func, var, func_sym, var_sym) {
return Some(result);
}
if let Some(result) = try_bernoulli(arena, expr, func, var, func_sym, var_sym) {
return Some(result);
}
if let Some(result) = try_homogeneous_coefficient(arena, expr, func, var, func_sym, var_sym) {
return Some(result);
}
if let Some(result) = try_nth_order_reducible(arena, expr, func, var, func_sym, var_sym) {
return Some(result);
}
if let Some(result) = try_full_separable(arena, expr, func, var, func_sym, var_sym) {
return Some(result);
}
if let Some(result) = try_simple_separable(arena, expr, func, var, func_sym, var_sym) {
return Some(result);
}
None
}
fn try_simple_separable(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
_var_sym: SymbolId,
) -> Option<OdeResult> {
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
if let ExprNode::Add(ref children) = arena.node(expr).clone() {
let mut has_deriv = false;
let mut _deriv_coeff = None;
let mut other_terms: Vec<ExprId> = Vec::new();
for &child in children {
if child == dy_dx {
has_deriv = true;
_deriv_coeff = Some(arena.one);
} else if !contains_sym(arena, child, func_sym) {
other_terms.push(child);
} else {
return None;
}
}
if has_deriv {
let rhs = if other_terms.is_empty() {
arena.zero
} else {
let sum = arena.add(&other_terms);
arena.neg(sum)
};
let integral = crate::transforms::integrate::integrate(arena, rhs, var);
let c1 = arena.symbol("C1");
let solution = arena.add(&[integral, c1]);
return Some(OdeResult {
solution,
constants: vec![c1],
});
}
}
if expr == dy_dx {
let c1 = arena.symbol("C1");
return Some(OdeResult {
solution: c1,
constants: vec![c1],
});
}
None
}
fn try_second_order_const_coeff(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
_func_sym: SymbolId,
_var_sym: SymbolId,
) -> Option<OdeResult> {
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let d2y_dx2 = arena.intern(ExprNode::Derivative(dy_dx, var));
let children = match arena.node(expr).clone() {
ExprNode::Add(children) => children,
_ => return None,
};
let mut a_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut b_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut c_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
for &child in &children {
let (coeff, term) = arena.as_coeff_term(child);
if term == d2y_dx2 {
a_coeff += coeff;
} else if term == dy_dx {
b_coeff += coeff;
} else if term == func {
c_coeff += coeff;
} else {
return None; }
}
use num_traits::Zero;
if a_coeff.is_zero() {
return None; }
let b = &b_coeff / &a_coeff;
let c = &c_coeff / &a_coeff;
{
let four_r = num_rational::Ratio::<num_bigint::BigInt>::from_integer(4.into());
let disc = &b * &b - &four_r * &c;
if disc.is_negative() {
return build_trig_homogeneous_solution(arena, &b, &disc, var);
}
}
let r_var = arena.symbol("__r");
let two = arena.int(2);
let b_id = {
let nid = arena.intern_num(b.clone());
arena.intern(ExprNode::Num(nid))
};
let c_id = {
let nid = arena.intern_num(c.clone());
arena.intern(ExprNode::Num(nid))
};
let r_var_sq = arena.pow(r_var, two);
let b_r = arena.mul(&[b_id, r_var]);
let char_eq = arena.add(&[r_var_sq, b_r, c_id]);
let roots = crate::transforms::solve::solve(arena, char_eq, r_var);
let c1 = arena.symbol("C1");
let c2 = arena.symbol("C2");
match roots.len() {
2 => {
let r1 = roots[0].value;
let r2 = roots[1].value;
if r1 == r2 {
let rx = arena.mul(&[r1, var]);
let exp_rx = arena.exp(rx);
let c2_x = arena.mul(&[c2, var]);
let inner = arena.add(&[c1, c2_x]);
let solution = arena.mul(&[inner, exp_rx]);
Some(OdeResult {
solution,
constants: vec![c1, c2],
})
} else {
let r1x = arena.mul(&[r1, var]);
let r2x = arena.mul(&[r2, var]);
let exp_r1x = arena.exp(r1x);
let exp_r2x = arena.exp(r2x);
let term1 = arena.mul(&[c1, exp_r1x]);
let term2 = arena.mul(&[c2, exp_r2x]);
let solution = arena.add(&[term1, term2]);
Some(OdeResult {
solution,
constants: vec![c1, c2],
})
}
}
1 => {
let r = roots[0].value;
let rx = arena.mul(&[r, var]);
let exp_rx = arena.exp(rx);
let c2_x = arena.mul(&[c2, var]);
let inner = arena.add(&[c1, c2_x]);
let solution = arena.mul(&[inner, exp_rx]);
Some(OdeResult {
solution,
constants: vec![c1, c2],
})
}
_ => None,
}
}
fn solve_characteristic_equation(
arena: &mut Arena,
b: num_rational::Ratio<num_bigint::BigInt>,
c: num_rational::Ratio<num_bigint::BigInt>,
var: ExprId,
) -> Option<OdeResult> {
{
let four_r = num_rational::Ratio::<num_bigint::BigInt>::from_integer(4.into());
let disc = &b * &b - &four_r * &c;
if disc.is_negative() {
return build_trig_homogeneous_solution(arena, &b, &disc, var);
}
}
let r_var = arena.symbol("__r");
let two = arena.int(2);
let b_id = {
let nid = arena.intern_num(b);
arena.intern(ExprNode::Num(nid))
};
let c_id = {
let nid = arena.intern_num(c);
arena.intern(ExprNode::Num(nid))
};
let r_var_sq = arena.pow(r_var, two);
let b_r = arena.mul(&[b_id, r_var]);
let char_eq = arena.add(&[r_var_sq, b_r, c_id]);
let roots = crate::transforms::solve::solve(arena, char_eq, r_var);
let c1 = arena.symbol("C1");
let c2 = arena.symbol("C2");
match roots.len() {
2 => {
let r1 = roots[0].value;
let r2 = roots[1].value;
if r1 == r2 {
let rx = arena.mul(&[r1, var]);
let exp_rx = arena.exp(rx);
let c2_x = arena.mul(&[c2, var]);
let inner = arena.add(&[c1, c2_x]);
let solution = arena.mul(&[inner, exp_rx]);
Some(OdeResult {
solution,
constants: vec![c1, c2],
})
} else {
let r1x = arena.mul(&[r1, var]);
let r2x = arena.mul(&[r2, var]);
let exp_r1x = arena.exp(r1x);
let exp_r2x = arena.exp(r2x);
let term1 = arena.mul(&[c1, exp_r1x]);
let term2 = arena.mul(&[c2, exp_r2x]);
let solution = arena.add(&[term1, term2]);
Some(OdeResult {
solution,
constants: vec![c1, c2],
})
}
}
1 => {
let r = roots[0].value;
let rx = arena.mul(&[r, var]);
let exp_rx = arena.exp(rx);
let c2_x = arena.mul(&[c2, var]);
let inner = arena.add(&[c1, c2_x]);
let solution = arena.mul(&[inner, exp_rx]);
Some(OdeResult {
solution,
constants: vec![c1, c2],
})
}
_ => None,
}
}
fn try_second_order_cc_nonhomogeneous(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
_var_sym: SymbolId,
) -> Option<OdeResult> {
use num_traits::Zero;
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let d2y_dx2 = arena.intern(ExprNode::Derivative(dy_dx, var));
let children = match arena.node(expr).clone() {
ExprNode::Add(children) => children,
_ => return None,
};
let mut a_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut b_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut c_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut f_of_x_terms: Vec<ExprId> = Vec::new();
for &child in &children {
let (coeff, term) = arena.as_coeff_term(child);
if term == d2y_dx2 {
a_coeff += coeff;
} else if term == dy_dx {
b_coeff += coeff;
} else if term == func {
c_coeff += coeff;
} else if !contains_sym(arena, child, func_sym) {
f_of_x_terms.push(child);
} else {
return None; }
}
if a_coeff.is_zero() {
return None; }
if f_of_x_terms.is_empty() {
return None; }
let b = &b_coeff / &a_coeff;
let c = &c_coeff / &a_coeff;
let f_sum = if f_of_x_terms.len() == 1 {
f_of_x_terms[0]
} else {
arena.add(&f_of_x_terms)
};
let y_p_poly = if let Some(f_coeffs_expr) = arena.coefficients_of(f_sum, var) {
if f_coeffs_expr.is_empty() {
None
} else {
let mut rhs_coeffs: Vec<num_rational::Ratio<num_bigint::BigInt>> = Vec::new();
let mut all_numeric = true;
for &cid in &f_coeffs_expr {
if let Some(val) = arena.as_num(cid) {
rhs_coeffs.push(-val.clone() / &a_coeff);
} else {
all_numeric = false;
break;
}
}
if all_numeric {
find_particular_polynomial(&b, &c, &rhs_coeffs)
.map(|pc| build_polynomial_expr(arena, &pc, var))
} else {
None
}
}
} else {
None
};
let y_p = if let Some(yp) = y_p_poly {
yp
} else {
let neg_f_sum = arena.neg(f_sum);
let a_id = ode_ratio_to_expr(arena, &a_coeff);
let rhs_expr = arena.div(neg_f_sum, a_id);
let rhs_expr = crate::transforms::eval::eval(arena, rhs_expr);
try_undetermined_trig_exp(arena, rhs_expr, &b, &c, var)?
};
let homo_result = solve_characteristic_equation(arena, b, c, var)?;
let solution = arena.add(&[homo_result.solution, y_p]);
Some(OdeResult {
solution,
constants: homo_result.constants,
})
}
fn find_particular_polynomial(
b: &num_rational::Ratio<num_bigint::BigInt>,
c: &num_rational::Ratio<num_bigint::BigInt>,
rhs_coeffs: &[num_rational::Ratio<num_bigint::BigInt>],
) -> Option<Vec<num_rational::Ratio<num_bigint::BigInt>>> {
use num_bigint::BigInt;
use num_rational::Ratio;
use num_traits::Zero;
let n = rhs_coeffs.len() - 1;
if !c.is_zero() {
let mut a = vec![Ratio::<BigInt>::zero(); n + 1];
for j in (0..=n).rev() {
let a_j2 = if j + 2 <= n {
a[j + 2].clone()
} else {
Ratio::zero()
};
let a_j1 = if j < n {
a[j + 1].clone()
} else {
Ratio::zero()
};
let factor2 = Ratio::from_integer(BigInt::from(((j + 2) * (j + 1)) as i64)) * a_j2;
let factor1 = Ratio::from_integer(BigInt::from((j + 1) as i64)) * b.clone() * a_j1;
a[j] = (rhs_coeffs[j].clone() - factor2 - factor1) / c.clone();
}
Some(a)
} else if !b.is_zero() {
let mut bb = vec![Ratio::<BigInt>::zero(); n + 1];
for j in (0..=n).rev() {
let deriv_term = if j < n {
Ratio::from_integer(BigInt::from(((j + 2) * (j + 1)) as i64)) * bb[j + 1].clone()
} else {
Ratio::zero()
};
let denom = Ratio::from_integer(BigInt::from((j + 1) as i64)) * b.clone();
if denom.is_zero() {
return None;
}
bb[j] = (rhs_coeffs[j].clone() - deriv_term) / denom;
}
let mut result = vec![Ratio::zero()];
result.extend(bb);
Some(result)
} else {
let mut result = vec![Ratio::<BigInt>::zero(); 2];
for (j, r_j) in rhs_coeffs.iter().enumerate() {
let denom = Ratio::from_integer(BigInt::from(((j + 1) * (j + 2)) as i64));
result.push(r_j.clone() / denom);
}
Some(result)
}
}
fn build_polynomial_expr(
arena: &mut Arena,
coeffs: &[num_rational::Ratio<num_bigint::BigInt>],
var: ExprId,
) -> ExprId {
use num_traits::Zero;
let mut terms = Vec::new();
for (j, coeff) in coeffs.iter().enumerate() {
if coeff.is_zero() {
continue;
}
let x_pow_j = if j == 0 {
arena.one
} else if j == 1 {
var
} else {
let exp = arena.int(j as i64);
arena.pow(var, exp)
};
let term = arena.make_coeff_term(coeff.clone(), x_pow_j);
terms.push(term);
}
if terms.is_empty() {
arena.zero
} else if terms.len() == 1 {
terms[0]
} else {
arena.add(&terms)
}
}
fn try_first_order_linear(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
_var_sym: SymbolId,
) -> Option<OdeResult> {
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let children = match arena.node(expr).clone() {
ExprNode::Add(children) => children,
_ => return None,
};
let mut has_dy = false;
let mut a_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut f_of_x = Vec::new();
for &child in &children {
let (coeff, term) = arena.as_coeff_term(child);
if term == dy_dx {
has_dy = true;
if !coeff.is_one() {
return None; }
} else if term == func {
a_coeff += coeff;
} else if !contains_sym(arena, child, func_sym) {
f_of_x.push(child);
} else {
return None; }
}
use num_traits::Zero;
if !has_dy {
return None;
}
let a_id = {
let nid = arena.intern_num(a_coeff.clone());
arena.intern(ExprNode::Num(nid))
};
let neg_a = arena.neg(a_id);
let ax = arena.mul(&[a_id, var]);
let neg_ax = arena.mul(&[neg_a, var]);
let exp_ax = arena.exp(ax);
let exp_neg_ax = arena.exp(neg_ax);
if a_coeff.is_zero() {
let f = if f_of_x.is_empty() {
arena.zero
} else {
let sum = arena.add(&f_of_x);
arena.neg(sum)
};
let integral = crate::transforms::integrate::integrate(arena, f, var);
let c1 = arena.symbol("C1");
let solution = arena.add(&[integral, c1]);
return Some(OdeResult {
solution,
constants: vec![c1],
});
}
let neg_f = if f_of_x.is_empty() {
arena.zero
} else {
let sum = arena.add(&f_of_x);
arena.neg(sum)
};
let integrand = if let ExprNode::Exp(neg_f_inner) = arena.node(neg_f).clone() {
let combined_arg = arena.add(&[neg_f_inner, ax]);
let combined_arg = crate::transforms::eval::eval(arena, combined_arg);
arena.exp(combined_arg)
} else {
arena.mul(&[neg_f, exp_ax])
};
let integrand = crate::transforms::eval::eval(arena, integrand);
let integral = crate::transforms::integrate::integrate(arena, integrand, var);
let c1 = arena.symbol("C1");
let inner = arena.add(&[integral, c1]);
let solution = arena.mul(&[exp_neg_ax, inner]);
Some(OdeResult {
solution,
constants: vec![c1],
})
}
fn try_full_separable(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
var_sym: SymbolId,
) -> Option<OdeResult> {
tracing::debug!("ode: trying full separable");
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let children = match arena.node(expr).clone() {
ExprNode::Add(children) => children,
_ => return None,
};
let mut has_dy = false;
let mut dy_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut other_terms: Vec<ExprId> = Vec::new();
for &child in &children {
let (coeff, term) = arena.as_coeff_term(child);
if term == dy_dx {
has_dy = true;
dy_coeff += coeff;
} else {
other_terms.push(child);
}
}
use num_traits::Zero;
if !has_dy || dy_coeff.is_zero() {
return None;
}
if !dy_coeff.is_one() {
let inv_coeff = num_rational::Ratio::<num_bigint::BigInt>::one() / &dy_coeff;
let inv_id = {
let nid = arena.intern_num(inv_coeff);
arena.intern(ExprNode::Num(nid))
};
let mut scaled = Vec::new();
for &t in &other_terms {
scaled.push(arena.mul(&[inv_id, t]));
}
other_terms = scaled;
}
let rhs = if other_terms.is_empty() {
return None; } else {
let sum = arena.add(&other_terms);
arena.neg(sum)
};
if !contains_sym(arena, rhs, func_sym) {
return None; }
let factors = collect_mul_factors(arena, rhs);
let mut x_factors: Vec<ExprId> = Vec::new();
let mut y_factors: Vec<ExprId> = Vec::new();
for &factor in &factors {
let has_x = contains_sym(arena, factor, var_sym);
let has_y = contains_sym(arena, factor, func_sym);
if has_x && has_y {
return None;
} else if has_y {
y_factors.push(factor);
} else {
x_factors.push(factor);
}
}
if y_factors.is_empty() {
return None; }
let f_x = if x_factors.is_empty() {
arena.one
} else if x_factors.len() == 1 {
x_factors[0]
} else {
arena.mul(&x_factors)
};
let g_y = if y_factors.len() == 1 {
y_factors[0]
} else {
arena.mul(&y_factors)
};
if g_y == func {
let integral_fx = crate::transforms::integrate::integrate(arena, f_x, var);
let c1 = arena.symbol("C1");
let exponent = arena.add(&[integral_fx, c1]);
let solution = arena.exp(exponent);
return Some(OdeResult {
solution,
constants: vec![c1],
});
}
{
let (coeff, base) = arena.as_coeff_term(g_y);
if base == func {
let coeff_id = {
let nid = arena.intern_num(coeff);
arena.intern(ExprNode::Num(nid))
};
let scaled_fx = arena.mul(&[coeff_id, f_x]);
let integral_fx = crate::transforms::integrate::integrate(arena, scaled_fx, var);
let c1 = arena.symbol("C1");
let exponent = arena.add(&[integral_fx, c1]);
let solution = arena.exp(exponent);
return Some(OdeResult {
solution,
constants: vec![c1],
});
}
}
let neg_one = arena.int(-1);
let inv_gy = arena.pow(g_y, neg_one);
let lhs_integral = arena.intern(ExprNode::Integral(inv_gy, func));
let rhs_integral = crate::transforms::integrate::integrate(arena, f_x, var);
let c1 = arena.symbol("C1");
let neg_rhs = arena.neg(rhs_integral);
let neg_c1 = arena.neg(c1);
let solution = arena.add(&[lhs_integral, neg_rhs, neg_c1]);
Some(OdeResult {
solution,
constants: vec![c1],
})
}
fn collect_mul_factors(arena: &mut Arena, expr: ExprId) -> Vec<ExprId> {
match arena.node(expr).clone() {
ExprNode::Mul(children) => children.to_vec(),
ExprNode::Neg(inner) => {
let neg_one = arena.int(-1);
let mut factors = vec![neg_one];
factors.extend(collect_mul_factors(arena, inner));
factors
}
_ => vec![expr],
}
}
fn try_first_order_linear_general(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
var_sym: SymbolId,
) -> Option<OdeResult> {
tracing::debug!("ode: trying variable-coefficient first-order linear");
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let children = match arena.node(expr).clone() {
ExprNode::Add(children) => children,
_ => return None,
};
let mut has_dy = false;
let mut dy_coeff_rational = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut y_terms: Vec<ExprId> = Vec::new(); let mut free_terms: Vec<ExprId> = Vec::new();
for &child in &children {
let (coeff, term) = arena.as_coeff_term(child);
if term == dy_dx {
has_dy = true;
dy_coeff_rational += coeff;
} else if contains_sym(arena, child, func_sym) {
y_terms.push(child);
} else {
free_terms.push(child);
}
}
use num_traits::Zero;
if !has_dy || dy_coeff_rational.is_zero() {
return None;
}
if y_terms.is_empty() {
return None;
}
let mut p_x_terms: Vec<ExprId> = Vec::new();
for &yt in &y_terms {
{
let px = extract_coeff_of_func(arena, yt, func, func_sym, var_sym)?;
p_x_terms.push(px);
}
}
if !dy_coeff_rational.is_one() {
let inv_coeff = num_rational::Ratio::<num_bigint::BigInt>::one() / &dy_coeff_rational;
let inv_id = {
let nid = arena.intern_num(inv_coeff);
arena.intern(ExprNode::Num(nid))
};
let mut scaled_p = Vec::new();
for &p in &p_x_terms {
scaled_p.push(arena.mul(&[inv_id, p]));
}
p_x_terms = scaled_p;
let mut scaled_f = Vec::new();
for &f in &free_terms {
scaled_f.push(arena.mul(&[inv_id, f]));
}
free_terms = scaled_f;
}
let p_x = if p_x_terms.is_empty() {
arena.zero
} else if p_x_terms.len() == 1 {
p_x_terms[0]
} else {
arena.add(&p_x_terms)
};
if contains_sym(arena, p_x, func_sym) {
return None;
}
if let Some(_num_val) = arena.as_num(p_x) {
return try_first_order_linear(arena, expr, func, var, func_sym, var_sym);
}
let q_x = if free_terms.is_empty() {
arena.zero
} else {
let sum = arena.add(&free_terms);
arena.neg(sum)
};
let int_px = crate::transforms::integrate::integrate(arena, p_x, var);
if let ExprNode::Integral(_, _) = arena.node(int_px).clone() {
return None;
}
let mu = exp_of_log_sum(arena, int_px);
let neg_one = arena.int(-1);
let inv_mu = arena.pow(mu, neg_one);
let c1 = arena.symbol("C1");
if arena.is_zero_structural(q_x) {
let solution = arena.mul(&[c1, inv_mu]);
return Some(OdeResult {
solution,
constants: vec![c1],
});
}
let integrand = arena.mul(&[q_x, mu]);
let integrand = crate::transforms::eval::eval(arena, integrand);
let integral = crate::transforms::integrate::integrate(arena, integrand, var);
let inner = arena.add(&[integral, c1]);
let solution = arena.mul(&[inv_mu, inner]);
Some(OdeResult {
solution,
constants: vec![c1],
})
}
fn extract_coeff_of_func(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
func_sym: SymbolId,
_var_sym: SymbolId,
) -> Option<ExprId> {
if expr == func {
return Some(arena.one);
}
if let ExprNode::Neg(inner) = arena.node(expr).clone() {
if inner == func {
let neg_one = arena.int(-1);
return Some(neg_one);
}
if let Some(inner_coeff) = extract_coeff_of_func(arena, inner, func, func_sym, _var_sym) {
let result = arena.neg(inner_coeff);
return Some(result);
}
return None;
}
if let ExprNode::Mul(ref children) = arena.node(expr).clone() {
let mut found_y = false;
let mut other_factors: Vec<ExprId> = Vec::new();
let mut y_count = 0;
for &child in children {
if child == func {
y_count += 1;
if y_count > 1 {
return None; }
found_y = true;
} else if contains_sym(arena, child, func_sym) {
return None;
} else {
other_factors.push(child);
}
}
if found_y {
let coeff = if other_factors.is_empty() {
arena.one
} else if other_factors.len() == 1 {
other_factors[0]
} else {
arena.mul(&other_factors)
};
return Some(coeff);
}
}
{
let (coeff, term) = arena.as_coeff_term(expr);
if !coeff.is_one()
&& term != expr
&& let Some(inner_coeff) = extract_coeff_of_func(arena, term, func, func_sym, _var_sym)
{
let coeff_id = {
let nid = arena.intern_num(coeff);
arena.intern(ExprNode::Num(nid))
};
let result = arena.mul(&[coeff_id, inner_coeff]);
return Some(result);
}
}
None
}
fn extract_m_n(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
) -> Option<(ExprId, ExprId)> {
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
if expr == dy_dx {
return Some((arena.zero, arena.one));
}
let children = match arena.node(expr).clone() {
ExprNode::Add(c) => c,
_ => return None,
};
let mut m_terms: Vec<ExprId> = Vec::new();
let mut n_terms: Vec<ExprId> = Vec::new();
for &child in &children {
let (coeff, term) = arena.as_coeff_term(child);
if term == dy_dx {
let coeff_id = {
let nid = arena.intern_num(coeff);
arena.intern(ExprNode::Num(nid))
};
n_terms.push(coeff_id);
} else if expr_contains(arena, child, dy_dx) {
if let ExprNode::Mul(ref mul_children) = arena.node(child).clone() {
let mut found_dy = false;
let mut other_factors: Vec<ExprId> = Vec::new();
for &mc in mul_children.iter() {
if mc == dy_dx && !found_dy {
found_dy = true;
} else {
other_factors.push(mc);
}
}
if found_dy {
let n_factor = match other_factors.len() {
0 => arena.one,
1 => other_factors[0],
_ => arena.mul(&other_factors),
};
n_terms.push(n_factor);
} else {
return None; }
} else {
return None;
}
} else {
m_terms.push(child);
}
}
if n_terms.is_empty() {
return None; }
let m_expr = match m_terms.len() {
0 => arena.zero,
1 => m_terms[0],
_ => arena.add(&m_terms),
};
let n_expr = match n_terms.len() {
1 => n_terms[0],
_ => arena.add(&n_terms),
};
Some((m_expr, n_expr))
}
fn try_exact_ode(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
var_sym: SymbolId,
) -> Option<OdeResult> {
tracing::debug!("ode: trying exact ODE");
let (m_expr, n_expr) = extract_m_n(arena, expr, func, var)?;
if !contains_sym(arena, m_expr, func_sym) && !contains_sym(arena, n_expr, func_sym) {
return None;
}
let dm_dy = crate::transforms::diff::diff(arena, m_expr, func);
let dn_dx = crate::transforms::diff::diff(arena, n_expr, var);
let diff_check = arena.sub(dm_dy, dn_dx);
let diff_eval = crate::transforms::eval::eval(arena, diff_check);
let diff_expanded = crate::transforms::expand::expand(arena, diff_eval);
let diff_simplified = crate::transforms::eval::eval(arena, diff_expanded);
if diff_simplified != arena.zero {
return None; }
let integral_m = crate::transforms::integrate::integrate(arena, m_expr, var);
if matches!(arena.node(integral_m), ExprNode::Integral(_, _)) {
return None; }
let d_intm_dy = crate::transforms::diff::diff(arena, integral_m, func);
let g_prime = arena.sub(n_expr, d_intm_dy);
let g_prime = crate::transforms::eval::eval(arena, g_prime);
let g_prime = crate::transforms::expand::expand(arena, g_prime);
let g_prime = crate::transforms::eval::eval(arena, g_prime);
if contains_sym(arena, g_prime, var_sym) {
return None;
}
let g_y = crate::transforms::integrate::integrate(arena, g_prime, func);
if matches!(arena.node(g_y), ExprNode::Integral(_, _)) {
return None;
}
let potential = arena.add(&[integral_m, g_y]);
let potential = crate::transforms::eval::eval(arena, potential);
let c1 = arena.symbol("C1");
let f_minus_c1 = arena.sub(potential, c1);
let solutions = crate::transforms::solve::solve(arena, f_minus_c1, func);
if solutions.len() == 1 {
return Some(OdeResult {
solution: solutions[0].value,
constants: vec![c1],
});
}
Some(OdeResult {
solution: potential,
constants: vec![c1],
})
}
fn try_integrating_factor_ode(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
var_sym: SymbolId,
) -> Option<OdeResult> {
tracing::debug!("ode: trying integrating factor for non-exact ODE");
let (m_expr, n_expr) = extract_m_n(arena, expr, func, var)?;
if !contains_sym(arena, m_expr, func_sym) && !contains_sym(arena, n_expr, func_sym) {
return None;
}
let dm_dy = crate::transforms::diff::diff(arena, m_expr, func);
let dn_dx = crate::transforms::diff::diff(arena, n_expr, var);
let diff_mn = arena.sub(dm_dy, dn_dx); let diff_eval = crate::transforms::eval::eval(arena, diff_mn);
let diff_expanded = crate::transforms::expand::expand(arena, diff_eval);
let diff_simplified = crate::transforms::eval::eval(arena, diff_expanded);
if diff_simplified == arena.zero {
return try_exact_ode(arena, expr, func, var, func_sym, var_sym);
}
{
let ratio = arena.div(diff_simplified, n_expr);
let ratio = crate::transforms::eval::eval(arena, ratio);
let ratio = crate::transforms::expand::expand(arena, ratio);
let ratio = crate::transforms::eval::eval(arena, ratio);
let ratio_cancelled = arena.cancel_expr(ratio, var);
if !contains_sym(arena, ratio_cancelled, func_sym) {
let int_ratio = crate::transforms::integrate::integrate(arena, ratio_cancelled, var);
if !matches!(arena.node(int_ratio), ExprNode::Integral(_, _)) {
let mu = exp_of_log_sum(arena, int_ratio);
let new_m = arena.mul(&[mu, m_expr]);
let new_n = arena.mul(&[mu, n_expr]);
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let n_dy = arena.mul(&[new_n, dy_dx]);
let new_expr = arena.add(&[new_m, n_dy]);
let new_expr = crate::transforms::eval::eval(arena, new_expr);
if let Some(result) = try_exact_ode(arena, new_expr, func, var, func_sym, var_sym) {
return Some(result);
}
}
}
}
{
let neg_diff = arena.neg(diff_simplified); let ratio = arena.div(neg_diff, m_expr);
let ratio = crate::transforms::eval::eval(arena, ratio);
let ratio = crate::transforms::expand::expand(arena, ratio);
let ratio = crate::transforms::eval::eval(arena, ratio);
let ratio_cancelled = arena.cancel_expr(ratio, func);
if !contains_sym(arena, ratio_cancelled, var_sym) {
let int_ratio = crate::transforms::integrate::integrate(arena, ratio_cancelled, func);
if !matches!(arena.node(int_ratio), ExprNode::Integral(_, _)) {
let mu = exp_of_log_sum(arena, int_ratio);
let new_m = arena.mul(&[mu, m_expr]);
let new_n = arena.mul(&[mu, n_expr]);
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let n_dy = arena.mul(&[new_n, dy_dx]);
let new_expr = arena.add(&[new_m, n_dy]);
let new_expr = crate::transforms::eval::eval(arena, new_expr);
if let Some(result) = try_exact_ode(arena, new_expr, func, var, func_sym, var_sym) {
return Some(result);
}
}
}
}
None
}
fn try_homogeneous_coefficient(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
var_sym: SymbolId,
) -> Option<OdeResult> {
tracing::debug!("ode: trying homogeneous coefficient");
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let d2y_dx2 = arena.intern(ExprNode::Derivative(dy_dx, var));
if expr_contains(arena, expr, d2y_dx2) {
return None;
}
let children = match arena.node(expr).clone() {
ExprNode::Add(children) => children,
_ => return None,
};
let mut has_dy = false;
let mut dy_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut other_terms: Vec<ExprId> = Vec::new();
for &child in &children {
let (coeff, term) = arena.as_coeff_term(child);
if term == dy_dx {
has_dy = true;
dy_coeff += coeff;
} else {
other_terms.push(child);
}
}
use num_traits::Zero;
if !has_dy || dy_coeff.is_zero() {
return None;
}
let rhs = if other_terms.is_empty() {
return None;
} else {
let sum = arena.add(&other_terms);
arena.neg(sum)
};
let rhs = if !dy_coeff.is_one() {
let inv_id = ode_ratio_to_expr(
arena,
&(num_rational::Ratio::<num_bigint::BigInt>::one() / &dy_coeff),
);
let s = arena.mul(&[inv_id, rhs]);
crate::transforms::eval::eval(arena, s)
} else {
rhs
};
if !contains_sym(arena, rhs, func_sym) || !contains_sym(arena, rhs, var_sym) {
return None;
}
let v = arena.symbol("__v");
let rhs_sub = crate::transforms::subs::subs(arena, rhs, func, v);
let rhs_sub = crate::transforms::subs::subs(arena, rhs_sub, var, arena.one);
let rhs_sub = crate::transforms::eval::eval(arena, rhs_sub);
let rhs_sub = crate::transforms::expand::expand(arena, rhs_sub);
let rhs_sub = crate::transforms::eval::eval(arena, rhs_sub);
if contains_sym(arena, rhs_sub, var_sym) {
return None;
}
{
let vx = arena.mul(&[v, var]);
let probe = crate::transforms::subs::subs(arena, rhs, func, vx);
let probe = crate::transforms::eval::eval(arena, probe);
let probe = crate::transforms::expand::expand(arena, probe);
let probe = crate::transforms::eval::eval(arena, probe);
let probe = arena.cancel_expr(probe, var);
let probe = crate::transforms::eval::eval(arena, probe);
if contains_sym(arena, probe, var_sym) {
return None;
}
}
let f_v_minus_v = arena.sub(rhs_sub, v);
let f_v_minus_v = crate::transforms::eval::eval(arena, f_v_minus_v);
if f_v_minus_v == arena.zero {
let c1 = arena.symbol("C1");
let solution = arena.mul(&[c1, var]);
return Some(OdeResult {
solution,
constants: vec![c1],
});
}
let neg_one = arena.int(-1);
let inv_fv = arena.pow(f_v_minus_v, neg_one);
let lhs_integral = crate::transforms::integrate::integrate(arena, inv_fv, v);
if matches!(arena.node(lhs_integral), ExprNode::Integral(_, _)) {
return None;
}
let abs_x = arena.abs(var);
let ln_abs_x = arena.ln(abs_x);
let c1 = arena.symbol("C1");
let rhs_eq = arena.add(&[ln_abs_x, c1]);
let y_over_x = arena.div(func, var);
let lhs_backsub = crate::transforms::subs::subs(arena, lhs_integral, v, y_over_x);
let lhs_backsub = crate::transforms::eval::eval(arena, lhs_backsub);
let implicit = arena.sub(lhs_backsub, rhs_eq);
let implicit = crate::transforms::eval::eval(arena, implicit);
let solutions = crate::transforms::solve::solve(arena, implicit, func);
if solutions.len() == 1 {
let sol = crate::transforms::eval::eval(arena, solutions[0].value);
return Some(OdeResult {
solution: sol,
constants: vec![c1],
});
}
Some(OdeResult {
solution: implicit,
constants: vec![c1],
})
}
fn try_nth_order_reducible(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
_func_sym: SymbolId,
var_sym: SymbolId,
) -> Option<OdeResult> {
tracing::debug!("ode: trying nth-order reducible (missing x)");
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let d2y_dx2 = arena.intern(ExprNode::Derivative(dy_dx, var));
if !expr_contains(arena, expr, d2y_dx2) {
return None;
}
let d2_placeholder = arena.symbol("__d2");
let d1_placeholder = arena.symbol("__d1");
let stripped = crate::transforms::subs::subs(arena, expr, d2y_dx2, d2_placeholder);
let stripped = crate::transforms::subs::subs(arena, stripped, dy_dx, d1_placeholder);
if contains_sym(arena, stripped, var_sym) {
return None; }
let p = arena.symbol("__p");
let dp_dy = arena.intern(ExprNode::Derivative(p, func));
let p_dp_dy = arena.mul(&[p, dp_dy]);
let reduced = crate::transforms::subs::subs(arena, expr, d2y_dx2, p_dp_dy);
let reduced = crate::transforms::subs::subs(arena, reduced, dy_dx, p);
let reduced = crate::transforms::eval::eval(arena, reduced);
let p_result = dsolve(arena, reduced, p, func)?;
let c2 = arena.symbol("C2");
let p_sol = p_result.solution;
let neg_one_id = arena.int(-1);
let inv_p = arena.pow(p_sol, neg_one_id);
let inv_p = crate::transforms::eval::eval(arena, inv_p);
let lhs_integral = crate::transforms::integrate::integrate(arena, inv_p, func);
if matches!(arena.node(lhs_integral), ExprNode::Integral(_, _)) {
return None; }
let rhs = arena.add(&[var, c2]);
let implicit = arena.sub(lhs_integral, rhs);
let implicit = crate::transforms::eval::eval(arena, implicit);
let solutions = crate::transforms::solve::solve(arena, implicit, func);
if solutions.len() == 1 {
let sol = crate::transforms::eval::eval(arena, solutions[0].value);
let mut constants = p_result.constants;
constants.push(c2);
return Some(OdeResult {
solution: sol,
constants,
});
}
let mut constants = p_result.constants;
constants.push(c2);
Some(OdeResult {
solution: implicit,
constants,
})
}
fn contains_sym(arena: &Arena, expr: ExprId, sym: SymbolId) -> bool {
let mut stack = vec![expr];
while let Some(id) = stack.pop() {
match arena.node(id) {
ExprNode::Symbol(s) => {
if *s == sym {
return true;
}
}
other => other.for_each_child(|c| stack.push(c)),
}
}
false
}
fn exp_of_log_sum(arena: &mut Arena, integral: ExprId) -> ExprId {
let terms: Vec<ExprId> = match arena.node(integral).clone() {
ExprNode::Add(c) => c.to_vec(),
_ => vec![integral],
};
let mut factors: Vec<ExprId> = Vec::new();
let mut leftover: Vec<ExprId> = Vec::new();
for t in terms {
let (coeff, term) = arena.as_coeff_term(t);
let inner = match arena.node(term).clone() {
ExprNode::Ln(inner) => Some(inner),
_ => None,
};
match inner {
Some(inner) => {
let base = match arena.node(inner).clone() {
ExprNode::Abs(a) => a,
_ => inner,
};
if coeff.is_one() {
factors.push(base);
} else {
let c_id = ode_ratio_to_expr(arena, &coeff);
factors.push(arena.pow(base, c_id));
}
}
None => leftover.push(t),
}
}
if factors.is_empty() {
return arena.exp(integral);
}
if !leftover.is_empty() {
let rest = arena.add(&leftover);
factors.push(arena.exp(rest));
}
let prod = arena.mul(&factors);
crate::transforms::eval::eval(arena, prod)
}
fn ode_ratio_to_expr(arena: &mut Arena, r: &num_rational::Ratio<num_bigint::BigInt>) -> ExprId {
let nid = arena.intern_num(r.clone());
arena.intern(ExprNode::Num(nid))
}
fn build_trig_homogeneous_solution(
arena: &mut Arena,
b: &num_rational::Ratio<num_bigint::BigInt>,
disc: &num_rational::Ratio<num_bigint::BigInt>,
var: ExprId,
) -> Option<OdeResult> {
use num_traits::Zero;
let c1 = arena.symbol("C1");
let c2 = arena.symbol("C2");
let two_r = num_rational::Ratio::<num_bigint::BigInt>::from_integer(2.into());
let alpha = -(b.clone()) / &two_r;
let neg_disc = -(disc.clone());
let neg_disc_id = ode_ratio_to_expr(arena, &neg_disc);
let half = arena.rational(1, 2);
let sqrt_neg_disc = arena.pow(neg_disc_id, half);
let two_id = arena.int(2);
let beta_expr = arena.div(sqrt_neg_disc, two_id);
let beta_expr = crate::transforms::eval::eval(arena, beta_expr);
let beta_x = arena.mul(&[beta_expr, var]);
let cos_bx = arena.cos(beta_x);
let sin_bx = arena.sin(beta_x);
let c1_cos = arena.mul(&[c1, cos_bx]);
let c2_sin = arena.mul(&[c2, sin_bx]);
let trig_part = arena.add(&[c1_cos, c2_sin]);
let solution = if alpha.is_zero() {
trig_part
} else {
let alpha_id = ode_ratio_to_expr(arena, &alpha);
let alpha_x = arena.mul(&[alpha_id, var]);
let exp_ax = arena.exp(alpha_x);
arena.mul(&[exp_ax, trig_part])
};
Some(OdeResult {
solution,
constants: vec![c1, c2],
})
}
fn try_undetermined_trig_exp(
arena: &mut Arena,
rhs: ExprId,
b: &num_rational::Ratio<num_bigint::BigInt>,
c: &num_rational::Ratio<num_bigint::BigInt>,
var: ExprId,
) -> Option<ExprId> {
use num_traits::Zero;
let (coeff_r, term) = arena.as_coeff_term(rhs);
let node = arena.node(term).clone();
match node {
ExprNode::Sin(inner) => {
let (omega, constant) = extract_linear_numeric(arena, inner, var)?;
if !constant.is_zero() {
return None;
}
let zero_r = num_rational::Ratio::<num_bigint::BigInt>::zero();
try_trig_particular(arena, &coeff_r, &zero_r, &omega, b, c, var)
}
ExprNode::Cos(inner) => {
let (omega, constant) = extract_linear_numeric(arena, inner, var)?;
if !constant.is_zero() {
return None;
}
let zero_r = num_rational::Ratio::<num_bigint::BigInt>::zero();
try_trig_particular(arena, &zero_r, &coeff_r, &omega, b, c, var)
}
ExprNode::Exp(inner) => {
let (r_val, constant) = extract_linear_numeric(arena, inner, var)?;
if !constant.is_zero() {
return None;
}
try_exp_particular(arena, &coeff_r, &r_val, b, c, var)
}
_ => None,
}
}
fn extract_linear_numeric(
arena: &Arena,
expr: ExprId,
var: ExprId,
) -> Option<(
num_rational::Ratio<num_bigint::BigInt>,
num_rational::Ratio<num_bigint::BigInt>,
)> {
let poly = crate::poly::polybridge::expr_to_poly(arena, expr, var)?;
if poly.degree()? != 1 {
return None;
}
Some((poly.coeff(1), poly.coeff(0)))
}
fn try_trig_particular(
arena: &mut Arena,
p: &num_rational::Ratio<num_bigint::BigInt>,
q: &num_rational::Ratio<num_bigint::BigInt>,
omega: &num_rational::Ratio<num_bigint::BigInt>,
b: &num_rational::Ratio<num_bigint::BigInt>,
c: &num_rational::Ratio<num_bigint::BigInt>,
var: ExprId,
) -> Option<ExprId> {
use num_traits::Zero;
let omega_sq = omega * omega;
let d = c - &omega_sq; let bw = b * omega;
let det = &d * &d + &bw * &bw;
if !det.is_zero() {
let alpha = (&d * p + &bw * q) / &det;
let beta = (&d * q - &bw * p) / &det;
let omega_id = ode_ratio_to_expr(arena, omega);
let omega_x = arena.mul(&[omega_id, var]);
let mut terms = Vec::new();
if !alpha.is_zero() {
let alpha_id = ode_ratio_to_expr(arena, &alpha);
let sin_wx = arena.sin(omega_x);
terms.push(arena.mul(&[alpha_id, sin_wx]));
}
if !beta.is_zero() {
let beta_id = ode_ratio_to_expr(arena, &beta);
let cos_wx = arena.cos(omega_x);
terms.push(arena.mul(&[beta_id, cos_wx]));
}
match terms.len() {
0 => Some(arena.zero),
1 => Some(terms[0]),
_ => Some(arena.add(&terms)),
}
} else {
if omega.is_zero() {
return None;
}
let two_omega = num_rational::Ratio::from_integer(num_bigint::BigInt::from(2)) * omega;
let alpha = q / &two_omega;
let beta = -(p / &two_omega);
let omega_id = ode_ratio_to_expr(arena, omega);
let omega_x = arena.mul(&[omega_id, var]);
let mut inner_terms = Vec::new();
if !alpha.is_zero() {
let alpha_id = ode_ratio_to_expr(arena, &alpha);
let sin_wx = arena.sin(omega_x);
inner_terms.push(arena.mul(&[alpha_id, sin_wx]));
}
if !beta.is_zero() {
let beta_id = ode_ratio_to_expr(arena, &beta);
let cos_wx = arena.cos(omega_x);
inner_terms.push(arena.mul(&[beta_id, cos_wx]));
}
if inner_terms.is_empty() {
Some(arena.zero)
} else {
let inner = if inner_terms.len() == 1 {
inner_terms[0]
} else {
arena.add(&inner_terms)
};
Some(arena.mul(&[var, inner]))
}
}
}
fn try_exp_particular(
arena: &mut Arena,
coeff_r: &num_rational::Ratio<num_bigint::BigInt>,
r: &num_rational::Ratio<num_bigint::BigInt>,
b: &num_rational::Ratio<num_bigint::BigInt>,
c: &num_rational::Ratio<num_bigint::BigInt>,
var: ExprId,
) -> Option<ExprId> {
use num_traits::Zero;
let char_val = r * r + b * r + c;
let r_id = ode_ratio_to_expr(arena, r);
let rx = arena.mul(&[r_id, var]);
let exp_rx = arena.exp(rx);
if !char_val.is_zero() {
let a_val = coeff_r / &char_val;
let a_id = ode_ratio_to_expr(arena, &a_val);
Some(arena.mul(&[a_id, exp_rx]))
} else {
let deriv_val = num_rational::Ratio::from_integer(num_bigint::BigInt::from(2)) * r + b;
if !deriv_val.is_zero() {
let a_val = coeff_r / &deriv_val;
let a_id = ode_ratio_to_expr(arena, &a_val);
let x_exp = arena.mul(&[var, exp_rx]);
Some(arena.mul(&[a_id, x_exp]))
} else {
let two_r_val = num_rational::Ratio::from_integer(num_bigint::BigInt::from(2));
let a_val = coeff_r / &two_r_val;
let a_id = ode_ratio_to_expr(arena, &a_val);
let two_id = arena.int(2);
let x_sq = arena.pow(var, two_id);
let x2_exp = arena.mul(&[x_sq, exp_rx]);
Some(arena.mul(&[a_id, x2_exp]))
}
}
}
fn try_bernoulli(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
var_sym: SymbolId,
) -> Option<OdeResult> {
tracing::debug!("ode: trying Bernoulli");
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let d2y_dx2 = arena.intern(ExprNode::Derivative(dy_dx, var));
if expr_contains(arena, expr, d2y_dx2) {
return None;
}
let children = match arena.node(expr).clone() {
ExprNode::Add(children) => children,
_ => return None,
};
use num_traits::Zero;
let mut dy_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut p_x_terms: Vec<ExprId> = Vec::new();
let mut q_x_terms: Vec<(ExprId, num_rational::Ratio<num_bigint::BigInt>)> = Vec::new();
for &child in &children {
let (coeff, term) = arena.as_coeff_term(child);
if term == dy_dx {
dy_coeff += coeff;
} else if !contains_sym(arena, child, func_sym) {
return None; } else if let Some(px) = extract_coeff_of_func(arena, child, func, func_sym, var_sym) {
p_x_terms.push(px);
} else if let Some((qx, n)) = extract_bernoulli_term(arena, child, func, func_sym, var_sym)
{
q_x_terms.push((qx, n));
} else {
return None;
}
}
if dy_coeff.is_zero() || q_x_terms.is_empty() {
return None;
}
let n_val = q_x_terms[0].1.clone();
if n_val.is_zero() || n_val.is_one() {
return None;
}
for (_, n) in &q_x_terms[1..] {
if *n != n_val {
return None;
}
}
let inv_dy = num_rational::Ratio::<num_bigint::BigInt>::one() / &dy_coeff;
let inv_dy_id = ode_ratio_to_expr(arena, &inv_dy);
let p_raw = if p_x_terms.is_empty() {
arena.zero
} else if p_x_terms.len() == 1 {
p_x_terms[0]
} else {
arena.add(&p_x_terms)
};
let p_x = if dy_coeff.is_one() {
p_raw
} else {
let s = arena.mul(&[inv_dy_id, p_raw]);
crate::transforms::eval::eval(arena, s)
};
let q_sum: Vec<ExprId> = q_x_terms.iter().map(|(qx, _)| *qx).collect();
let q_raw = if q_sum.len() == 1 {
q_sum[0]
} else {
arena.add(&q_sum)
};
let neg_q_raw = arena.neg(q_raw);
let q_x = if dy_coeff.is_one() {
neg_q_raw
} else {
let s = arena.mul(&[inv_dy_id, neg_q_raw]);
crate::transforms::eval::eval(arena, s)
};
if contains_sym(arena, p_x, func_sym) || contains_sym(arena, q_x, func_sym) {
return None;
}
let one_minus_n = num_rational::Ratio::<num_bigint::BigInt>::one() - &n_val;
let one_minus_n_id = ode_ratio_to_expr(arena, &one_minus_n);
let new_p = arena.mul(&[one_minus_n_id, p_x]);
let new_p = crate::transforms::eval::eval(arena, new_p);
let new_q = arena.mul(&[one_minus_n_id, q_x]);
let new_q = crate::transforms::eval::eval(arena, new_q);
let v = arena.symbol("__v");
let dv = arena.intern(ExprNode::Derivative(v, var));
let pv = arena.mul(&[new_p, v]);
let neg_new_q = arena.neg(new_q);
let linear_expr = arena.add(&[dv, pv, neg_new_q]);
let v_sym = match arena.node(v) {
ExprNode::Symbol(sid) => *sid,
_ => return None,
};
let v_result = try_first_order_linear_general(arena, linear_expr, v, var, v_sym, var_sym)
.or_else(|| try_simple_separable(arena, linear_expr, v, var, v_sym, var_sym))?;
let inv_one_minus_n = num_rational::Ratio::<num_bigint::BigInt>::one() / &one_minus_n;
let inv_id = ode_ratio_to_expr(arena, &inv_one_minus_n);
let solution = arena.pow(v_result.solution, inv_id);
let solution = crate::transforms::eval::eval(arena, solution);
Some(OdeResult {
solution,
constants: v_result.constants,
})
}
fn extract_bernoulli_term(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
func_sym: SymbolId,
_var_sym: SymbolId,
) -> Option<(ExprId, num_rational::Ratio<num_bigint::BigInt>)> {
if let ExprNode::Pow(base, exp) = arena.node(expr).clone()
&& base == func
{
let n = arena.as_num(exp)?.clone();
return Some((arena.one, n));
}
if let ExprNode::Mul(ref factors) = arena.node(expr).clone() {
let mut yn_idx = None;
let mut yn_exp = None;
for (i, &f) in factors.iter().enumerate() {
if let ExprNode::Pow(base, exp) = arena.node(f).clone()
&& base == func
&& let Some(n) = arena.as_num(exp)
{
yn_idx = Some(i);
yn_exp = Some(n.clone());
break;
}
}
if let (Some(idx), Some(n)) = (yn_idx, yn_exp) {
let mut other: Vec<ExprId> = Vec::new();
for (i, &f) in factors.iter().enumerate() {
if i != idx {
if contains_sym(arena, f, func_sym) {
return None;
}
other.push(f);
}
}
let qx = match other.len() {
0 => arena.one,
1 => other[0],
_ => arena.mul(&other),
};
return Some((qx, n));
}
}
{
let (coeff, term) = arena.as_coeff_term(expr);
if !coeff.is_one()
&& term != expr
&& let Some((inner_qx, n)) =
extract_bernoulli_term(arena, term, func, func_sym, _var_sym)
{
let coeff_id = ode_ratio_to_expr(arena, &coeff);
let qx = arena.mul(&[coeff_id, inner_qx]);
return Some((qx, n));
}
}
if let ExprNode::Neg(inner) = arena.node(expr).clone()
&& let Some((inner_qx, n)) = extract_bernoulli_term(arena, inner, func, func_sym, _var_sym)
{
let neg_qx = arena.neg(inner_qx);
return Some((neg_qx, n));
}
None
}
fn try_euler_cauchy(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
var_sym: SymbolId,
) -> Option<OdeResult> {
tracing::debug!("ode: trying Euler-Cauchy");
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let d2y_dx2 = arena.intern(ExprNode::Derivative(dy_dx, var));
if !expr_contains(arena, expr, d2y_dx2) {
return None;
}
let children = match arena.node(expr).clone() {
ExprNode::Add(children) => children,
_ => return None,
};
use num_traits::Zero;
let mut a_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut b_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut c_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let two_id = arena.int(2);
let x_sq = arena.pow(var, two_id);
for &child in &children {
let (coeff, term) = arena.as_coeff_term(child);
if term == func {
c_coeff += coeff;
} else if !contains_sym(arena, child, func_sym) {
return None; } else if let ExprNode::Mul(ref factors) = arena.node(term).clone() {
let has_d2 = factors.contains(&d2y_dx2);
let has_d1 = factors.contains(&dy_dx);
let has_xsq = factors.contains(&x_sq);
let has_xvar = factors.contains(&var);
let other: Vec<ExprId> = factors
.iter()
.copied()
.filter(|&f| f != d2y_dx2 && f != dy_dx && f != x_sq && f != var)
.collect();
for &of in &other {
if contains_sym(arena, of, func_sym) || contains_sym(arena, of, var_sym) {
return None;
}
}
let extra = if other.is_empty() {
num_rational::Ratio::<num_bigint::BigInt>::from_integer(1.into())
} else {
let e = if other.len() == 1 {
other[0]
} else {
arena.mul(&other)
};
arena.as_num(e)?.clone()
};
if has_d2 && has_xsq && !has_d1 && !has_xvar {
a_coeff += &coeff * &extra;
} else if has_d1 && has_xvar && !has_d2 && !has_xsq {
b_coeff += &coeff * &extra;
} else {
return None;
}
} else {
return None;
}
}
if a_coeff.is_zero() {
return None;
}
let p = (&b_coeff - &a_coeff) / &a_coeff;
let q = &c_coeff / &a_coeff;
let four = num_rational::Ratio::<num_bigint::BigInt>::from_integer(4.into());
let disc = &p * &p - &four * &q;
let c1 = arena.symbol("C1");
let c2 = arena.symbol("C2");
if disc.is_positive() {
let r_var = arena.symbol("__r");
let r_sq = arena.pow(r_var, two_id);
let p_id = ode_ratio_to_expr(arena, &p);
let q_id = ode_ratio_to_expr(arena, &q);
let p_r = arena.mul(&[p_id, r_var]);
let char_eq = arena.add(&[r_sq, p_r, q_id]);
let roots = crate::transforms::solve::solve(arena, char_eq, r_var);
if roots.len() >= 2 {
let r1 = roots[0].value;
let r2 = roots[1].value;
let x_r1 = arena.pow(var, r1);
let x_r2 = arena.pow(var, r2);
let t1 = arena.mul(&[c1, x_r1]);
let t2 = arena.mul(&[c2, x_r2]);
let solution = arena.add(&[t1, t2]);
Some(OdeResult {
solution,
constants: vec![c1, c2],
})
} else {
None
}
} else if disc.is_zero() {
let two_r = num_rational::Ratio::<num_bigint::BigInt>::from_integer(2.into());
let r = -(&p) / &two_r;
let r_id = ode_ratio_to_expr(arena, &r);
let x_r = arena.pow(var, r_id);
let ln_x = arena.ln(var);
let c2_ln = arena.mul(&[c2, ln_x]);
let inner = arena.add(&[c1, c2_ln]);
let solution = arena.mul(&[inner, x_r]);
Some(OdeResult {
solution,
constants: vec![c1, c2],
})
} else {
let two_r = num_rational::Ratio::<num_bigint::BigInt>::from_integer(2.into());
let alpha = -(&p) / &two_r;
let neg_disc = -disc;
let neg_disc_id = ode_ratio_to_expr(arena, &neg_disc);
let half = arena.rational(1, 2);
let sqrt_neg_disc = arena.pow(neg_disc_id, half);
let two_expr = arena.int(2);
let beta = arena.div(sqrt_neg_disc, two_expr);
let beta = crate::transforms::eval::eval(arena, beta);
let ln_x = arena.ln(var);
let beta_ln_x = arena.mul(&[beta, ln_x]);
let cos_part = arena.cos(beta_ln_x);
let sin_part = arena.sin(beta_ln_x);
let c1_cos = arena.mul(&[c1, cos_part]);
let c2_sin = arena.mul(&[c2, sin_part]);
let trig_part = arena.add(&[c1_cos, c2_sin]);
let solution = if alpha.is_zero() {
trig_part
} else {
let alpha_id = ode_ratio_to_expr(arena, &alpha);
let x_alpha = arena.pow(var, alpha_id);
arena.mul(&[x_alpha, trig_part])
};
Some(OdeResult {
solution,
constants: vec![c1, c2],
})
}
}
fn try_variation_of_parameters(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
var_sym: SymbolId,
) -> Option<OdeResult> {
tracing::debug!("ode: trying variation of parameters");
use num_traits::Zero;
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let d2y_dx2 = arena.intern(ExprNode::Derivative(dy_dx, var));
let children = match arena.node(expr).clone() {
ExprNode::Add(children) => children,
_ => return None,
};
let mut a_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut b_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut c_coeff = num_rational::Ratio::<num_bigint::BigInt>::zero();
let mut f_terms: Vec<ExprId> = Vec::new();
for &child in &children {
let (coeff, term) = arena.as_coeff_term(child);
if term == d2y_dx2 {
a_coeff += coeff;
} else if term == dy_dx {
b_coeff += coeff;
} else if term == func {
c_coeff += coeff;
} else if !contains_sym(arena, child, func_sym) {
f_terms.push(child);
} else {
return None;
}
}
if a_coeff.is_zero() || f_terms.is_empty() {
return None;
}
let b = &b_coeff / &a_coeff;
let c = &c_coeff / &a_coeff;
let f_sum = if f_terms.len() == 1 {
f_terms[0]
} else {
arena.add(&f_terms)
};
let neg_f = arena.neg(f_sum);
let a_id = ode_ratio_to_expr(arena, &a_coeff);
let g_x = arena.div(neg_f, a_id);
let g_x = crate::transforms::eval::eval(arena, g_x);
let homo = solve_characteristic_equation(arena, b.clone(), c.clone(), var)?;
let (y1, y2) = extract_fundamental_solutions(arena, homo.solution, &homo.constants)?;
let y1_prime = crate::transforms::diff::diff(arena, y1, var);
let y2_prime = crate::transforms::diff::diff(arena, y2, var);
let w_term1 = arena.mul(&[y1, y2_prime]);
let w_term2 = arena.mul(&[y2, y1_prime]);
let wronskian = arena.sub(w_term1, w_term2);
let wronskian = crate::transforms::eval::eval(arena, wronskian);
let wronskian = crate::transforms::expand::expand(arena, wronskian);
let wronskian = crate::transforms::eval::eval(arena, wronskian);
let wronskian = if contains_sym(arena, wronskian, var_sym) {
let w0 = crate::transforms::subs::subs(arena, wronskian, var, arena.zero);
let w0 = crate::transforms::eval::eval(arena, w0);
if w0 == arena.zero {
return None;
}
if b.is_zero() {
w0
} else {
let neg_b_id = ode_ratio_to_expr(arena, &(-b.clone()));
let neg_bx = arena.mul(&[neg_b_id, var]);
let exp_nbx = arena.exp(neg_bx);
arena.mul(&[w0, exp_nbx])
}
} else {
wronskian
};
if wronskian == arena.zero {
return None;
}
let y2_g = arena.mul(&[y2, g_x]);
let integrand1 = arena.div(y2_g, wronskian);
let integrand1 = crate::transforms::eval::eval(arena, integrand1);
let integral1 = crate::transforms::integrate::integrate(arena, integrand1, var);
if matches!(arena.node(integral1), ExprNode::Integral(_, _)) {
return None;
}
let y1_g = arena.mul(&[y1, g_x]);
let integrand2 = arena.div(y1_g, wronskian);
let integrand2 = crate::transforms::eval::eval(arena, integrand2);
let integral2 = crate::transforms::integrate::integrate(arena, integrand2, var);
if matches!(arena.node(integral2), ExprNode::Integral(_, _)) {
return None;
}
let term1 = arena.mul(&[y1, integral1]);
let neg_term1 = arena.neg(term1);
let term2 = arena.mul(&[y2, integral2]);
let y_p = arena.add(&[neg_term1, term2]);
let y_p = crate::transforms::eval::eval(arena, y_p);
let solution = arena.add(&[homo.solution, y_p]);
let solution = crate::transforms::eval::eval(arena, solution);
Some(OdeResult {
solution,
constants: homo.constants,
})
}
fn extract_fundamental_solutions(
arena: &mut Arena,
homo_solution: ExprId,
constants: &[ExprId],
) -> Option<(ExprId, ExprId)> {
if constants.len() != 2 {
return None;
}
let c1 = constants[0];
let c2 = constants[1];
let zero = arena.zero;
let one = arena.one;
let y1 = crate::transforms::subs::subs(arena, homo_solution, c1, one);
let y1 = crate::transforms::subs::subs(arena, y1, c2, zero);
let y1 = crate::transforms::eval::eval(arena, y1);
let y2 = crate::transforms::subs::subs(arena, homo_solution, c1, zero);
let y2 = crate::transforms::subs::subs(arena, y2, c2, one);
let y2 = crate::transforms::eval::eval(arena, y2);
if y1 == arena.zero || y2 == arena.zero {
return None;
}
Some((y1, y2))
}
type Rat = num_rational::Ratio<num_bigint::BigInt>;
const MAX_ODE_ORDER: usize = 12;
struct LinearCcOde {
coeffs: Vec<Rat>,
forcing: Vec<ExprId>,
}
fn derivative_chain(arena: &mut Arena, func: ExprId, var: ExprId) -> Vec<ExprId> {
let mut chain = vec![func];
for k in 1..=MAX_ODE_ORDER {
let d = arena.intern(ExprNode::Derivative(chain[k - 1], var));
chain.push(d);
}
chain
}
fn extract_linear_cc(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
) -> Option<LinearCcOde> {
use num_traits::Zero;
let chain = derivative_chain(arena, func, var);
let children: Vec<ExprId> = match arena.node(expr).clone() {
ExprNode::Add(c) => c.to_vec(),
_ => vec![expr],
};
let mut coeffs = vec![Rat::zero(); MAX_ODE_ORDER + 1];
let mut forcing = Vec::new();
for child in children {
let (coeff, term) = arena.as_coeff_term(child);
if let Some(k) = chain.iter().position(|&d| d == term) {
coeffs[k] += coeff;
} else if !contains_sym(arena, child, func_sym) {
forcing.push(child);
} else {
return None;
}
}
while coeffs.len() > 1 && coeffs.last().is_some_and(|c| c.is_zero()) {
coeffs.pop();
}
if coeffs.len() < 2 {
return None; }
Some(LinearCcOde { coeffs, forcing })
}
#[derive(Clone, Debug, PartialEq)]
struct CQ {
re: Rat,
im: Rat,
}
impl CQ {
fn new(re: Rat, im: Rat) -> Self {
Self { re, im }
}
fn real(re: Rat) -> Self {
Self {
re,
im: num_traits::Zero::zero(),
}
}
fn is_zero(&self) -> bool {
num_traits::Zero::is_zero(&self.re) && num_traits::Zero::is_zero(&self.im)
}
fn add(&self, o: &CQ) -> CQ {
CQ::new(&self.re + &o.re, &self.im + &o.im)
}
fn sub(&self, o: &CQ) -> CQ {
CQ::new(&self.re - &o.re, &self.im - &o.im)
}
fn mul(&self, o: &CQ) -> CQ {
CQ::new(
&self.re * &o.re - &self.im * &o.im,
&self.re * &o.im + &self.im * &o.re,
)
}
fn div(&self, o: &CQ) -> CQ {
let denom = &o.re * &o.re + &o.im * &o.im;
let num = self.mul(&CQ::new(o.re.clone(), -o.im.clone()));
CQ::new(num.re / &denom, num.im / &denom)
}
fn scale(&self, r: &Rat) -> CQ {
CQ::new(&self.re * r, &self.im * r)
}
}
fn shifted_coefficient(p: &[Rat], j: usize, lambda: &CQ) -> CQ {
let mut acc = CQ::real(num_traits::Zero::zero());
let mut lambda_pow = CQ::real(num_traits::One::one());
for (k, a_k) in p.iter().enumerate().skip(j) {
let binom = binomial_rat(k, j);
let term = lambda_pow.scale(&(a_k * binom));
acc = acc.add(&term);
lambda_pow = lambda_pow.mul(lambda);
}
acc
}
fn binomial_rat(n: usize, k: usize) -> Rat {
let mut r = Rat::from_integer(1.into());
for i in 0..k {
r *= Rat::from_integer(((n - i) as i64).into());
r /= Rat::from_integer(((i + 1) as i64).into());
}
r
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
enum TrigKind {
None,
Cos,
Sin,
}
struct ForcingTerm {
coeff: Rat,
degree: usize,
a: Rat,
b: Rat,
kind: TrigKind,
}
fn parse_forcing_term(arena: &mut Arena, term: ExprId, var: ExprId) -> Option<ForcingTerm> {
use num_traits::{Signed, Zero};
let (mut coeff, t) = arena.as_coeff_term(term);
let factors: Vec<ExprId> = if t == arena.one {
Vec::new()
} else {
match arena.node(t).clone() {
ExprNode::Mul(c) => c.to_vec(),
_ => vec![t],
}
};
let mut degree = 0usize;
let mut a = Rat::zero();
let mut b = Rat::zero();
let mut kind = TrigKind::None;
let mut seen_exp = false;
for f in factors {
if f == var {
degree += 1;
continue;
}
match arena.node(f).clone() {
ExprNode::Pow(base, exp) if base == var => {
let n = arena.as_num(exp)?.clone();
if !n.is_integer() || n.is_negative() {
return None;
}
degree += usize::try_from(n.to_integer()).ok()?;
}
ExprNode::Exp(inner) => {
if seen_exp {
return None;
}
let (rate, c) = extract_linear_numeric(arena, inner, var)?;
if !c.is_zero() {
return None;
}
seen_exp = true;
a = rate;
}
ExprNode::Sin(inner) => {
if kind != TrigKind::None {
return None;
}
let (freq, c) = extract_linear_numeric(arena, inner, var)?;
if !c.is_zero() || freq.is_zero() {
return None;
}
kind = TrigKind::Sin;
if freq.is_negative() {
coeff = -coeff;
b = -freq;
} else {
b = freq;
}
}
ExprNode::Cos(inner) => {
if kind != TrigKind::None {
return None;
}
let (freq, c) = extract_linear_numeric(arena, inner, var)?;
if !c.is_zero() || freq.is_zero() {
return None;
}
kind = TrigKind::Cos;
b = freq.abs();
}
_ => {
let r = arena.as_num(f)?;
coeff *= r.clone();
}
}
}
Some(ForcingTerm {
coeff,
degree,
a,
b,
kind,
})
}
fn rat_poly_expr(arena: &mut Arena, coeffs: &[Rat], var: ExprId) -> ExprId {
build_polynomial_expr(arena, coeffs, var)
}
fn particular_for_group(
arena: &mut Arena,
p: &[Rat],
q: &[Rat],
a: &Rat,
b: &Rat,
kind: TrigKind,
var: ExprId,
) -> Option<ExprId> {
use num_traits::Zero;
let lambda = CQ::new(a.clone(), b.clone());
let n = p.len() - 1;
let c: Vec<CQ> = (0..=n)
.map(|j| shifted_coefficient(p, j, &lambda))
.collect();
let s = c.iter().position(|cj| !cj.is_zero())?;
let d = q.len().checked_sub(1)?;
let m = d + s;
let mut w = vec![CQ::real(Rat::zero()); m + 1];
for k in (0..=d).rev() {
let mut rhs = CQ::real(q[k].clone());
for j in (s + 1)..=n {
if k + j > m {
continue;
}
let fact = falling_factorial_rat(k + j, j);
rhs = rhs.sub(&c[j].mul(&w[k + j]).scale(&fact));
}
let lead = c[s].scale(&falling_factorial_rat(k + s, s));
w[k + s] = rhs.div(&lead);
}
let wr: Vec<Rat> = w.iter().map(|z| z.re.clone()).collect();
let wi: Vec<Rat> = w.iter().map(|z| z.im.clone()).collect();
let wr_x = rat_poly_expr(arena, &wr, var);
let wi_x = rat_poly_expr(arena, &wi, var);
let exp_ax = if a.is_zero() {
arena.one
} else {
let a_id = ode_ratio_to_expr(arena, a);
let ax = arena.mul(&[a_id, var]);
arena.exp(ax)
};
let body = match kind {
TrigKind::None => wr_x,
TrigKind::Cos | TrigKind::Sin => {
let b_id = ode_ratio_to_expr(arena, b);
let bx = arena.mul(&[b_id, var]);
let cos_bx = arena.cos(bx);
let sin_bx = arena.sin(bx);
if kind == TrigKind::Cos {
let t1 = arena.mul(&[wr_x, cos_bx]);
let t2 = arena.mul(&[wi_x, sin_bx]);
arena.sub(t1, t2)
} else {
let t1 = arena.mul(&[wr_x, sin_bx]);
let t2 = arena.mul(&[wi_x, cos_bx]);
arena.add(&[t1, t2])
}
}
};
let y_p = arena.mul(&[exp_ax, body]);
let y_p = crate::transforms::expand::expand(arena, y_p);
Some(crate::transforms::eval::eval(arena, y_p))
}
fn falling_factorial_rat(n: usize, j: usize) -> Rat {
let mut r = Rat::from_integer(1.into());
for i in 0..j {
r *= Rat::from_integer(((n - i) as i64).into());
}
r
}
fn homogeneous_basis_cc(arena: &mut Arena, coeffs: &[Rat], var: ExprId) -> Option<Vec<ExprId>> {
let p = crate::poly::Poly::from_coeffs(coeffs.to_vec());
let r_sym = arena.symbol("__r_cc");
let mut basis: Vec<ExprId> = Vec::new();
for (factor, mult) in p.squarefree_factors() {
let f_expr = crate::poly::polybridge::poly_to_expr(arena, &factor, r_sym);
let roots = crate::transforms::solve::solve(arena, f_expr, r_sym);
if roots.is_empty() {
return None;
}
let i_unit = arena.i_unit;
let mut used = vec![false; roots.len()];
for i in 0..roots.len() {
if used[i] {
continue;
}
used[i] = true;
let root = roots[i].value;
let is_complex = crate::base::walk::contains(arena, root, i_unit);
if !is_complex {
for j in 0..mult {
basis.push(mode_real(arena, root, j, var));
}
continue;
}
let (re, im) = arena.as_real_imag_expr(root);
let re = crate::transforms::eval::eval(arena, re);
let im = crate::transforms::eval::eval(arena, im);
if arena.is_zero_structural(im) {
for j in 0..mult {
basis.push(mode_real(arena, re, j, var));
}
continue;
}
let neg_im = arena.neg(im);
let neg_im = crate::transforms::eval::eval(arena, neg_im);
for (k, other) in roots.iter().enumerate() {
if used[k] {
continue;
}
let (re2, im2) = arena.as_real_imag_expr(other.value);
let re2 = crate::transforms::eval::eval(arena, re2);
let im2 = crate::transforms::eval::eval(arena, im2);
if re2 == re && im2 == neg_im {
used[k] = true;
break;
}
}
let beta = if arena
.as_num(im)
.is_some_and(num_traits::Signed::is_negative)
{
neg_im
} else {
im
};
for j in 0..mult {
let (c, s) = mode_complex(arena, re, beta, j, var);
basis.push(c);
basis.push(s);
}
}
}
Some(basis)
}
fn mode_real(arena: &mut Arena, r: ExprId, j: usize, var: ExprId) -> ExprId {
let x_pow = match j {
0 => arena.one,
1 => var,
_ => {
let e = arena.int(j as i64);
arena.pow(var, e)
}
};
if arena.is_zero_structural(r) {
return x_pow;
}
let rx = arena.mul(&[r, var]);
let exp_rx = arena.exp(rx);
let m = arena.mul(&[x_pow, exp_rx]);
crate::transforms::eval::eval(arena, m)
}
fn mode_complex(
arena: &mut Arena,
alpha: ExprId,
beta: ExprId,
j: usize,
var: ExprId,
) -> (ExprId, ExprId) {
let envelope = mode_real(arena, alpha, j, var);
let bx = arena.mul(&[beta, var]);
let cos_bx = arena.cos(bx);
let sin_bx = arena.sin(bx);
let c = arena.mul(&[envelope, cos_bx]);
let s = arena.mul(&[envelope, sin_bx]);
(
crate::transforms::eval::eval(arena, c),
crate::transforms::eval::eval(arena, s),
)
}
fn try_nth_order_linear_const_coeff(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
) -> Option<OdeResult> {
use num_traits::Zero;
let ode = extract_linear_cc(arena, expr, func, var, func_sym)?;
let n = ode.coeffs.len() - 1;
if n < 2 {
return None;
}
tracing::debug!(order = n, "ode: nth-order linear constant-coefficient");
let mut groups: Vec<((Rat, Rat, TrigKind), Vec<Rat>)> = Vec::new();
for &term in &ode.forcing {
let ft = parse_forcing_term(arena, term, var)?;
let key = (ft.a.clone(), ft.b.clone(), ft.kind);
let entry = match groups.iter_mut().find(|(k, _)| *k == key) {
Some(e) => e,
None => {
groups.push((key, Vec::new()));
groups.last_mut()?
}
};
if entry.1.len() <= ft.degree {
entry.1.resize(ft.degree + 1, Rat::zero());
}
entry.1[ft.degree] -= ft.coeff; }
let basis = homogeneous_basis_cc(arena, &ode.coeffs, var)?;
if basis.len() != n {
return None;
}
let mut constants = Vec::with_capacity(n);
let mut terms = Vec::with_capacity(n + groups.len());
for (i, &b) in basis.iter().enumerate() {
let c = arena.symbol(&format!("C{}", i + 1));
constants.push(c);
terms.push(arena.mul(&[c, b]));
}
for ((a, b, kind), q) in &groups {
if q.iter().all(Zero::is_zero) {
continue;
}
let y_p = particular_for_group(arena, &ode.coeffs, q, a, b, *kind, var)?;
terms.push(y_p);
}
let solution = arena.add(&terms);
let solution = crate::transforms::eval::eval(arena, solution);
Some(OdeResult {
solution,
constants,
})
}
fn clairaut_f(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
var_sym: SymbolId,
p: ExprId,
) -> Option<ExprId> {
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let d2y_dx2 = arena.intern(ExprNode::Derivative(dy_dx, var));
if !expr_contains(arena, expr, dy_dx) || expr_contains(arena, expr, d2y_dx2) {
return None;
}
let g = crate::transforms::subs::subs(arena, expr, dy_dx, p);
let coeffs = crate::transforms::solve::symbolic_poly_coeffs(arena, g, func)?;
if coeffs.len() != 2 {
return None;
}
let c_y = coeffs[1];
let c_0 = coeffs[0];
let p_sym = match arena.node(p) {
ExprNode::Symbol(s) => *s,
_ => return None,
};
if contains_sym(arena, c_y, var_sym) || contains_sym(arena, c_y, p_sym) {
return None;
}
let ratio = arena.div(c_0, c_y);
let neg_ratio = arena.neg(ratio);
let xp = arena.mul(&[var, p]);
let f = arena.sub(neg_ratio, xp);
let f = crate::transforms::eval::eval(arena, f);
let f = crate::transforms::expand::expand(arena, f);
let f = crate::transforms::eval::eval(arena, f);
if contains_sym(arena, f, var_sym) || contains_sym(arena, f, func_sym) {
return None;
}
Some(f)
}
fn try_clairaut(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
func_sym: SymbolId,
var_sym: SymbolId,
) -> Option<OdeResult> {
let p = arena.symbol("__clairaut_p");
let f = clairaut_f(arena, expr, func, var, func_sym, var_sym, p)?;
tracing::debug!("ode: Clairaut form recognised");
let c1 = arena.symbol("C1");
let f_c = crate::transforms::subs::subs(arena, f, p, c1);
let cx = arena.mul(&[c1, var]);
let solution = arena.add(&[cx, f_c]);
let solution = crate::transforms::eval::eval(arena, solution);
Some(OdeResult {
solution,
constants: vec![c1],
})
}
fn riccati_coeffs(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
) -> Option<(ExprId, ExprId, ExprId)> {
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let d2y_dx2 = arena.intern(ExprNode::Derivative(dy_dx, var));
if expr_contains(arena, expr, d2y_dx2) {
return None;
}
let (rest, c) = extract_m_n(arena, expr, func, var)?;
if crate::base::walk::contains(arena, c, func) || crate::base::walk::contains(arena, c, var) {
return None;
}
if crate::base::walk::contains(arena, rest, dy_dx) {
return None;
}
let neg_rest = arena.neg(rest);
let rhs = arena.div(neg_rest, c);
let rhs = crate::transforms::eval::eval(arena, rhs);
let coeffs = crate::transforms::solve::symbolic_poly_coeffs(arena, rhs, func)?;
if coeffs.len() != 3 {
return None;
}
let (q0, q1, q2) = (coeffs[0], coeffs[1], coeffs[2]);
if arena.is_zero_structural(q0) || arena.is_zero_structural(q2) {
return None;
}
Some((q0, q1, q2))
}
pub fn solve_riccati(
arena: &mut Arena,
expr: ExprId,
func: ExprId,
var: ExprId,
particular: ExprId,
) -> Option<OdeResult> {
let (q0, q1, q2) = riccati_coeffs(arena, expr, func, var)?;
if !checkodesol(arena, expr, particular, func, var) {
let yp_prime = crate::transforms::diff::diff(arena, particular, var);
let yp2 = arena.mul(&[particular, particular]);
let t1 = arena.mul(&[q1, particular]);
let t2 = arena.mul(&[q2, yp2]);
let rhs = arena.add(&[q0, t1, t2]);
let residual = arena.sub(yp_prime, rhs);
let residual = crate::transforms::eval::eval(arena, residual);
let residual = crate::transforms::expand::expand(arena, residual);
let residual = crate::transforms::eval::eval(arena, residual);
if !arena.is_zero_structural(residual) {
return None;
}
}
let _ = q0;
let v = arena.symbol("__riccati_v");
let dv = arena.intern(ExprNode::Derivative(v, var));
let two = arena.int(2);
let two_q2_yp = arena.mul(&[two, q2, particular]);
let coeff = arena.add(&[q1, two_q2_yp]);
let coeff = crate::transforms::eval::eval(arena, coeff);
let coeff_v = arena.mul(&[coeff, v]);
let lin = arena.add(&[dv, coeff_v, q2]);
let lin = crate::transforms::eval::eval(arena, lin);
let v_res = dsolve(arena, lin, v, var)?;
if crate::base::walk::contains(arena, v_res.solution, v) {
return None; }
let neg_one = arena.neg_one;
let inv_v = arena.pow(v_res.solution, neg_one);
let solution = arena.add(&[particular, inv_v]);
let solution = crate::transforms::eval::eval(arena, solution);
Some(OdeResult {
solution,
constants: v_res.constants,
})
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum OdeType {
SimpleSeparable,
FullSeparable,
FirstOrderLinearCC,
FirstOrderLinearVC,
ExactFirstOrder,
SecondOrderLinearCCHomogeneous,
SecondOrderLinearCCNonHomogeneous,
Bernoulli,
EulerCauchy,
VariationOfParameters,
HomogeneousCoefficient,
NthOrderReducible,
NthOrderLinearConstCoeff,
IntegratingFactor,
Clairaut,
Riccati,
Unknown,
}
pub fn classify_ode(arena: &mut Arena, expr: ExprId, func: ExprId, var: ExprId) -> OdeType {
let var_sym = match arena.node(var) {
ExprNode::Symbol(sid) => *sid,
_ => return OdeType::Unknown,
};
let func_sym = match arena.node(func) {
ExprNode::Symbol(sid) => *sid,
_ => return OdeType::Unknown,
};
let dy_dx = arena.intern(ExprNode::Derivative(func, var));
let d2y_dx2 = arena.intern(ExprNode::Derivative(dy_dx, var));
let d3y_dx3 = arena.intern(ExprNode::Derivative(d2y_dx2, var));
if expr_contains(arena, expr, d3y_dx3) {
if let Some(ode) = extract_linear_cc(arena, expr, func, var, func_sym)
&& ode.coeffs.len() >= 4
{
return OdeType::NthOrderLinearConstCoeff;
}
return OdeType::Unknown;
}
if expr_contains(arena, expr, d2y_dx2) {
if let ExprNode::Add(ref children) = arena.node(expr).clone() {
let mut all_const_coeff = true;
let mut has_y2 = false;
let mut forcing: Vec<ExprId> = Vec::new();
for &child in children {
let (_coeff, term) = arena.as_coeff_term(child);
if term == d2y_dx2 {
has_y2 = true;
} else if term == dy_dx {
} else if term == func {
} else if !contains_sym(arena, child, func_sym) {
forcing.push(child);
} else {
all_const_coeff = false;
break;
}
}
if has_y2 && all_const_coeff {
if forcing.is_empty() {
return OdeType::SecondOrderLinearCCHomogeneous;
}
let all_uc = forcing
.iter()
.all(|&t| parse_forcing_term(arena, t, var).is_some());
if all_uc {
return OdeType::SecondOrderLinearCCNonHomogeneous;
}
if try_variation_of_parameters(arena, expr, func, var, func_sym, var_sym).is_some()
{
return OdeType::VariationOfParameters;
}
return OdeType::Unknown;
}
}
if try_euler_cauchy(arena, expr, func, var, func_sym, var_sym).is_some() {
return OdeType::EulerCauchy;
}
{
let d2_ph = arena.symbol("__d2_cls");
let d1_ph = arena.symbol("__d1_cls");
let stripped = crate::transforms::subs::subs(arena, expr, d2y_dx2, d2_ph);
let stripped = crate::transforms::subs::subs(arena, stripped, dy_dx, d1_ph);
if !contains_sym(arena, stripped, var_sym) {
return OdeType::NthOrderReducible;
}
}
return OdeType::Unknown;
}
if expr_contains(arena, expr, dy_dx) {
if !contains_sym_outside_deriv(arena, expr, func_sym, dy_dx) {
return OdeType::SimpleSeparable;
}
if let ExprNode::Add(ref children) = arena.node(expr).clone() {
let mut is_linear = true;
let mut has_var_coeff = false;
for &child in children {
let (coeff, term) = arena.as_coeff_term(child);
if term == dy_dx {
} else if term == func {
} else if contains_sym(arena, child, func_sym) {
if let Some(px) = extract_coeff_of_func(arena, child, func, func_sym, var_sym) {
let _ = coeff; if contains_sym(arena, px, var_sym) {
has_var_coeff = true;
}
} else {
is_linear = false;
break;
}
}
}
if is_linear {
if has_var_coeff {
return OdeType::FirstOrderLinearVC;
}
return OdeType::FirstOrderLinearCC;
}
}
if let Some((m_ex, n_ex)) = extract_m_n(arena, expr, func, var)
&& (contains_sym(arena, m_ex, func_sym) || contains_sym(arena, n_ex, func_sym))
{
let dm_dy = crate::transforms::diff::diff(arena, m_ex, func);
let dn_dx = crate::transforms::diff::diff(arena, n_ex, var);
let check = arena.sub(dm_dy, dn_dx);
let check = crate::transforms::eval::eval(arena, check);
let check = crate::transforms::expand::expand(arena, check);
let check = crate::transforms::eval::eval(arena, check);
if check == arena.zero {
return OdeType::ExactFirstOrder;
}
}
{
let p = arena.symbol("__clairaut_p_cls");
if clairaut_f(arena, expr, func, var, func_sym, var_sym, p).is_some() {
return OdeType::Clairaut;
}
}
if let ExprNode::Add(ref bn_children) = arena.node(expr).clone() {
let mut bn_has_dy = false;
let mut bn_has_yn = false;
let mut bn_ok = true;
let mut bn_has_free = false;
for &child in bn_children {
let (_, term) = arena.as_coeff_term(child);
if term == dy_dx {
bn_has_dy = true;
} else if !contains_sym(arena, child, func_sym) {
bn_has_free = true;
} else if extract_coeff_of_func(arena, child, func, func_sym, var_sym).is_some() {
} else if extract_bernoulli_term(arena, child, func, func_sym, var_sym).is_some() {
bn_has_yn = true;
} else {
bn_ok = false;
break;
}
}
if bn_has_dy && bn_has_yn && bn_ok && !bn_has_free {
return OdeType::Bernoulli;
}
}
if let ExprNode::Add(ref hc_children) = arena.node(expr).clone() {
let mut hc_has_dy = false;
let mut hc_other: Vec<ExprId> = Vec::new();
for &child in hc_children {
let (_, term) = arena.as_coeff_term(child);
if term == dy_dx {
hc_has_dy = true;
} else {
hc_other.push(child);
}
}
if hc_has_dy && !hc_other.is_empty() {
let hc_rhs = if hc_other.len() == 1 {
arena.neg(hc_other[0])
} else {
let s = arena.add(&hc_other);
arena.neg(s)
};
if contains_sym(arena, hc_rhs, func_sym) && contains_sym(arena, hc_rhs, var_sym) {
let v_cls = arena.symbol("__v_cls");
let vx = arena.mul(&[v_cls, var]);
let sub = crate::transforms::subs::subs(arena, hc_rhs, func, vx);
let sub = crate::transforms::eval::eval(arena, sub);
let sub = crate::transforms::expand::expand(arena, sub);
let sub = crate::transforms::eval::eval(arena, sub);
let sub = arena.cancel_expr(sub, var);
let sub = crate::transforms::eval::eval(arena, sub);
if !contains_sym(arena, sub, var_sym) {
return OdeType::HomogeneousCoefficient;
}
}
}
}
if let ExprNode::Add(ref children) = arena.node(expr).clone() {
let mut has_deriv = false;
let mut other_terms: Vec<ExprId> = Vec::new();
for &child in children {
let (_coeff, term) = arena.as_coeff_term(child);
if term == dy_dx {
has_deriv = true;
} else {
other_terms.push(child);
}
}
if has_deriv && !other_terms.is_empty() {
let rhs = if other_terms.len() == 1 {
arena.neg(other_terms[0])
} else {
let sum = arena.add(&other_terms);
arena.neg(sum)
};
let factors = collect_mul_factors(arena, rhs);
let mut can_separate = true;
for &factor in &factors {
let has_x = contains_sym(arena, factor, var_sym);
let has_y = contains_sym(arena, factor, func_sym);
if has_x && has_y {
can_separate = false;
break;
}
}
if can_separate {
return OdeType::FullSeparable;
}
}
}
if try_integrating_factor_ode(arena, expr, func, var, func_sym, var_sym).is_some() {
return OdeType::IntegratingFactor;
}
if riccati_coeffs(arena, expr, func, var).is_some() {
return OdeType::Riccati;
}
return OdeType::Unknown;
}
OdeType::Unknown
}
fn expr_contains(arena: &Arena, haystack: ExprId, needle: ExprId) -> bool {
if haystack == needle {
return true;
}
let node = arena.node(haystack).clone();
for &child in node.children().iter() {
if expr_contains(arena, child, needle) {
return true;
}
}
false
}
fn contains_sym_outside_deriv(
arena: &Arena,
expr: ExprId,
sym: SymbolId,
deriv_node: ExprId,
) -> bool {
if expr == deriv_node {
return false;
}
match arena.node(expr).clone() {
ExprNode::Symbol(s) => s == sym,
other => {
for &child in other.children().iter() {
if contains_sym_outside_deriv(arena, child, sym, deriv_node) {
return true;
}
}
false
}
}
}
pub fn checkodesol(
arena: &mut Arena,
ode_expr: ExprId,
solution: ExprId,
func: ExprId,
var: ExprId,
) -> bool {
let chain = derivative_chain(arena, func, var);
let max_order = (1..chain.len())
.rev()
.find(|&k| expr_contains(arena, ode_expr, chain[k]))
.unwrap_or(0);
let mut sol_derivs = vec![solution];
for k in 1..=max_order {
let d = crate::transforms::diff::diff(arena, sol_derivs[k - 1], var);
sol_derivs.push(d);
}
let mut result = ode_expr;
for k in (0..=max_order).rev() {
result = crate::transforms::subs::subs(arena, result, chain[k], sol_derivs[k]);
}
result = crate::transforms::eval::eval(arena, result);
result = crate::transforms::expand::expand(arena, result);
result = crate::transforms::eval::eval(arena, result);
if result == arena.zero {
return true;
}
result = crate::transforms::expand::expand(arena, result);
result = crate::transforms::eval::eval(arena, result);
if result == arena.zero {
return true;
}
let simplified = crate::simplify::simplify_engine::unified_simplify(
arena,
result,
&crate::simplify::simplify_engine::SimplifyOpts::default(),
);
arena.is_zero_structural(simplified.expr)
}
pub fn solve_ode_system(a_matrix: &Matrix, t_var: &Ex) -> Option<Vec<Ex>> {
let n = a_matrix.nrows();
if !a_matrix.is_square() || n == 0 {
return None;
}
for i in 0..n {
for j in 0..n {
if a_matrix.get(i, j).contains(t_var) {
return None;
}
}
}
if ode_system_is_diagonal(a_matrix, n) {
return Some(solve_ode_system_diagonal(a_matrix, t_var, n));
}
if let Some(sol) = solve_ode_system_eigen(a_matrix, t_var, n) {
return Some(sol);
}
Some(solve_ode_system_series(a_matrix, t_var, n))
}
pub fn solve_ode_system_ivp(
a_matrix: &Matrix,
t_var: &Ex,
x0: &[Ex],
) -> Result<Vec<Ex>, crate::base::errors::SymplexError> {
use crate::base::errors::SymplexError;
let n = a_matrix.nrows();
if !a_matrix.is_square() || n == 0 {
return Err(SymplexError::InvalidArgument {
operation: "solve_ode_system_ivp",
reason: "coefficient matrix must be square and non-empty".into(),
});
}
if x0.len() != n {
return Err(SymplexError::InvalidArgument {
operation: "solve_ode_system_ivp",
reason: format!("expected {n} initial values, got {}", x0.len()),
});
}
let general =
solve_ode_system(a_matrix, t_var).ok_or_else(|| SymplexError::ComputationFailed {
operation: "solve_ode_system_ivp",
reason: "could not solve the homogeneous system".into(),
})?;
let ctx = t_var.context();
let constants: Vec<Ex> = (1..=n).map(|i| ctx.symbol(&format!("C{i}"))).collect();
let zero = ctx.int(0);
let eqs: Vec<Ex> = general
.iter()
.zip(x0)
.map(|(g, v)| (g.subs(t_var, &zero).eval() - v).eval())
.collect();
let sol = crate::api::expr_solve_ext::linsolve(&eqs, &constants).map_err(|e| {
SymplexError::ComputationFailed {
operation: "solve_ode_system_ivp",
reason: format!("could not fit initial values: {e}"),
}
})?;
let pairs = match sol {
crate::api::expr_solve_ext::LinearSolution::Inconsistent => {
return Err(SymplexError::NoSolution {
operation: "solve_ode_system_ivp",
reason: "initial values are inconsistent with the general solution".into(),
});
}
crate::api::expr_solve_ext::LinearSolution::Unique(p) => p,
crate::api::expr_solve_ext::LinearSolution::Parametric { solution, free } => solution
.into_iter()
.filter(|(c, _)| !free.contains(c))
.collect(),
};
Ok(general
.iter()
.map(|g| {
let mut e = g.clone();
for (c, v) in &pairs {
e = e.subs(c, v);
}
e.eval().simplify()
})
.collect())
}
pub fn solve_ode_system_nonhomogeneous(
a_matrix: &Matrix,
b_vec: &[Ex],
t_var: &Ex,
) -> Option<Vec<Ex>> {
let n = a_matrix.nrows();
if !a_matrix.is_square() || n == 0 || b_vec.len() != n {
return None;
}
let x_h = solve_ode_system(a_matrix, t_var)?;
let ctx = t_var.context();
let neg_one = ctx.int(-1);
let neg_a = a_matrix.scale(&neg_one);
let neg_at = neg_a.scale(t_var);
let exp_neg_at = neg_at.matrix_exp().unwrap_or_else(|_| {
neg_at
.exp_series(12)
.expect("exp_series: matrix must be square")
});
let b_col = Matrix::col_vector(b_vec.to_vec());
let integrand_matrix = exp_neg_at
.matmul(&b_col)
.expect("matmul: dimension mismatch")
.eval();
let mut integrated = Vec::with_capacity(n);
for i in 0..n {
integrated.push(integrand_matrix.get(i, 0).integrate(t_var).eval());
}
let integrated_col = Matrix::col_vector(integrated);
let at = a_matrix.scale(t_var);
let exp_at = at.matrix_exp().unwrap_or_else(|_| {
at.exp_series(12)
.expect("exp_series: matrix must be square")
});
let particular = exp_at
.matmul(&integrated_col)
.expect("matmul: dimension mismatch")
.eval();
let mut solution = Vec::with_capacity(n);
for (i, x_h_i) in x_h.iter().enumerate() {
let xi: Ex = x_h_i + particular.get(i, 0);
solution.push(xi.eval());
}
Some(solution)
}
pub fn classify_ode_system_is_constant(a_matrix: &Matrix, t_var: &Ex) -> bool {
if !a_matrix.is_square() {
return false;
}
let n = a_matrix.nrows();
for i in 0..n {
for j in 0..n {
if a_matrix.get(i, j).contains(t_var) {
return false;
}
}
}
true
}
fn ode_system_is_diagonal(m: &Matrix, n: usize) -> bool {
for i in 0..n {
for j in 0..n {
if i != j && !m.get(i, j).is_zero_structural() {
return false;
}
}
}
true
}
fn solve_ode_system_diagonal(a_matrix: &Matrix, t_var: &Ex, n: usize) -> Vec<Ex> {
let ctx = t_var.context();
(0..n)
.map(|i| {
let ci = ctx.symbol(&format!("C{}", i + 1));
let aii = a_matrix.get(i, i);
if aii.is_zero_structural() {
ci } else {
let exp_term = (aii * t_var).exp();
&ci * &exp_term
}
})
.collect()
}
fn solve_ode_system_series(a_matrix: &Matrix, t_var: &Ex, n: usize) -> Vec<Ex> {
let ctx = t_var.context();
let m = a_matrix.scale(t_var);
let exp_m = m
.matrix_exp()
.unwrap_or_else(|_| m.exp_series(12).expect("exp_series: matrix must be square"));
let constants: Vec<Ex> = (1..=n).map(|i| ctx.symbol(&format!("C{i}"))).collect();
let c_vec = Matrix::col_vector(constants);
let result = exp_m.matmul(&c_vec).expect("matmul: dimension mismatch");
(0..n).map(|i| result.get(i, 0).eval()).collect()
}
fn solve_ode_system_eigen(a_matrix: &Matrix, t_var: &Ex, n: usize) -> Option<Vec<Ex>> {
let ctx = t_var.context();
let eigenvalues = match a_matrix.eigenvals() {
Ok(ev) => ev,
Err(_) => return None,
};
if eigenvalues.len() < n {
return None;
}
for i in 0..eigenvalues.len() {
if eigenvalues[i + 1..].contains(&eigenvalues[i]) {
return None;
}
}
let i_unit = ctx.i_unit();
let zero_ex = ctx.int(0);
let neg_i = -&i_unit;
let identity = Matrix::identity(&ctx, n);
let mut solution: Vec<Ex> = (0..n).map(|_| ctx.int(0)).collect();
let mut const_idx = 1_usize;
let mut used = vec![false; eigenvalues.len()];
for idx in 0..eigenvalues.len() {
if used[idx] {
continue;
}
used[idx] = true;
let ev = &eigenvalues[idx];
if ev.contains(&i_unit) {
let alpha = ev.subs(&i_unit, &zero_ex).eval().simplify();
let ev_minus_alpha = ev - α
let beta = (&ev_minus_alpha * &neg_i).eval().simplify();
for j in (idx + 1)..eigenvalues.len() {
if !used[j] && eigenvalues[j].contains(&i_unit) {
let alpha_j = eigenvalues[j].subs(&i_unit, &zero_ex).eval().simplify();
let ej_diff = &eigenvalues[j] - &alpha_j;
let beta_j = (&ej_diff * &neg_i).eval().simplify();
let beta_sum = (&beta + &beta_j).eval().simplify();
if beta_sum.is_zero_structural() {
used[j] = true;
break;
}
}
}
let ev_identity = identity.scale(ev);
let a_shifted = a_matrix
.sub(&ev_identity)
.expect("sub: shape mismatch")
.eval()
.simplify();
let null_basis = a_shifted.nullspace();
if null_basis.is_empty() {
return None;
}
let mut u_re = Vec::with_capacity(n);
let mut w_im = Vec::with_capacity(n);
for row in 0..n {
let vi = null_basis[0].get(row, 0).eval().simplify();
let re = vi.subs(&i_unit, &zero_ex).eval().simplify();
let vi_minus_re = &vi - &re;
let im = (&vi_minus_re * &neg_i).eval().simplify();
u_re.push(re);
w_im.push(im);
}
let c_a = ctx.symbol(&format!("C{const_idx}"));
let c_b = ctx.symbol(&format!("C{}", const_idx + 1));
const_idx += 2;
let exp_alpha_t = if alpha.is_zero_structural() {
ctx.int(1)
} else {
(&alpha * t_var).exp()
};
let cos_beta_t = (&beta * t_var).cos();
let sin_beta_t = (&beta * t_var).sin();
for row in 0..n {
let cu = &cos_beta_t * &u_re[row];
let sw = &sin_beta_t * &w_im[row];
let su = &sin_beta_t * &u_re[row];
let cw = &cos_beta_t * &w_im[row];
let m1 = &exp_alpha_t * &(&cu - &sw);
let m2 = &exp_alpha_t * &(&su + &cw);
let ca_m1 = &c_a * &m1;
let cb_m2 = &c_b * &m2;
let contrib = &ca_m1 + &cb_m2;
solution[row] = &solution[row] + &contrib;
}
} else {
let ev_identity = identity.scale(ev);
let a_shifted = a_matrix
.sub(&ev_identity)
.expect("sub: shape mismatch")
.eval()
.simplify();
let null_basis = a_shifted.nullspace();
if null_basis.is_empty() {
return None;
}
let ci = ctx.symbol(&format!("C{const_idx}"));
const_idx += 1;
let exp_ev_t = if ev.is_zero_structural() {
ctx.int(1)
} else {
(ev * t_var).exp()
};
for (row, sol_row) in solution.iter_mut().enumerate().take(n) {
let vi = null_basis[0].get(row, 0).eval().simplify();
if !vi.is_zero_structural() {
let exp_vi = &exp_ev_t * &vi;
let ci_exp_vi = &ci * &exp_vi;
*sol_row = &*sol_row + &ci_exp_vi;
}
}
}
}
let solution: Vec<Ex> = solution.into_iter().map(|s| s.eval()).collect();
Some(solution)
}
#[cfg(test)]
mod tests {
use super::*;
fn sym(a: &mut Arena, name: &str) -> ExprId {
a.symbol(name)
}
fn display(a: &Arena, id: ExprId) -> String {
a.display(id).to_string()
}
#[test]
fn solve_dy_dx_eq_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let dy = a.intern(ExprNode::Derivative(y, x));
let expr = a.sub(dy, x);
let result = dsolve(&mut a, expr, y, x).expect("should solve");
let s = display(&a, result.solution);
assert!(s.contains("C1"), "should have constant: {s}");
assert!(s.contains("x"), "should contain x: {s}");
}
#[test]
fn solve_dy_dx_eq_0() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let dy = a.intern(ExprNode::Derivative(y, x));
let result = dsolve(&mut a, dy, y, x).expect("should solve");
let s = display(&a, result.solution);
assert!(s.contains("C1"), "should be constant: {s}");
}
#[test]
fn solve_second_order_y_plus_y_eq_0() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let dy = a.intern(ExprNode::Derivative(y, x));
let d2y = a.intern(ExprNode::Derivative(dy, x));
let expr = a.add(&[d2y, y]);
let r = dsolve(&mut a, expr, y, x).expect("should solve y'' + y = 0");
let s = display(&a, r.solution);
assert!(
s.contains("C1") && s.contains("C2"),
"should have two constants: {s}"
);
assert!(
s.contains("cos") && s.contains("sin"),
"should use trig form (cos and sin): {s}"
);
}
#[test]
fn solve_y_prime_plus_2y_eq_0() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let dy = a.intern(ExprNode::Derivative(y, x));
let two = a.int(2);
let two_y = a.mul(&[two, y]);
let expr = a.add(&[dy, two_y]);
let result = dsolve(&mut a, expr, y, x).expect("should solve");
let s = display(&a, result.solution);
assert!(s.contains("C1"), "should have constant: {s}");
assert!(s.contains("exp"), "should contain exp: {s}");
}
#[test]
fn solve_full_separable_xy() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let dy = a.intern(ExprNode::Derivative(y, x));
let xy = a.mul(&[x, y]);
let expr = a.sub(dy, xy); let result = dsolve(&mut a, expr, y, x).expect("should solve y' = xy");
let s = display(&a, result.solution);
assert!(s.contains("C1"), "should have constant: {s}");
assert!(s.contains("exp"), "should contain exp: {s}");
}
#[test]
fn solve_variable_coeff_linear_2xy() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let dy = a.intern(ExprNode::Derivative(y, x));
let two = a.int(2);
let two_x_y = a.mul(&[two, x, y]);
let expr = a.add(&[dy, two_x_y]); let result = dsolve(&mut a, expr, y, x).expect("should solve y' + 2xy = 0");
let s = display(&a, result.solution);
assert!(s.contains("C1"), "should have constant: {s}");
assert!(s.contains("exp"), "should contain exp: {s}");
}
}