import numpy as np
import pymc3 as pm
import pickle


def normalize(vector):
    return vector / sum(vector)


def geometric_mean(vector):
    """takes a numpy array object and calculates geometric mean"""
    return vector.prod() ** (1.0 / len(vector))


# noinspection PyTypeChecker
def priority_vector(matrix):
    return normalize([geometric_mean(row) for row in matrix])


def theoretical_matrix(pv, logs=True):
    n_rows = n_cols = len(pv)
    mu = np.empty((n_rows, n_cols))
    for i in range(n_rows):
        for j in range(n_cols):
            if logs:
                mu[i, j] = pv[i] - pv[j]
            else:
                mu[i, j] = pv[i] / pv[j]
    return mu


def ensure_reciprocals(mu_mat, logs=True):
    n_rows, n_cols = mu_mat.shape
    for i in range(n_rows):
        for j in range(n_cols):
            if i < j:
                mu_mat[i, j] = mu_mat[i, j]
            elif i > j:
                if logs:
                    mu_mat[i, j] = - mu_mat[j, i]
                else:
                    mu_mat[i, j] = 1 / mu_mat[j, i]
            else:  # i == j
                if logs:
                    mu_mat[i, j] = 0
                else:
                    mu_mat[i, j] = 1
    return mu_mat


def rank_data(vector):
    seq = sorted(vector, reverse=True)
    index = [seq.index(elem) + 1 for elem in vector]
    return index


def triangle_indices(matrix, triangle="lower"):
    """ Returns list of upper / lower triangle indices for a given matrix """
    if triangle == "lower":
        return list(zip(*np.tril_indices_from(matrix, k=-1)))
    elif triangle == "upper":
        return list(zip(*np.triu_indices_from(matrix, k=1)))
    elif triangle == "both":
        return triangle_indices(matrix, triangle="lower") + triangle_indices(matrix, triangle="upper")
    else:
        raise ValueError


def readable_matrix(matrix):
    return np.round(matrix, 3)


def numpy_to_latex(matrix):
    """Returns a LaTeX bmatrix
    :a: numpy array
    :returns: LaTeX bmatrix as a string
    """
    if len(matrix.shape) > 2:
        raise ValueError('bmatrix can at most display two dimensions')
    lines = str(matrix).replace('[', '').replace(']', '').splitlines()
    rv = [r'\begin{bmatrix}']
    rv += ['  ' + ' & '.join(l.split()) + r'\\' for l in lines]
    rv += [r'\end{bmatrix}']
    return '\n'.join(rv)


def save_pymc_trace(path, model, trace):
    """
    https://stackoverflow.com/questions/44764932/can-a-pymc3-trace-be-loaded-and-values-accessed-without-the-original-model-in-me
    """
    with open(path, 'wb') as buff:
        pickle.dump({'model': model, 'trace': trace}, buff)


def load_pymc_trace(path):
    with open(path, 'rb') as buff:
        data = pickle.load(buff)

    model, trace = data['model'], data['trace']
    return model, trace


def inference_step(model, samples=22_000, burn=2_000, cores=4):
    with model:
        trace = pm.sample(samples, cores=cores, tune=burn, progressbar=False)
    return trace
