# Copyright (c) 2021 Graphcore Ltd. All rights reserved.
import numpy as np
import pytest
import popart
import json

# `import test_util` requires adding to sys.path
import sys
from pathlib import Path

sys.path.append(str(Path(__file__).resolve().parent.parent))
import test_util as tu


def grad(dY, X):
    axes = []
    offset = dY.ndim - X.ndim
    for i in range(0, dY.ndim):
        if i < offset or X.shape[i - offset] == 1:
            axes.append(int(i))
    axis_tup = tuple(axis for axis in axes)
    dX = np.sum(a=dY, axis=axis_tup)
    dX = dX.reshape(X.shape)
    return dX


def run_and_test_value(op_tester, inplace, init_builder, reference, mode):
    if inplace:
        op_tester.setPatterns(["InPlace"], enableRuntimeAsserts=False)
    session = op_tester.run(
        init_builder, reference, mode, {"ai.onnx": 8, "ai.graphcore": 1}
    )
    if inplace:
        ir = json.loads(session._serializeIr(popart.IrSerializationFormat.JSON))
        graph = ir["maingraph"]

        inplace = [op for op in graph if op["type"] == "ExpandInplace"]
        assert len(inplace) == 1


def expand(op_tester, inplace, int_type):
    d1 = np.random.rand(3, 1).astype(np.float32)
    d2 = np.array([2, 1, 6]).astype(int_type)

    def init_builder(builder):
        i1 = builder.addInputTensor(d1)
        c = builder.aiOnnx.constant(d2)
        o = builder.aiOnnx.expand([i1, c])
        if inplace:
            o_identity = builder.aiOnnx.identity([o])
            builder.addOutputTensor(o_identity)
            return [o_identity]
        else:
            builder.addOutputTensor(o)
            return [
                o,
                popart.reservedGradientPrefix() + o,
                popart.reservedGradientPrefix() + i1,
            ]

    def reference(ref_data):
        expanded = d1 * np.ones(d2, dtype=np.float32)
        if inplace:
            return [expanded]
        dY = ref_data.getOutputTensorGrad(0)
        dX = grad(dY, d1)
        return [expanded, dY, dX]

    if (
        inplace
    ):  # grad is tested only as part of non invariant version to avoid duplication
        mode = "infer"
    else:
        mode = "train"
    run_and_test_value(op_tester, inplace, init_builder, reference, mode)


def expand_unwind(op_tester, inplace, int_type):
    d1 = np.random.rand(6).astype(np.float32)
    d2 = np.array([3, 6]).astype(int_type)
    d3 = np.random.rand(6, 3).astype(np.float32)

    def init_builder(builder):
        i1 = builder.addInputTensor(d1)
        c = builder.aiOnnx.constant(d2)
        e = builder.aiOnnx.expand([i1, c])
        i2 = builder.addInputTensor(d3)
        m = builder.aiOnnx.matmul([e, i2])

        if inplace:
            m_identity = builder.aiOnnx.identity([m])
        else:
            m_identity = m
        builder.addOutputTensor(m_identity)
        return [m_identity]

    def reference(_):  # ref_data is an unused argument
        expanded = d1 * np.ones(d2, dtype=np.float32)
        m = expanded.dot(d3)
        return [m]

    run_and_test_value(op_tester, inplace, init_builder, reference, "infer")


def expand_scalar(op_tester, inplace, int_type):
    d1 = np.array((3.0), dtype=np.float32)
    d2 = np.array([2, 1, 6]).astype(int_type)

    def init_builder(builder):
        i1 = builder.addInputTensor(d1)
        c = builder.aiOnnx.constant(d2)
        o = builder.aiOnnx.expand([i1, c])
        if inplace:
            o_identity = builder.aiOnnx.identity([o])
            builder.addOutputTensor(o_identity)
            return [o_identity]
        else:
            builder.addOutputTensor(o)
            return [
                o,
                popart.reservedGradientPrefix() + o,
                popart.reservedGradientPrefix() + i1,
            ]

    def reference(ref_data):
        expanded = d1 * np.ones(d2, dtype=np.float32)
        if inplace:
            return [expanded]
        dY = ref_data.getOutputTensorGrad(0)
        dX = grad(dY, d1)
        return [expanded, dY, dX]

    if (
        inplace
    ):  # grad is tested only as part of non invariant version to avoid duplication
        mode = "infer"
    else:
        mode = "train"
    run_and_test_value(op_tester, inplace, init_builder, reference, mode)


def expand_smaller_output(op_tester, inplace, int_type):
    d1 = np.random.rand(3, 1).astype(np.float32)
    d2 = np.array([2, 1, 1]).astype(int_type)

    def init_builder(builder):
        i1 = builder.addInputTensor(d1)
        c = builder.aiOnnx.constant(d2)
        o = builder.aiOnnx.expand([i1, c])
        if inplace:
            o_identity = builder.aiOnnx.identity([o])
            builder.addOutputTensor(o_identity)
            return [o_identity]
        else:
            builder.addOutputTensor(o)
            return [
                o,
                popart.reservedGradientPrefix() + o,
                popart.reservedGradientPrefix() + i1,
            ]

    def reference(ref_data):
        expanded = d1 * np.ones(d2, dtype=np.float32)
        if inplace:
            return [expanded]
        dY = ref_data.getOutputTensorGrad(0)
        dX = grad(dY, d1)
        return [expanded, dY, dX]

    if (
        inplace
    ):  # grad is tested only as part of non invariant version to avoid duplication
        mode = "infer"
    else:
        mode = "train"
    run_and_test_value(op_tester, inplace, init_builder, reference, mode)


