import csv
import string
from typing import Callable, List, Tuple

import gensim.models
import matplotlib.pyplot as plt
import numpy as np
from gensim.models.keyedvectors import KeyedVectors


def generate_2d_xor_dataset(d: float = 4, seed: int = 1) -> Tuple[np.ndarray,
                                                                  np.ndarray]:
    """
    Generates a 2D dataset with the following properties:

    Two classes, indicated by labels 0 and 1.

    The examples come from 4 clusters:
        1. A 2D normal distribution centered at (-d, d) with unit variances.
        2. A 2D normal distribution centered at (d, d) with unit variances.
        3. A 2D normal distribution centered at (-d, 3d) with unit variances.
        4. A 2D normal distribution centered at (d, 3d) with unit variances.

    Points from clusters 1 and 2 have label 1, and points from clusters 3 and 4
    have label 0.

    Args:
        d (float): Parameter determining the positions of the cluster centers.
        seed (int): The seed to use to set the randomness for np.random.

    Returns:
        Tuple[np.ndarray, np.ndarray]: The first output is an array containing
                                        all of the data points/examples, with
                                        shape (num_examples, 2).
                                        The second output is an array containing
                                        the corresponding labels (0 or 1),
                                        with shape
                                        (num_examples, 1).
    """
    np.random.seed(seed)

    cluster_means_to_label = {(-d, d):     0,
                              (d, d):      1,
                              (-d, 3 * d): 1,
                              (d, 3 * d):  0}
    covariance = np.array([[1, 0],
                           [0, 1]])

    points_per_cluster = 100

    clusters_x = []
    clusters_y = []

    for mean, label in cluster_means_to_label.items():
        clusters_x.append(np.random.multivariate_normal(mean, covariance,
                                                        points_per_cluster))
        clusters_y.append(np.full((points_per_cluster,), label))

    X = np.concatenate(clusters_x, axis=0)
    y = np.expand_dims(np.concatenate(clusters_y, axis=0), axis=1)

    return X, y


def plot_2d_dataset_points(X: np.ndarray, y: np.ndarray) -> None:
    """
    Plots the 2D toy dataset with labels.

    Args:
        X (np.ndarray): Array of the examples/data points of shape
                        (num_examples, 2). Generated by
                        generate_2d_xor_dataset().
        y (np.ndarray): Array of the corresponding labels of shape
                        (num_examples, 1). Generated by
                        generate_2d_xor_dataset().
    """
    # Get points in class 1 and class 0 separately
    X_c1 = X[y[:, 0] == 1]
    X_c0 = X[y[:, 0] == 0]

    # Plot them as colored points
    plt.scatter(X_c1[:, 0], X_c1[:, 1], c="Red", label="Class 1")
    plt.scatter(X_c0[:, 0], X_c0[:, 1], c="Blue", label="Class 0")
    # Add axes labels and legend
    plt.xlabel("Average Sentiment of Words in Sentence")
    plt.ylabel("Sarcasm level in Sentence")
    plt.legend()
    plt.show()


def plot_points_with_classifier_predictions(
        X: np.ndarray,
        y: np.ndarray,
        classifier,
        h: float = 0.02,
        alpha: float = 0.3) -> None:
    """
    Plots the data points X, labeling based on their true labels y.

    Then, overlays the prediction of the provided classifier over the data,
    showing the regions predicted as each class with shaded colors.

    Args:
        X (np.ndarray): Array of the examples/data points of shape
                        (num_examples, 2). Generated by
                        generate_2d_xor_dataset().
        y (np.ndarray): Array of the corresponding labels of shape
                        (num_examples, 1). Generated by
                        generate_2d_xor_dataset().
        classifier: An object which implements the method predict(X), which
                    takes in the array of inputs of shape (num_examples,
                    input_size) and outputs an array of predictions of
                    shape (num_examples,).
        h (float): Parameter controlling the density of sampling used to
                    produce the colored classifier prediction regions.
        alpha (float): Parameter controlling the opacity of the colored
                        classifier prediction regions.

    Credit to the excellent plotting example from Scikit-Learn here:
    http://scikit-learn.org/stable/auto_examples/classification/plot_classifier_comparison.html
    """

    # Plot the points in class 1 and class 0 as colored dots
    X_c1 = X[y[:, 0] == 1]
    X_c0 = X[y[:, 0] == 0]
    plt.scatter(X_c1[:, 0], X_c1[:, 1], c="Red", label="Class 1")
    plt.scatter(X_c0[:, 0], X_c0[:, 1], c="Blue", label="Class 0")
    plt.xlabel("Average Sentiment of Words in Sentence")
    plt.ylabel("Sarcasm level in Sentence")

    # Shade the plot based on the model's predictions
    x_min, x_max = X[:, 0].min() - 1, X[:, 0].max() + 1
    y_min, y_max = X[:, 1].min() - 1, X[:, 1].max() + 1
    xx, yy = np.meshgrid(np.arange(x_min, x_max, h),
                         np.arange(y_min, y_max, h))

    Z = classifier.predict(np.c_[xx.ravel(), yy.ravel()])
    Z = Z.reshape(xx.shape)
    plt.contourf(xx, yy, Z, cmap=plt.cm.coolwarm, alpha=alpha)

    plt.legend()
    plt.show()

