Deriving gradients by hand for every new loss is slow and error-prone. You will write the machine that does it: given a formula and values for its variables, return the formula's value and its partial derivative with respect to every variable.
A formula is a nested tuple. The node types are:
("const", c) the number c
("var", name) the variable called name (a string)
("add", a, b) ("sub", a, b) ("mul", a, b) ("div", a, b)
("pow", a, k) a to the power k, where k is a plain number, not a node
("exp", a) ("log", a) ("sin", a) ("cos", a) ("sigmoid", a) with sigmoid(t) = 1 / (1 + e^(-t))
Write gradient(expr, values) where values maps every variable name that occurs in expr (and possibly some that do not) to a number. Return the tuple (value, grads): the value of the formula as a float, and a dict with one entry for every name in values, its partial derivative as a float (0.0 for names that do not occur).
Formulas may share subexpressions: the same tuple object can appear in several places, so a formula of a few hundred tuples can stand for an expansion with astronomically many nodes. The helpers mse_expr(X, y) and log_loss_expr(X, y) build the mean squared error and the average log loss of a linear model over the variables "w0", "w1", ... and "b", so you can compare with gradients you derived yourself.
Examples
Input: expr = ("mul", ("var", "x"), ("sin", ("var", "y"))), values = {"x": 2.0, "y": 0.5}
Output: (0.958851077208406, {'x': 0.479425538604203, 'y': 1.7551651237807455})
Explanation: the value is 2·sin(0.5); the slope in x is sin(0.5) and in y it is 2·cos(0.5).
Input: expr = ("pow", ("sub", ("add", ("mul", ("var", "w"), ("const", 3)), ("var", "b")), ("const", 11)), 2),
values = {"w": 2.0, "b": 1.0}
Output: (16.0, {'w': -24.0, 'b': -8.0})
Explanation: the squared error (3w + b - 11)² with error -4: the slopes are 2·(-4)·3 and 2·(-4).
Constraints
- at most
10**4distinct tuple objects in a formula, nested at most 300 deep; the expanded formula can be far larger - every
loggets a positive argument and everydiva non-zero divisor - floats are compared with a tolerance of
1e-6
Goals
- Apply the chain rule mechanically to a formula given as a tree
- Add up the contributions of a quantity that is used in several places
- Handle shared subexpressions without recomputing them