"""Compare seeded initial embeddings with the saved Word2Vec input vectors.

Usage: python word2vec-compare.py /path/to/02-word2vec/data [--plot out.png]
Requires NumPy (and Matplotlib for --plot) and the C.npy/vocab.txt produced by word2vec.py at:
https://github.com/PragalvaXFREZ/embedding-models-from-scratch/blob/0ea72056dafa78e06553cf3c9f4192e83f841b49/02-word2vec/word2vec.py

The trainer draws C twice from default_rng(0): once for the commented-out
baseline, then again for hierarchical softmax. Both draws must be reproduced.
This script reads saved data and does not train or overwrite any files.
"""

import argparse
from pathlib import Path

import numpy as np


def cosine(a, b):
    return float(a @ b / (np.linalg.norm(a) * np.linalg.norm(b)))


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("data_dir", type=Path)
    parser.add_argument("--plot", type=Path, help="also save a before/after chart")
    args = parser.parse_args()
    trained = np.load(args.data_dir / "C.npy", allow_pickle=False)
    vocab = (args.data_dir / "vocab.txt").read_text().splitlines()
    if trained.shape != (21_681, 100) or len(vocab) != trained.shape[0]:
        raise ValueError("Expected the article's 21,681-word, 100-dimensional run")
    word2id = {word: i for i, word in enumerate(vocab)}

    rng = np.random.default_rng(0)
    rng.random(trained.shape)  # Baseline draw advances the generator.
    initial = (rng.random(trained.shape) - 0.5) / 100

    print("Cosine similarity of input embeddings")
    print("Before: seeded initialization (reconstructed)")
    print("After:  saved C.npy, one training pass")
    print()
    print(f'{"Pair":<18} {"Before":>8} {"After":>8}')
    rows = []
    pairs = [("three", "four"), ("three", "five"), ("three", "six"),
             ("france", "spain"), ("france", "italy"),
             ("king", "prince"), ("computer", "software"),
             ("three", "computer")]
    for a, b in pairs:
        i, j = word2id[a], word2id[b]
        before = cosine(initial[i], initial[j])
        after = cosine(trained[i], trained[j])
        label = f"{a} / {b}"
        rows.append((label, before, after))
        print(f"{label:<18} {before:>8.3f} {after:>8.3f}")

    if args.plot:
        plot(rows, args.plot)


def plot(rows, out):
    import matplotlib
    matplotlib.use("Agg")
    import matplotlib.pyplot as plt

    fig, ax = plt.subplots(figsize=(8, 4.8))
    for y, (label, before, after) in enumerate(reversed(rows)):
        ax.annotate("", xy=(after, y), xytext=(before, y),
                    arrowprops=dict(arrowstyle="->", color="gray", lw=1.2))
        ax.plot(before, y, "o", color="tab:blue", label="before training" if y == 0 else None)
        ax.plot(after, y, "o", color="tab:orange", label="after one pass" if y == 0 else None)
        ax.text(max(before, after) + 0.03, y, f"{after:.3f}", va="center", fontsize=9)
    ax.set_yticks(range(len(rows)), [label for label, _, _ in reversed(rows)])
    ax.axvline(0, color="gray", lw=0.8)
    ax.set_xlim(-0.2, 1.1)
    ax.set_xlabel("cosine similarity of input embeddings")
    ax.legend(loc="lower right")
    fig.tight_layout()
    fig.savefig(out, dpi=150)


if __name__ == "__main__":
    main()