def expand_non_one_smaller_output(op_tester, inplace, int_type):
    d1 = np.random.rand(5, 4, 3).astype(np.float32)
    d2 = np.array([3, 4, 2]).astype(int_type)

    def init_builder(builder):
        i1 = builder.addInputTensor(d1)
        c = builder.aiOnnx.constant(d2)
        o = builder.aiOnnx.expand([i1, c])
        if inplace:
            o_identity = builder.aiOnnx.identity([o])
        else:
            o_identity = o
        builder.addOutputTensor(o_identity)
        return [
            o_identity,
            popart.reservedGradientPrefix() + o,
            popart.reservedGradientPrefix() + i1,
        ]

    with pytest.raises(popart.popart_exception) as e_info:
        if inplace:
            op_tester.setPatterns(["InPlace"], enableRuntimeAsserts=False)
        session = op_tester.run(init_builder, None, "infer")
        if inplace:
            ir = json.loads(session._serializeIr(popart.IrSerializationFormat.JSON))
            graph = ir["maingraph"]

            inplace = [op for op in graph if op["type"] == "ExpandInplace"]
            assert len(inplace) == 1

    assert (
        "corresponding dimensions must have the same value or one of them must be 1"
    ) in str(e_info.value)


def test_expand(op_tester):
    expand(op_tester, False, np.int64)
    expand(op_tester, False, np.int32)


def test_expand_inplace(op_tester):
    expand(op_tester, True, np.int64)
    expand(op_tester, True, np.int32)


def test_expand_smaller_output(op_tester):
    expand_smaller_output(op_tester, False, np.int64)
    expand_smaller_output(op_tester, False, np.int32)


def test_expand_smaller_output_inplace(op_tester):
    expand_smaller_output(op_tester, True, np.int64)
    expand_smaller_output(op_tester, True, np.int32)


def test_expand_non_one_smaller_output(op_tester):
    expand_non_one_smaller_output(op_tester, False, np.int64)
    expand_non_one_smaller_output(op_tester, False, np.int32)


def test_expand_non_one_smaller_output_inplace(op_tester):
    expand_non_one_smaller_output(op_tester, True, np.int64)
    expand_non_one_smaller_output(op_tester, True, np.int32)


def test_expand_scalar(op_tester):
    expand_scalar(op_tester, False, np.int64)
    expand_scalar(op_tester, False, np.int32)


def test_expand_scalar_inplace(op_tester):
    expand_scalar(op_tester, True, np.int64)
    expand_scalar(op_tester, True, np.int32)


def test_expand_unwind(op_tester):
    expand_unwind(op_tester, False, np.int64)
    expand_unwind(op_tester, False, np.int32)


#  Model1:            Model2:
#
#  i1   w             i1   w
#   |   |              |   |
#   |   |              |   |
#   |   |              |   |
#   matmul             matmul
#    |                  |
#    |       c          |   ones
#    |      /           |  /
#   expand -           mul-
#    |                  |
#    |                  |
#    |                  |
#   out                out
#
# Where ones = np.ones(shape=c)
@pytest.mark.parametrize("int_type", [np.int32, np.int64])
def test_expand_mul(int_type):
    in_ = np.random.random_sample([3, 1]).astype(np.float32)
    w = np.random.random_sample([3, 3]).astype(np.float32)
    const = np.array([2, 1, 6]).astype(int_type)

    def get_session(expand_op=True, device=None):
        builder = popart.Builder()

        i1 = builder.addInputTensor(popart.TensorInfo("FLOAT", in_.shape))
        w_in = builder.addInitializedInputTensor(w)
        mul = builder.aiOnnx.matmul([w_in, i1])

        if expand_op:
            c = builder.aiOnnx.constant(const)
            o = builder.aiOnnx.expand([mul, c])
        else:
            ones = np.ones(shape=const).astype(np.float32)
            ones_in = builder.aiOnnx.constant(ones)
            o = builder.aiOnnx.mul([mul, ones_in])

        builder.addOutputTensor(o)
        loss = builder.aiGraphcore.identityloss([o])

        anchor_returns = [
            o,
            popart.reservedGradientPrefix() + o,
            popart.reservedGradientPrefix() + i1,
            popart.reservedGradientPrefix() + w_in,
        ]

        opts = popart.SessionOptions()
        session = popart.TrainingSession(
            fnModel=builder.getModelProto(),
            dataFlow=popart.DataFlow(1, anchor_returns),
            deviceInfo=device,
            optimizer=popart.ConstSGD(0.1),
            loss=loss,
            userOptions=opts,
        )

        anchors = session.initAnchorArrays()
        inputs = {i1: in_}
        stepio = popart.PyStepIO(inputs, anchors)

        session.prepareDevice()
        session.weightsFromHost()
        return session, stepio, anchors, anchor_returns

    with tu.create_test_device() as d1, tu.create_test_device() as d2:
        sess1, stepio1, anch1, arts1 = get_session(True, d1)
        sess2, stepio2, anch2, arts2 = get_session(False, d2)

        for i in range(5):
            print(f"Step {i}")
            sess1.run(stepio1)
            sess2.run(stepio2)
            for (k1, k2) in zip(arts1, arts2):
                assert np.allclose(anch1[k1], anch2[k2])
