"""Original worked-example checks for houhuiyang.com's MicroGPT article.

Run: python3 check_math.py
Standard library only. This is NOT a GPT implementation or training benchmark.
"""
import math


class Scalar:
    """A small reverse-mode graph for addition and multiplication examples."""

    def __init__(self, number, edges=()):
        self.number = float(number)
        self.edges = edges
        self.gradient = 0.0

    def __add__(self, rhs):
        return Scalar(self.number + rhs.number, ((self, 1.0), (rhs, 1.0)))

    def __mul__(self, rhs):
        return Scalar(self.number * rhs.number,
                      ((self, rhs.number), (rhs, self.number)))

    def backward(self):
        ordered, seen = [], set()

        def visit(node):
            if node in seen:
                return
            seen.add(node)
            for parent, _ in node.edges:
                visit(parent)
            ordered.append(node)

        visit(self)
        for node in ordered:
            node.gradient = 0.0
        self.gradient = 1.0
        for node in reversed(ordered):
            for parent, derivative in node.edges:
                parent.gradient += node.gradient * derivative


def close(actual, expected, tolerance=1e-7):
    assert math.isclose(actual, expected, rel_tol=tolerance, abs_tol=tolerance), (
        actual, expected)


def softmax(numbers):
    shifted = [math.exp(number - max(numbers)) for number in numbers]
    return [value / sum(shifted) for value in shifted]


def difference(function, values, index, h=1e-5):
    left, right = list(values), list(values)
    left[index] -= h
    right[index] += h
    return (function(right) - function(left)) / (2 * h)


def main():
    a, b = Scalar(2), Scalar(3)
    loss = a * b + a
    loss.backward()
    close(loss.number, 8)
    close(a.gradient, 4)
    close(b.gradient, 2)
    for index, expected in enumerate([a.gradient, b.gradient]):
        close(difference(lambda x: x[0] * x[1] + x[0], [2, 3], index), expected)
    square = a * a
    square.backward()
    close(a.gradient, 4)
    print('PASS shared graph: L=8, gradients=(4, 2); repeated input: 4')

    query = [1, 0, 1, 0]
    keys = [[1, 0, 0, 0], [0, 1, 0, 0], [1, 0, 1, 0]]
    values = [[1, 0, 0, 0], [0, 2, 0, 0], [0, 0, 3, 0]]
    scores = [sum(q * k for q, k in zip(query, key)) / math.sqrt(4) for key in keys]
    weights = softmax(scores)
    output = [sum(w * value[j] for w, value in zip(weights, values)) for j in range(4)]
    close(sum(weights), 1)
    for actual, expected in zip(output, [0.3072, 0.3726, 1.5194, 0]):
        close(actual, expected, 5e-5)
    print('PASS attention:', [round(x, 4) for x in weights],
          'output:', [round(x, 4) for x in output])

    logits = [math.log(p) for p in [0.2, 0.5, 0.3]]
    expected_gradients = [-0.8, 0.5, 0.3]
    for i, expected in enumerate(expected_gradients):
        close(difference(lambda z: -math.log(softmax(z)[0]), logits, i), expected)
    print('PASS cross-entropy central differences:', expected_gradients)

    gradient, beta1, beta2 = 0.2, 0.85, 0.99
    first = (1 - beta1) * gradient
    second = (1 - beta2) * gradient ** 2
    first_hat = first / (1 - beta1)
    second_hat = second / (1 - beta2)
    delta = 0.01 * first_hat / (math.sqrt(second_hat) + 1e-8)
    close(first, 0.03)
    close(second, 0.0004)
    close(delta, 0.01)
    print('PASS Adam first-step decrement:', round(delta, 10))

    close(math.log(27), 3.295836866004329)
    assert 32 * 27 + 3328 == 4192
    print('PASS parameter count: 4192; uniform loss:', round(math.log(27), 4))


if __name__ == '__main__':
    main()