def check_parameter_shapes(stacked_logreg_model_class):
    test_network = stacked_logreg_model_class(10, 20, 30)
    errors = []
    if test_network.W_1.shape != (10, 20):
        errors.append("W_1")
    if test_network.W_2.shape != (20, 30):
        errors.append("W_2")
    if test_network.W_3.shape != (30, 1):
        errors.append("W_3")
    if test_network.b_1.shape != (1, 20):
        errors.append("b_1")   
    if test_network.b_2.shape != (1, 30):
        errors.append("b_2")   
    if test_network.b_3.shape != (1, 1):
        errors.append("b_3")   
       
    if len(errors) == 0:
        print("All parameter shapes are correct.")
    else:
        print(", ".join(errors)+" have incorrect shapes!")
    
    W_errors = []
    b_errors = []
    for W, name in [(test_network.W_1, 'W_1'), (test_network.W_2, 'W_2'), (test_network.W_3, 'W_3')]:
        if not np.any(W):
            W_errors.append(name)
    
    for b, name in [(test_network.b_1, 'b_1'), (test_network.b_2, 'b_2'), (test_network.b_3, 'b_3')]:
        if np.any(b):
            b_errors.append(name)
    
    if len(W_errors) > 0:
        print(", ".join(W_errors) + ' should not be initialized as zero matrices! Please use the `initialize_weights()` method.')
    if len(b_errors) > 0:
        print(', '.join(b_errors) + ' should be initialized as zero matrices!')
        
def check_predict(stacked_logreg_model_class,
                  X: np.ndarray) -> None:

    # Seed so we always get the same results
    np.random.seed(1)

    test_network = stacked_logreg_model_class(2, 10, 1)
    _, _, raw_output = test_network.forward_pass(X)
    correct_output = np.ones(raw_output.shape) * (raw_output >= 0.5)
    label_output = test_network.predict(X)
    if label_output.shape != (X.shape[0], 1):
        print("Wrong prediction shape. Should be", (X.shape[0], 1), 
              "but got", label_output.shape, "instead!")
    elif np.array_equal(correct_output, label_output):
        print("Correct.")
    elif not issubclass(label_output.dtype.type, np.integer):
        print("You did not return an integer from predict().")
    else:
        print("Incorrect implementation of predict.")
    
    
