import numpy as np
from pathlib import Path

def read(name):
    return np.fromfile(Path(__file__).parent / ('mlp-' + name + '.bin'), dtype='<f4').reshape(64, 64)

def half(x):
    return x.astype(np.float16).astype(np.float32)

# Positive Graphcore F143 values, bias 8, at operand scale -2.
codes = np.arange(128)
levels = np.where(codes < 8, codes * 2.**-10,
                  (1 + (codes % 8) / 8) * 2.**((codes // 8) - 8)) / 4

def quantize(x):
    distance = abs(abs(x)[..., None] - levels)
    minimum = distance.min(axis=-1, keepdims=True)
    # Nearest, ties to an even encoding.
    score = np.where(distance == minimum, codes % 2, 2)
    return (levels[score.argmin(axis=-1)] * np.sign(x)).astype(np.float32)

def gelu(x):
    return half(.5 * x * (1 + np.tanh(.7978846 * (x + .044715 * x**3))))

def report(label, actual, expected):
    error = abs(actual - expected)
    print(label, 'max', error.max(), 'rms', np.sqrt(np.mean(error**2)),
          'outside tolerance', np.count_nonzero(error > .001 + .02 * abs(expected)))

actual = [read('actual-' + str(i)) for i in range(4)]
reference = [read('reference-' + str(i)) for i in range(4)]
inputs = [read('input-' + str(i)) for i in range(3)]
for i in range(4):
    report('checkpoint ' + str(i), actual[i], reference[i])
report('GELU applied to actual GEMM1', actual[1], gelu(actual[0]))
qa, qr = quantize(actual[1]), quantize(reference[1])
different = qa != qr
print('FP8 threshold crossings', different.sum(), 'rows', np.flatnonzero(different.any(axis=1)))
print('crossing indices', np.argwhere(different).tolist())
predicted = half(qa @ inputs[2])
report('GEMM2 using actual GELU1, quantized', actual[2], predicted)
report('final using actual GELU1, quantized', actual[3], gelu(predicted))
report('final GELU using actual GEMM2', actual[3], gelu(actual[2]))
print('final mismatch rows', np.flatnonzero((abs(actual[3]-reference[3]) > .001 + .02 * abs(reference[3])).any(axis=1)))
print('cosine', np.sum(actual[3]*reference[3]) / np.sqrt(np.sum(actual[3]**2) * np.sum(reference[3]**2)))
print('minimum row cosine', np.min(np.sum(actual[3]*reference[3], axis=1) /
      np.sqrt(np.sum(actual[3]**2, axis=1) * np.sum(reference[3]**2, axis=1))))
