Building a symbolic math REPL in C
I spent a few days writing this small derivation engine in C. You type expressions, it computes derivatives, evaluates them, does integration and Taylor series. The kind of thing that looks simple until you actually try to build it.
>> f = x^3 + 2*x
x^3 + 2 * x
>> f'(x)
3 * x^2 + 2
>> f'(x)'(x)
6 * x
>> taylor(sin(x), 0, 5)
0.00833333 * x^5 + -0.166667 * x^3 + x
Note: the simplifier isn't perfect. Some results may be messier than they should. But the core works, and building it was a great excercise in recursion, memory management and parsing.
Representing expressions#
Everything in this program is an expr. A number, a variable, , . All the same struct
typedef struct expr {
exprtype type;
int mindepth, maxdepth, nodes;
int ref;
union {
char *symbol;
double num;
struct {
struct expr *left;
struct expr *right;
};
struct expr *unary;
};
} expr;
The union is the interesting part Depending on the type, an epxression is either:
- a leaf: just a number or a symbol (x, 3.14)
- a unary: just one child (sin(x), -x, log(x))
- a binary: just two children (x+1, x^2)
So the expression becomes a tree like this:
ADD
/ \
EXP 1
/ \
x 2
Every node is the same struct. You navigate it recursively. This pattern shows up everywhere in the codebase.
The bitmask trick#
To know what kind of expression you're dealing with, you check expr->type. But instead of comparing agains every possible type, there's a neat trick in the enum:
typedef enum {
EXPR_LEAF_MASK = 1 << 8,
EXPR_SYM,
EXPR_NUM,
EXPR_UNARY_MASK = 1 << 9,
EXPR_LOG,
EXPR_SIN,
EXPR_COS,
EXPR_NEG,
EXPR_BINARY_MASK = 1 << 10,
EXPR_ADD,
EXPR_MUL,
EXPR_SUB,
EXPR_FRAC,
EXPR_EXP
} exprtype;
EXPR_SYM is (1 << 8) + 1, EXPR_NUM is (1 << 8) + 2. EXPR_LOG is (1 << 9) + 1 and so on.
Because each group starts right after a power-of-two mask, every value in that group has that bit set.
So instead of writing:
if (type == EXPR_ADD || type == EXPR_MUL || type == EXPR_SUB || ...)
you just write:
if (type & EXPR_BINARY_MASK)
This makes the recursive tree traversals really clean. For example, freeing an expression:
void freeexpr(expr *f) {
if (f->type == EXPR_SYM) {
free(f->symbol);
} else if (f->type & EXPR_BINARY_MASK) {
release(f->left);
release(f->right);
} else if (f->type & EXPR_UNARY_MASK) {
release(f->unary);
}
free(f);
}
Three cases cover every node type in the entire tree. Note: this approach makes functions like this very clean but there's a catch. You have to remember what are left and right. For example when you evaluating an EXPR_FRAC left and right aren't expressive and you have to remember that left is the numerator and right the denominator.
Memory: reference counting#
This is where things get interesting. When you differentiate an expression, you often need to reuse subtrees. For example, the product rule:
$$
(f \cdot g)' = f' \cdot g + f \cdot g'
$$
Both and appear twice in the result. You could copy them, but that's a waste of memory. Instead, the program uses the count of reference (also known as reference counting).
Every expr has a ref field. When you create one, ref=1. When you share it, you call retain() to increment the count. When you're done with it you call release(), which decrement the count and frees the node only when it hits zero.
Note: in the implementation i 'cache' the nodes zero, one and two because i do a massive usage in the simplification process and having these special variables avoid reallocating these variables every time.
It's not perfect, you have to disciplined about calling retain and release in the right places. But it's much nicer than copying everything or manually tracking ownership.
Differentiation#
This is the fun part. Symbolic differentiation is just pattern matching on the expression tree. Each node type has a rule, and you apply it recursively. The main function is a switch:
expr *derive(expr *f, expr *sym) {
switch (f->type) {
case EXPR_ADD: return deriveAdd(f, sym);
case EXPR_MUL: return deriveMul(f, sym);
case EXPR_EXP: return deriveExp(f, sym);
case EXPR_SIN: return deriveSin(f, sym);
case EXPR_LOG: return deriveLog(f, sym);
case EXPR_SYM: return eq(f, sym) ? one : zero;
case EXPR_NUM: return zero;
// ...
}
}
The base casses are trivial: the derivative of a constant is zero and the derivative of x with respect to x is one (anything else is zero).
The interesting cases are the rules. Here's the product rule:
expr *deriveMul(expr *f, expr *sym) {
expr *left = mul(derive(f->left, sym), retain(f->right));
expr *right = mul(retain(f->left), derive(f->right, sym));
return add(left, right);
}
That's literally f' * g + f * g'. The retain calls are there because we're sharing the original subtrees, the derivative will own a reference to them.
The chain rule falls out naturally too.
When you differentiate sin(x^2), deriveSin calls derive on its argument, which recurses into the x^2subtree. You don't write chain rule handling separately, it's just recursion.
Simplification#
Differentiation produces correct but messy results. derive(x^2) gives you x^2 * (1 * log(x) + 2 * 1 / x) before any cleanup. Simplification is a second recursive pass over the tree that applies algebric identities.
Some cases are straightforward:
x * 0 = 0
1 * x = x
x / x = 1
x^1 = x
x^0 = 1
...
Others are more involved, combining like terms, pulling negations out of fractions, merging exponents (x^2 * x^3 -> x^5), constant folding. The simplifier recurses bottom-up: it simplifies the children first, then applies rules to the result.
This is also where the "not perfect" disclaimer comes in. Simplification is essentially rewriting, and you can always find expressions that don't reduce as far as they could. Getting it fully right would mean implementing a proper term rewritiing system, which is a rabbit hole I chose not to go down.
The parser#
The REPL reads a string and turns it into an expression tree. This is done with hand-written recursive descendant parser, no lexer generator, no parser library.
The grammar looks like this (there's actually a comment in the source):
unary = num | sym | "-" unary | "(" expr ")"
postfix = unary ("^" unary | "'(" sym ")" | "(" sym "=" num ")")*
factor = postfix (("*" | "/") postfix)*
term = factor (("+" | "-") factor)*
expr = term | sym "=" expr
Each level of the grammar is a function that calls one below it. parseTerm calls parseFator which calls parsePostfix which calls parseUnary. This naturally handles operator precedence, multiplication binds thighter than addition because parseFactor recurses before parseTerm collects its results.
The most unusual part is the postfix syntax. After parsing a base expression the parser tries to consume '(x) (differentiate) or (x=3) (evaluate). This what lets you chain derivatives:
>> f'(x)'(x)
It parses f, then sees '(x) and produces the derivative, then sees another '(x) and differentiates again. Just a loop that keeps trying to extend the expression with a postfix operation.
Problems#
A few things I'd do differently with more time:
- Better simplification: the current approach is a bag of pattern matching rules. A proper canonical form (like always sorting terms) would make it more predictable.
- Error messages: right now it just says "Invalid expression" when parsing fails, which is not very helpful.
- More functions: like
tanexp,absare alla missing. - LaTeX output: there's actually a --latex flag that partially works, but it needs polish.
The code is no GitHub if you want to look around or contribute here is the link: derive.c