def check_forward_pass(stacked_logreg_model_class,
                       X: np.ndarray,
                       tolerance: float = 1e-3) -> None:
    """
    Checks the StackedLogisticRegressionNetwork.forward_pass() method
    and prints the results.

    Args:
        stacked_logreg_model_class: The StackedLogisticRegressionNetwork
        class.
        X (np.ndarray): Array of the examples/data points of shape
                        (num_examples, 2). Generated by
                        generate_2d_xor_dataset().
        tolerance (float): Tolerance for equality comparison of the output
                            with the expected output.
    """

    # Seed so we always get the same results
    np.random.seed(1)

    test_network = stacked_logreg_model_class(2, 2, 2)
    a_1, a_2, a_3 = test_network.forward_pass(X[:3, :])

    test_network_outputs = {
        "a_1": a_1,
        "a_2": a_2,
        "a_3": a_3
    }

    # Pre-computed correct outputs for seed 1 (old default init seed)
    # seed_1_correct_outputs = {
    #     "a_1": np.array([[0.02494248, 0.05104596],
    #                      [0.06515092, 0.02061314],
    #                      [0.19114486, 0.07487885]]),
    #     "a_2": np.array([[0.48480197, 0.48894751],
    #                      [0.48195428, 0.48179426],
    #                      [0.44449568, 0.44541923]]),
    #     "a_3": np.array([[0.47804574],
    #                      [0.47805725],
    #                      [0.4797901]])
    # }
    
    # Pre-computed correct outputs for seed 1234 (current default init seed)
    correct_outputs = {
        "a_1": np.array([[0.78215436, 0.83993657],
                         [0.95143653, 0.66631161],
                         [0.89182168, 0.56208439]]), 
        "a_2": np.array([[0.51911413, 0.5462156 ], 
                         [0.57142199, 0.49067908],
                         [0.5753812 , 0.47972371]]), 
        "a_3": np.array([[0.77789913],
                         [0.77940367],
                         [0.77828089]])
    }

    for name, output in test_network_outputs.items():
        if np.linalg.norm(output - correct_outputs[name]) <= tolerance:
            print("Forward pass output {} is correct.".format(name))
        else:
            print("Forward pass output {} is incorrect!!".format(name))


def check_backward_pass(stacked_logreg_model_class,
                        X: np.ndarray,
                        y: np.ndarray,
                        tolerance: float = 1e-3) -> None:
    """
    Checks the StackedLogisticRegressionNetwork.backward_pass() method
    and prints the results.

    Args:
        stacked_logreg_model_class: The StackedLogisticRegressionNetwork
        class.
        X (np.ndarray): Array of the examples/data points of shape
                        (num_examples, 2). Generated by
                        generate_2d_xor_dataset().
        y (np.ndarray): Array of the corresponding labels of shape
                        (num_examples, 1). Generated by
                        generate_2d_xor_dataset().
        tolerance (float): Tolerance for equality comparison of the output
                            with the expected output.
    """
    # Seed so we always get the same results
    np.random.seed(1)

    test_network = stacked_logreg_model_class(2, 2, 2)
    a_1, a_2, a_3 = test_network.forward_pass(X[:3, :])
    test_network_gradients = test_network.backward_pass(X[:3, :], y[:3], a_1,
                                                        a_2, a_3)

    # Pre-computed correct outputs for seed 1 (old default init seed)
    # seed_1_correct_gradients = {
    #     "W_3": np.array([[0.22514129],
    #                      [0.22592408]]),
    #     "b_3": np.array([[0.47863103]]),
    #     "W_2": np.array([[-0.00325071, 0.00122269],
    #                      [-0.00169606, 0.00063797]]),
    #     "b_2": np.array([[-0.03476859, 0.01307702]]),
    #     "W_1": np.array([[-0.00463418, -0.00307033],
    #                      [0.00296584, 0.0024738]]),
    #     "b_1": np.array([[0.00135783, 0.00099955]])
    # }
    
    # Pre-computed correct outputs for seed 1234 (current default init seed)
    correct_gradients = {
        "W_3": np.array([[0.43233167],
                         [0.39356584]]), 
        "b_3": np.array([[0.77852789]]), 
        "W_2": np.array([[0.21729521, 0.18053744],
                      [0.17155242, 0.14209996]]), 
        "b_2": np.array([[0.24845651, 0.2062255]]), 
        "W_1": np.array([[-0.01709541, -0.01199752], 
                         [0.01662384, 0.00842318]]), 
        "b_1": np.array([[0.00588494, 0.00339362]])
    }

    for name, gradient in test_network_gradients.items():
        if np.linalg.norm(gradient - correct_gradients[name]) <= tolerance:
            print("Backward pass gradient {} is correct.".format(name))
        else:
            print("Backward pass gradient {} is incorrect!!".format(name))


