/////////////////////////////////////////////////////////////////////// // Math Type Library // $Id: simplifier.tcc,v 1.9 2002/05/05 10:49:09 cparpart Exp $ // (This file contains the expression simplifier template members) // // Copyright (c) 2002 by Christian Parpart // // This library is free software; you can redistribute it and/or // modify it under the terms of the GNU Library General Public // License as published by the Free Software Foundation; either // version 2 of the License, or (at your option) any later version. // // This library is distributed in the hope that it will be useful, // but WITHOUT ANY WARRANTY; without even the implied warranty of // MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the GNU // Library General Public License for more details. // // You should have received a copy of the GNU Library General Public License // along with this library; see the file COPYING.LIB. If not, write to // the Free Software Foundation, Inc., 59 Temple Place - Suite 330, // Boston, MA 02111-1307, USA. /////////////////////////////////////////////////////////////////////// #ifndef libmath_simplifier_h #error You may not include math++/simplify.tcc directly; include math++/simplify.h instead. #endif #include #include #include namespace math { template bool isConst(const TNode *ANode) { // note, that symbols may not be interpreted as const here, even if they're return !ANode || (ANode->nodeType() != TNode::PARAM_NODE && ANode->nodeType() != TNode::SYMBOL_NODE && ANode->nodeType() < TNode::FUNC_NODE && isConst(ANode->left()) && isConst(ANode->right())); } template TNode *TSimplifier::simplify(const TNode *AExpression) { TNode *oldResult = 0; TNode *newResult = AExpression->clone(); do { delete oldResult; oldResult = newResult; TSimplifier simplifier; oldResult->accept(simplifier); newResult = simplifier.FResult; } while (*newResult != *oldResult); delete oldResult; return newResult; } template TSimplifier::TSimplifier() : FResult(0) { } template T TSimplifier::calculate(const TNode *AExpr) const { static TLibrary library; if (!library.constants()) { library.insert(TConstant("pi", 3.1415));// M_PI)); // just in case they're used. :) library.insert(TConstant("e", 2.1718));// M_E)); } return TCalculator::calculate(TFunction("tmp", AExpr), T(), library); } template void TSimplifier::visit(TNumberNode *ANode) { // (-number) = -(number) T value(ANode->number()); if (value < 0) { FResult = new TNegNode( new TNumberNode(-value) ); return; } FResult = ANode->clone(); } template void TSimplifier::visit(TSymbolNode *ANode) { FResult = ANode->clone(); } template void TSimplifier::visit(TParamNode *ANode) { FResult = ANode->clone(); } template void TSimplifier::visit(TPlusNode *ANode) { std::auto_ptr > left(simplify(ANode->left())); std::auto_ptr > right(simplify(ANode->right())); // check constness if (isConst(left.get()) && isConst(right.get())) { // Do not use simple calculate(ANode): FResult = new TNumberNode(calculate(left.get()) + calculate(right.get())); return; } // 0+a = a if (left->nodeType() == TNode::NUMBER_NODE && calculate(left.get()) == T(0)) { FResult = right->clone(); return; } // a+0 = a if (right->nodeType() == TNode::NUMBER_NODE && calculate(right.get()) == T(0)) { FResult = left->clone(); return; } // a+(-0) = a if (right->nodeType() == TNode::NEG_NODE && right->right()->nodeType() == TNode::NUMBER_NODE && calculate(right->right()) == T(0)) { FResult = left->clone(); return; } // a + a = 2a if (*left == *right) { FResult = new TMulNode( new TNumberNode(T(2)), right->clone() ); return; } // a + (-a) = 0 if (right->nodeType() == TNode::NEG_NODE && *left == *right->right()) { FResult = new TNumberNode(T(0)); return; } // (a + b) + (-b) = b if (left->nodeType() == TNode::PLUS_NODE && right->nodeType() == TNode::NEG_NODE && *left->right() == *right->right()) { FResult = left->left()->clone(); return; } // a + -(a + b) = -b if (right->nodeType() == TNode::NEG_NODE && right->right()->nodeType() == TNode::PLUS_NODE && *left == *right->right()->left()) { FResult = new TNegNode( right->right()->right()->clone() ); return; } // (-a) + (-b) = -(a + b) if (left->nodeType() == TNode::NEG_NODE) { if (right->nodeType() == TNode::NEG_NODE) { FResult = new TNegNode( new TPlusNode( left->right()->clone(), right->right()->clone() ) ); } else { // (-a) + b = b + (-a) FORM TRANSFORMATION FResult = new TPlusNode( right->clone(), left->clone() ); } return; } // n * a + a = a(n + 1) if (left->nodeType() == TNode::MUL_NODE && *left->right() == *right) { FResult = new TMulNode( right->clone(), new TPlusNode( left->left()->clone(), new TNumberNode(T(1)) ) ); return; } // a*b + a*c = a(b + c) if (left->nodeType() == TNode::MUL_NODE && right->nodeType() == TNode::MUL_NODE && *left->left() == *right->left()) { FResult = new TMulNode( left->left()->clone(), new TPlusNode( left->right()->clone(), right->right()->clone() ) ); return; } // a*b + a*c = a*b + c*a = b*a + a*c = b*a + c*a = a(b + c) // WE MUST FIND AN ALGORITHM FOR GENERIC PATTERN MATCHING SOON !!! // (f * g) + (f / h) = f * (g + 1/h) if (left->nodeType() == TNode::MUL_NODE && right->nodeType() == TNode::DIV_NODE && *left->left() == *right->left()) { FResult = new TMulNode( left->left()->clone(), new TPlusNode( left->right()->clone(), new TDivNode( new TNumberNode(T(1)), right->right()->clone() ) ) ); return; } // nothing special found, just optimized left and right child nodes FResult = new TPlusNode(left->clone(), right->clone()); } template void TSimplifier::visit(TNegNode *ANode) { std::auto_ptr > node(simplify(ANode->node())); // -(-a) = a if (node->nodeType() == TNode::NEG_NODE) { FResult = node->right()->clone(); return; } // nothing special found, just optimized left and right child nodes FResult = new TNegNode(node->clone()); } template void TSimplifier::visit(TMulNode *ANode) { std::auto_ptr > left(simplify(ANode->left())); std::auto_ptr > right(simplify(ANode->right())); // check constness if (isConst(left.get()) && isConst(right.get())) { FResult = new TNumberNode(calculate(ANode)); return; } // 0*a = 0 if (left->nodeType() == TNode::NUMBER_NODE && calculate(left.get()) == T(0)) { FResult = new TNumberNode(T(0)); return; } // a*0 = 0 if (right->nodeType() == TNode::NUMBER_NODE && calculate(right.get()) == T()) { FResult = new TNumberNode(T(0)); return; } // 1*a = a, a if (left->nodeType() == TNode::NUMBER_NODE && calculate(left.get()) == T(1)) { FResult = right->clone(); return; } // a*1 = a, a if (right->nodeType() == TNode::NUMBER_NODE && calculate(right.get()) == T(1)) { FResult = left->clone(); return; } // a*a = a^2 if (*left == *right) { FResult = new TPowNode(left->clone(), new TNumberNode(T(2))); return; } // (-a) * b = -(a * b) if (left->nodeType() == TNode::NEG_NODE) { FResult = new TNegNode( new TMulNode( left->right()->clone(), right->clone() ) ); return; } // a * (-b) = -(a * b) if (right->nodeType() == TNode::NEG_NODE) { FResult = new TNegNode( new TMulNode( left->clone(), right->right()->clone() ) ); return; } // a^n*a = a^(n+1) if (left->nodeType() == TNode::POW_NODE && *left->left() == *right) { FResult = new TPowNode( right->clone(), new TPlusNode( left->right()->clone(), new TNumberNode(T(1)) ) ); return; } // (C*a)*D = (C*D)*a; C, D const. if (left->nodeType() == TNode::MUL_NODE && isConst(left->left()) && isConst(right.get())) { FResult = new TMulNode( new TNumberNode( T(calculate(left->left()) * calculate(right.get())) ), left->right()->clone() ); return; } // (a*b)*b = a*b^2 if (left->nodeType() == TNode::MUL_NODE && *left->right() == *right.get()) { FResult = new TMulNode( left->left()->clone(), new TPowNode( right->clone(), new TNumberNode(T(2)) ) ); return; } // (a*b)*a = b*a^2 if (left->nodeType() == TNode::MUL_NODE && *left->left() == *right.get()) { FResult = new TMulNode( left->right()->clone(), new TPowNode( right->clone(), new TNumberNode(T(2)) ) ); return; } // (a * b) * b^c = a * b^(c + 1) if (left->nodeType() == TNode::MUL_NODE && right->nodeType() == TNode::POW_NODE && *left->right() == *right->left()) { FResult = new TMulNode( left->left()->clone(), new TPowNode( left->right()->clone(), new TPlusNode( right->right()->clone(), new TNumberNode(T(1)) ) ) ); return; } // a^b*c/a = c*a^(b - 1) if (left->nodeType() == TNode::POW_NODE && right->nodeType() == TNode::DIV_NODE && *left->left() == *right->right()) { FResult = new TMulNode( right->left()->clone(), new TPowNode( left->left()->clone(), new TPlusNode( left->right()->clone(), new TNegNode(new TNumberNode(T(1))) ) ) ); return; } // (a/b)*c = (a*c)/b if (left->nodeType() == TNode::DIV_NODE) { FResult = new TDivNode( new TMulNode( left->left()->clone(), right->clone() ), left->right()->clone() ); return; } // a * a^b = a^(b+1) if (right->nodeType() == TNode::POW_NODE && *left == *right->left()) { FResult = new TPowNode( left->clone(), new TPlusNode( right->right()->clone(), new TNumberNode(T(1)) ) ); return; } // nothing special found, just optimized left and right child nodes FResult = new TMulNode(left->clone(), right->clone()); } template void TSimplifier::visit(TDivNode *ANode) { std::auto_ptr > left(simplify(ANode->left())); std::auto_ptr > right(simplify(ANode->right())); // check constness if (isConst(left.get()) && isConst(right.get())) { T divisor(calculate(right.get())); if (divisor == T(0)) // prevent division by zero, by no simplifying FResult = new TDivNode(left->clone(), right->clone()); else FResult = new TNumberNode(calculate(left.get()) / divisor); return; } // 0/a = 0 if (left->nodeType() == TNode::NUMBER_NODE && calculate(left.get()) == T(0)) { FResult = new TNumberNode(T(0)); return; } // a / a = 1 if (*left == *right) { FResult = new TNumberNode(T(1)); return; } // (-a) / b = -(a / b) FORM TRANSFORMATION if (left->nodeType() == TNode::NEG_NODE) { FResult = new TNegNode( new TDivNode( left->right()->clone(), right->clone() ) ); return; } // (a * b) / c = a * (b / c) FORM TRANSFORMATION if (left->nodeType() == TNode::MUL_NODE) { FResult = new TMulNode( left->left()->clone(), new TDivNode( left->right()->clone(), right->clone() ) ); return; } // a / (-b) = -(a / b) FORM TRANSFORMATION if (right->nodeType() == TNode::NEG_NODE) { FResult = new TNegNode( new TDivNode( left->clone(), right->right()->clone() ) ); return; } // a / b^c = a * b^(-c) if (right->nodeType() == TNode::POW_NODE) { FResult = new TMulNode( left->clone(), new TPowNode( right->left()->clone(), new TNegNode( right->right()->clone() ) ) ); return; } // (a^b)/a = a^(b-1) if (left->nodeType() == TNode::POW_NODE && *left->left() == *right.get()) { FResult = new TPowNode( right->clone(), new TPlusNode( left->right()->clone(), new TNegNode(new TNumberNode(T(1))) ) ); return; } // a^b / a^c = a^(b-c) if (left->nodeType() == TNode::POW_NODE && right->nodeType() == TNode::POW_NODE && *left->left() == *right->left()) { FResult = new TPowNode( left->left()->clone(), new TPlusNode( left->right()->clone(), new TNegNode(right->right()->clone()) ) ); return; } // nothing special found, just optimized left and right child nodes FResult = new TDivNode(left->clone(), right->clone()); } template void TSimplifier::visit(TPowNode *ANode) { std::auto_ptr > left(simplify(ANode->left())); std::auto_ptr > right(simplify(ANode->right())); // check constness if (isConst(left.get()) && isConst(right.get())) { FResult = new TNumberNode(calculate(ANode)); return; } // a^0 = 1 if (right->nodeType() == TNode::NUMBER_NODE && calculate(right.get()) == 0) { FResult = new TNumberNode(1); return; } // a^1 = a if (right->nodeType() == TNode::NUMBER_NODE && calculate(right.get()) == 1) { FResult = left->clone(); return; } // f^g^h = f^(g * h) if (left->nodeType() == TNode::POW_NODE) { FResult = new TPowNode( left->left()->clone(), new TMulNode( left->right()->clone(), right->clone() ) ); return; } // nothing special found, just optimized left and right child nodes FResult = new TPowNode(left->clone(), right->clone()); } template void TSimplifier::visit(TSqrtNode *ANode) { FResult = ANode->clone(); } template void TSimplifier::visit(TSinNode *ANode) { FResult = new TSinNode(simplify(ANode->right())); } template void TSimplifier::visit(TCosNode *ANode) { FResult = new TCosNode(simplify(ANode->right())); } template void TSimplifier::visit(TTanNode *ANode) { FResult = new TTanNode(simplify(ANode->right())); } template void TSimplifier::visit(TLnNode *ANode) { std::auto_ptr > node(simplify(ANode->node())); // ln(e) = 1 if (node->nodeType() == TNode::SYMBOL_NODE && static_cast *>(node.get())->symbol() == "e") { FResult = new TNumberNode(T(1)); return; } FResult = new TLnNode(node->clone()); } template void TSimplifier::visit(TFuncNode *ANode) { FResult = new TFuncNode(ANode->name(), simplify(ANode->node())); } template void TSimplifier::visit(TIfNode *ANode) { FResult = new TIfNode(simplify(ANode->condition()), simplify(ANode->trueExpr()), simplify(ANode->falseExpr())); } template void TSimplifier::visit(TEquNode *ANode) { FResult = new TEquNode(simplify(ANode->left()), simplify(ANode->right())); } template void TSimplifier::visit(TUnEquNode *ANode) { FResult = new TUnEquNode(simplify(ANode->left()), simplify(ANode->right())); } template void TSimplifier::visit(TGreaterNode *ANode) { FResult = new TGreaterNode(simplify(ANode->left()), simplify(ANode->right())); } template void TSimplifier::visit(TLessNode *ANode) { FResult = new TLessNode(simplify(ANode->left()), simplify(ANode->right())); } template void TSimplifier::visit(TGreaterEquNode *ANode) { FResult = new TGreaterEquNode(simplify(ANode->left()), simplify(ANode->right())); } template void TSimplifier::visit(TLessEquNode *ANode) { FResult = new TLessEquNode(simplify(ANode->left()), simplify(ANode->right())); } } // namespace math