def load_embeddings(embedding_filename: str) -> gensim.models.KeyedVectors:
    """
    Load the GloVe word embeddings from a file and add an "unknown"
    word embedding.

    Args:
        embedding_filename (str): The text file containing the embedding data,
                                    must be in the Word2Vec format (see
                                    Gensim documentation).
    Returns:
        gensim.models.KeyedVectors: A Gensim embeddings object.
    """
    embeddings = KeyedVectors.load_word2vec_format(
        embedding_filename, binary=False)

    # Create the "<unk>" token for unknown words and set it equal to the average
    # of all the other word embeddings.
    embeddings["<unk>"] = np.mean(embeddings.vectors, axis=0)
    return embeddings


def load_dataset(filename: str) -> Tuple[List[List[str]], List[int]]:
    """
    Load in a Yelp review dataset (pre-processed) from a custom-formatted
    CSV file.

    Args:
        filename (str): The text file containing the Yelp review data. Formatted
                        as a pipe ("|")-delimited CSV with a header row and two
                        columns ("Review", and "Label").
    Returns:
        Tuple[List[List[str]], List[int]]: The first output is a list of
        examples, where each example is a list of words and each word is a
        string. The second output is a list of labels for the corresponding
        examples, where each label is either 0 or 1. Both lists are of the
        same length (the number of examples).
    """
    examples = []
    labels = []
    with open(filename, "r") as f:
        reader = csv.DictReader(f, delimiter='|', quotechar='"')
        for row in reader:
            if not row["Review"].split():
                continue
            examples.append(row["Review"].split())
            labels.append(int(row["Label"]))

    return examples, labels


def check_examples_to_array(examples_to_array_fn: Callable,
                            embeddings: gensim.models.KeyedVectors,
                            tolerance: float = 1e-3) -> None:
    """
    Checks the examples_to_array() function and prints the results.

    Args:
        examples_to_array_fn (Callable): The function examples_to_array() to
        check.
        embeddings (gensim.models.KeyedVectors): The Gensim embeddings object,
        loaded using load_embeddings().
        tolerance (float): Tolerance for equality comparison of the output
                            with the expected output.
    """
    test_example = ["i", "loved", "the", "non_existant_word", "food", "it", "was", "fantastic"]

    test_outputs = {
        mode: examples_to_array_fn([test_example], embeddings, mode=mode) for
        mode in ["mean", "sum", "max"]
    }

    correct_outputs = {
        'mean': np.array(
            [[0.29803967, 0.04200414, -0.28900403, -0.2646109,
              0.5288827, -0.06764147, -0.54383063, -0.02759279]]),
        'sum':  np.array(
            [[2.3843174, 0.3360331, -2.3120322, -2.116887,
              4.2310615, -0.54113173, -4.350645, -0.22074229,]]),
        'max':  np.array(
            [[0.61183, 0.66826, -0.082073, 0.1217, 
              0.88331, 0.39783, 0.11975443, 0.66537]])
    }

    for mode, output in test_outputs.items():
        if np.linalg.norm(output[:, :8] - correct_outputs[mode]) <= tolerance:
            print("Output for mode {} is correct.".format(mode))
        else:
            print("Output for mode {} is incorrect!!".format(mode))


def check_forward_pass_yelp(yelp_network_class,
                            X: np.ndarray,
                            tolerance: float = 1e-3) -> None:
    """
    Checks the YelpClassificationNeuralNetwork.forward_pass() method
    and prints the results.

    Args:
        yelp_network_class: The YelpClassificationNeuralNetwork class.
        X (np.ndarray): Array of the examples/data points of shape
                        (num_examples, 2). Generated by
                        generate_2d_xor_dataset().
        tolerance (float): Tolerance for equality comparison of the output
                            with the expected output.
    """
    # Seed so we always get the same results
    np.random.seed(1)

    test_network = yelp_network_class(50, 2, 2)
    a_1, a_2, a_3 = test_network.forward_pass(X[:3, :])

    test_network_outputs = {
        "a_1": a_1,
        "a_2": a_2,
        "a_3": a_3
    }

    # Pre-computed correct outputs for seed 1 (old default init seed)
    # seed_1_correct_outputs = {
    #     'a_1': np.array([[0., 0.95457457],
    #                      [0., 0.85213306],
    #                      [0., 0.75179564]]),
    #     'a_2': np.array([[0.27235127, 0.40065728],
    #                      [0.24312351, 0.35766018],
    #                      [0.21449607, 0.31554622]]),
    #     'a_3': np.array([[0.45312089],
    #                      [0.45812682],
    #                      [0.46303814]])
    # }
    
    # Pre-computed correct outputs for seed 1234 (current default init seed)
    correct_outputs = {
        'a_1': np.array([[0., 0.4267531],
                         [0., 0.],
                         [0., 0.62810926]]),
        'a_2': np.array([[0.15196385, 0.08609083],
                         [0., 0.],
                         [0.2236654, 0.12671132]]),
        'a_3': np.array([[0.53890433],
                         [0.5],
                         [0.55712612]])
    }

    for name, output in test_network_outputs.items():
        if np.linalg.norm(output - correct_outputs[name]) <= tolerance:
            print("Forward pass output {} is correct.".format(name))
        else:
            print("Forward pass output {} is incorrect!!".format(name))


def check_backward_pass_yelp(yelp_network_class,
                             X: np.ndarray,
                             y: np.ndarray,
                             tolerance: float = 1e-3) -> None:
    """
    Checks the YelpClassificationNeuralNetwork.backward_pass() method
    and prints the results.

    Args:
        yelp_network_class: The YelpClassificationNeuralNetwork class.
        X (np.ndarray): Array of the examples/data points of shape
                        (num_examples, 2). Generated by
                        generate_2d_xor_dataset().
        y (np.ndarray): Array of the corresponding labels of shape
                        (num_examples, 1). Generated by
                        generate_2d_xor_dataset().
        tolerance (float): Tolerance for equality comparison of the output
                            with the expected output.
    """
    # Seed so we always get the same results
    np.random.seed(2)

    test_network = yelp_network_class(50, 2, 2)
    a_1, a_2, a_3 = test_network.forward_pass(X[:3, :])
    test_network_gradients = test_network.backward_pass(
        X[:3, :], y[:3], a_1, a_2, a_3)

    # Pre-computed correct outputs for seed 1 (old default init seed)
    # correct_gradients = {
    #     "W_3": np.array([-0.13195385]),
    #     "b_3": np.array([-0.54190472]),
    #     "W_2": np.array([0., 0.]),
    #     "b_2": np.array([0.48445836, -0.07494552]),
    #     "W_1": np.array([0., 0.03387387]),
    #     "b_1": np.array([0., 0.10676524])
    # }
    
    # Pre-computed correct outputs for seed 1234 (current default init seed)
    correct_gradients = {
        "W_3": np.array([-0.05637515]),
        "b_3": np.array([-0.46798985]),
        "W_2": np.array([0., 0.]),
        "b_2": np.array([-0.11914909, -0.33545647]),
        "W_1": np.array([0., -0.03424743]),
        "b_1": np.array([0., -0.11010132])
    }

    for name, gradient in test_network_gradients.items():
        if gradient.ndim > 1:
            gradient = gradient[0]
        if np.linalg.norm(gradient - correct_gradients[name]) <= tolerance:
            print("Backward pass gradient {} is correct.".format(name))
        else:
            print("Backward pass gradient {} is incorrect!!".format(name))


def evaluate(model,
             examples_to_array_fn: Callable,
             embeddings: gensim.models.KeyedVectors,
             review: str) -> None:
    """
    Receives a test review as a string and prints the trained model's
    predictions.

    Args:
        model: A YelpClassificationNeuralNetwork instance.
        examples_to_array_fn (Callable): The examples_to_array() function.
        embeddings (gensim.models.KeyedVectors): The Gensim embeddings object,
            loaded using load_embeddings().
        review (str): A sample "Yelp review" as a string.
    """
    print("Review: '{}'\n".format(review))
    # Tokenize, strip punctuation, to lower case
    review = review.translate(str.maketrans('', '', string.punctuation))
    review_words = review.lower().split()

    X_array = examples_to_array_fn([review_words], embeddings)

    _, _, predicted_scores = model.forward_pass(X_array)

    print("Predicted scores (class probabilities):")
    print("0: {}%".format((1.0 - predicted_scores[0][0]) * 100))
    print("1: {}%".format(predicted_scores[0][0] * 100))

    print("\nPredicted sentiment: {}".format(
        "Positive" if model.predict(X_array)[0][0] == 1 else "Negative"))
