import os
import math
import random
import urllib.request
import torch
import torch.nn as nn
import torch.nn.functional as F
import matplotlib.pyplot as plt
from torch.utils.data import Dataset, DataLoader



#-------------------------------------------------- USER SETTINGS --------------------------------------------------
SEED = 2026
EPOCHS = 10
BATCH_SIZE = 64
LEARNING_RATE = 1e-3
TRAIN_SPLIT = 0.8
SUBSET_CHARS = 300000
EMBEDDING_DIM = 64
DEFAULT_SEQUENCE_LENGTH = 50
DEFAULT_HIDDEN_SIZE = 64
NUM_LAYERS = 1
TEMPERATURE = 1.0
GENERATION_STEPS = 400
SEED_TEXT = "to be or not to be"
SAVE_PLOTS = True
OUTPUT_DIR = "hw4_outputs"
DATA_DIR = "data"
DATA_FILE = os.path.join(DATA_DIR, "tinyshakespeare.txt")
TINY_SHAKESPEARE_URL = "https://raw.githubusercontent.com/karpathy/char-rnn/master/data/tinyshakespeare/input.txt"



#-------------------------------------------------- SEED EVERYTHING --------------------------------------------------
def seed_everything(seed):
    random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    torch.backends.cudnn.deterministic = True
    torch.backends.cudnn.benchmark = False



#-------------------------------------------------- LOAD TINY SHAKESPEARE --------------------------------------------------
def download_tiny_shakespeare():
    os.makedirs(DATA_DIR, exist_ok=True)

    try:
        urllib.request.urlretrieve(TINY_SHAKESPEARE_URL, DATA_FILE)
    except Exception as error:
        raise RuntimeError("Could not download Tiny Shakespeare automatically") from error



def load_tiny_shakespeare(subset_chars=SUBSET_CHARS):
    download_tiny_shakespeare()

    with open(DATA_FILE, "r", encoding="utf-8") as file:
        text = file.read().lower()

    if subset_chars is not None:
        text = text[:subset_chars]

    print("Number of characters used:", len(text))

    return text



#-------------------------------------------------- VOCABULARY --------------------------------------------------
def create_vocabulary(text):
    characters = sorted(list(set(text)))
    char2idx = {char: idx for idx, char in enumerate(characters)}
    idx2char = {idx: char for char, idx in char2idx.items()}

    print("Vocabulary size:", len(characters))

    return char2idx, idx2char



def encode_text(text, char2idx):
    return torch.tensor([char2idx[char] for char in text], dtype=torch.long)



#-------------------------------------------------- SHAKESPEARE DATASET --------------------------------------------------
class ShakespeareDataset(Dataset):
    def __init__(self, encoded_text, sequence_length=DEFAULT_SEQUENCE_LENGTH):
        self.encoded_text = encoded_text
        self.sequence_length = sequence_length

    def __len__(self):
        return len(self.encoded_text) - self.sequence_length

    def __getitem__(self, idx):
        current_characters = self.encoded_text[idx:idx+self.sequence_length]
        next_characters = self.encoded_text[idx+1:idx+self.sequence_length+1]

        return current_characters, next_characters



#-------------------------------------------------- CREATE DATALOADERS --------------------------------------------------
def get_dataloaders(sequence_length=DEFAULT_SEQUENCE_LENGTH, batch_size=BATCH_SIZE, subset_chars=SUBSET_CHARS):
    text = load_tiny_shakespeare(subset_chars=subset_chars)
    char2idx, idx2char = create_vocabulary(text)
    encoded_text = encode_text(text, char2idx)

    split_idx = int(TRAIN_SPLIT * len(encoded_text))
    train_encoded = encoded_text[:split_idx]
    val_encoded = encoded_text[split_idx:]

    train_dataset = ShakespeareDataset(train_encoded, sequence_length=sequence_length)
    val_dataset = ShakespeareDataset(val_encoded, sequence_length=sequence_length)

    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
    val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)

    print("Sequence length:", sequence_length)
    print("Training sequences:", len(train_dataset))
    print("Validation sequences:", len(val_dataset))
    print("Training batch shape should be: (batch size, sequence length)")

    return train_loader, val_loader, char2idx, idx2char



#-------------------------------------------------- VANILLA RNN MODEL --------------------------------------------------
class VanillaRNN(nn.Module):
    def __init__(self, vocab_size, embedding_dim=EMBEDDING_DIM, hidden_size=DEFAULT_HIDDEN_SIZE, num_layers=NUM_LAYERS):
        super().__init__()

        self.hidden_size = hidden_size
        self.num_layers = num_layers

        self.embedding = nn.Embedding(vocab_size, embedding_dim)
        self.rnn = nn.RNN(embedding_dim, hidden_size, num_layers=num_layers, batch_first=True)
        self.output_layer = nn.Linear(hidden_size, vocab_size)

    def forward(self, x, hidden=None):
        embedded = self.embedding(x)
        output, hidden = self.rnn(embedded, hidden)
        logits = self.output_layer(output)

        return logits, hidden

    def count_parameters(self):
        return sum(p.numel() for p in self.parameters() if p.requires_grad)



#-------------------------------------------------- LSTM MODEL --------------------------------------------------
class LSTMModel(nn.Module):
    def __init__(self, vocab_size, embedding_dim=EMBEDDING_DIM, hidden_size=DEFAULT_HIDDEN_SIZE, num_layers=NUM_LAYERS):
        super().__init__()

        self.hidden_size = hidden_size
        self.num_layers = num_layers

        self.embedding = nn.Embedding(vocab_size, embedding_dim)
        self.lstm = nn.LSTM(embedding_dim, hidden_size, num_layers=num_layers, batch_first=True)
        self.output_layer = nn.Linear(hidden_size, vocab_size)

    def forward(self, x, hidden=None):
        embedded = self.embedding(x)
        output, hidden = self.lstm(embedded, hidden)
        logits = self.output_layer(output)

        return logits, hidden

    def count_parameters(self):
        return sum(p.numel() for p in self.parameters() if p.requires_grad)



#-------------------------------------------------- TRAIN ONE EPOCH --------------------------------------------------
def train_one_epoch(model, train_loader, optimizer, criterion, device):
    model.train()
    total_loss = 0

    for inputs, targets in train_loader:
        inputs, targets = inputs.to(device), targets.to(device)

        optimizer.zero_grad()

        outputs, _ = model(inputs, hidden=None)

        outputs = outputs.reshape(-1, outputs.shape[-1])
        targets = targets.reshape(-1)

        loss = criterion(outputs, targets)
        loss.backward()
        optimizer.step()

        total_loss += loss.item()

    return total_loss / len(train_loader)



#-------------------------------------------------- EVALUATION LOOP --------------------------------------------------
def evaluate(model, val_loader, criterion, device):
    model.eval()
    total_loss = 0

    with torch.no_grad():
        for inputs, targets in val_loader:
            inputs, targets = inputs.to(device), targets.to(device)

            outputs, _ = model(inputs, hidden=None)
            outputs = outputs.reshape(-1, outputs.shape[-1])
            targets = targets.reshape(-1)

            loss = criterion(outputs, targets)
            total_loss += loss.item()

    avg_loss = total_loss / len(val_loader)
    perplexity = math.exp(avg_loss)

    return avg_loss, perplexity



#-------------------------------------------------- TRAINING FUNCTION --------------------------------------------------
def train_model(model, train_loader, val_loader, lr=LEARNING_RATE, epochs=EPOCHS, device="cpu", debug=True):
    model = model.to(device)
    criterion = nn.CrossEntropyLoss()
    optimizer = torch.optim.Adam(model.parameters(), lr=lr)

    train_losses = []
    val_losses = []
    val_perplexities = []

    for epoch in range(epochs):
        train_loss = train_one_epoch(model, train_loader, optimizer, criterion, device)
        val_loss, val_perplexity = evaluate(model, val_loader, criterion, device)

        train_losses.append(train_loss)
        val_losses.append(val_loss)
        val_perplexities.append(val_perplexity)

        if debug:
            print(f"[{epoch+1}/{epochs}] Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | Val Perplexity: {val_perplexity:.2f}")

    return train_losses, val_losses, val_perplexities



#-------------------------------------------------- TEXT GENERATION --------------------------------------------------
def generate_text(model, seed_text, char2idx, idx2char, steps=GENERATION_STEPS, temperature=TEMPERATURE, device="cpu"):
    model.eval()
    generated_text = seed_text.lower()

    seed_indices = [char2idx[char] for char in generated_text if char in char2idx]

    if len(seed_indices) == 0:
        seed_indices = [0]
        generated_text = idx2char[0]

    current_input = torch.tensor(seed_indices, dtype=torch.long).unsqueeze(0).to(device)
    hidden = None

    with torch.no_grad():
        for _ in range(steps):
            outputs, hidden = model(current_input, hidden)

            last_step_logits = outputs[:, -1, :] / temperature
            probabilities = F.softmax(last_step_logits, dim=-1)
            next_idx = torch.multinomial(probabilities, num_samples=1)

            next_char = idx2char[next_idx.item()]
            generated_text += next_char

            current_input = next_idx.reshape(1, 1)

    return generated_text



#-------------------------------------------------- SAVE GENERATED TEXT --------------------------------------------------
def save_generated_text(experiment_name, sample_dict):
    os.makedirs(OUTPUT_DIR, exist_ok=True)
    output_path = os.path.join(OUTPUT_DIR, "generated_text_samples.txt")

    with open(output_path, "a", encoding="utf-8") as file:
        file.write("\n" + "="*90 + "\n")
        file.write(experiment_name + "\n")
        file.write("="*90 + "\n")

        for label, sample in sample_dict.items():
            file.write("\n" + label + "\n")
            file.write("-"*90 + "\n")
            file.write(sample + "\n")



#-------------------------------------------------- PLOTTING --------------------------------------------------
def plot_curves(curve_dict, title, xlabel="Epoch", ylabel="Loss", filename=None):
    plt.figure()

    for label, values in curve_dict.items():
        plt.plot(range(1, len(values)+1), values, label=label)

    plt.xlabel(xlabel)
    plt.ylabel(ylabel)
    plt.title(title)
    plt.legend()
    plt.tight_layout()

    if SAVE_PLOTS and filename is not None:
        os.makedirs(OUTPUT_DIR, exist_ok=True)
        plt.savefig(os.path.join(OUTPUT_DIR, filename), dpi=300)
        print("Saved plot:", os.path.join(OUTPUT_DIR, filename))

    plt.show(block=False)



#-------------------------------------------------- EXPERIMENT A: RNN VS LSTM --------------------------------------------------
def run_experiment_A(debug=True):
    print("========================= EXPERIMENT A: RNN VS LSTM COMPARISON =========================")
    seed_everything(SEED)
    device = "cuda" if torch.cuda.is_available() else "cpu"
    print("Using device:", device)

    train_loader, val_loader, char2idx, idx2char = get_dataloaders(sequence_length=50)
    vocab_size = len(char2idx)

    print("----------------------------------------------------------------------------------------")
    print("Training Vanilla RNN")
    seed_everything(SEED)
    rnn_model = VanillaRNN(vocab_size=vocab_size, hidden_size=64)
    print(f"RNN parameters: {rnn_model.count_parameters()}")
    rnn_train, rnn_val, rnn_ppl = train_model(rnn_model, train_loader, val_loader, lr=1e-3, epochs=EPOCHS, device=device, debug=debug)
    rnn_sample = generate_text(rnn_model, SEED_TEXT, char2idx, idx2char, device=device)

    print("----------------------------------------------------------------------------------------")
    print("Training LSTM")
    seed_everything(SEED)
    lstm_model = LSTMModel(vocab_size=vocab_size, hidden_size=64)
    print(f"LSTM parameters: {lstm_model.count_parameters()}")
    lstm_train, lstm_val, lstm_ppl = train_model(lstm_model, train_loader, val_loader, lr=1e-3, epochs=EPOCHS, device=device, debug=debug)
    lstm_sample = generate_text(lstm_model, SEED_TEXT, char2idx, idx2char, device=device)

    #------------------------- TRAINING LOSS -------------------------
    plot_curves({"RNN": rnn_train, "LSTM": lstm_train}, "Experiment A: RNN vs LSTM Training Loss", ylabel="Training Loss", filename="experiment_A_training_loss.png")

    #------------------------- VALIDATION PERPLEXITY -------------------------
    plot_curves({"RNN": rnn_ppl, "LSTM": lstm_ppl}, "Experiment A: RNN vs LSTM Validation Perplexity", ylabel="Validation Perplexity", filename="experiment_A_validation_perplexity.png")

    #------------------------- GENERATED TEXT -------------------------
    print("\nRNN Generated Text:")
    print(rnn_sample)
    print("\nLSTM Generated Text:")
    print(lstm_sample)

    save_generated_text("Experiment A: RNN vs LSTM", {"RNN Generated Text": rnn_sample, "LSTM Generated Text": lstm_sample})

    print("========================================================================================")



#-------------------------------------------------- EXPERIMENT B: SEQUENCE LENGTH EFFECT --------------------------------------------------
def run_experiment_B(debug=True):
    print("========================= EXPERIMENT B: SEQUENCE LENGTH EFFECT =========================")
    seed_everything(SEED)
    device = "cuda" if torch.cuda.is_available() else "cpu"
    print("Using device:", device)

    sequence_lengths = [25, 50]
    train_loss_dict = {}

    for sequence_length in sequence_lengths:
        print("----------------------------------------------------------------------------------------")
        print(f"Training RNN with sequence length = {sequence_length}")
        seed_everything(SEED)
        train_loader, val_loader, char2idx, idx2char = get_dataloaders(sequence_length=sequence_length)
        vocab_size = len(char2idx)

        model = VanillaRNN(vocab_size=vocab_size, hidden_size=64)
        train_losses, _, _ = train_model(model, train_loader, val_loader, lr=1e-3, epochs=EPOCHS, device=device, debug=debug)
        train_loss_dict[f"Sequence Length {sequence_length}"] = train_losses

    #------------------------- TRAINING LOSS -------------------------
    plot_curves(train_loss_dict, "Experiment B: Sequence Length Comparison", ylabel="Training Loss", filename="experiment_B_training_loss.png")

    print("========================================================================================")



#-------------------------------------------------- EXPERIMENT C: HIDDEN SIZE EFFECT --------------------------------------------------
def run_experiment_C(debug=True):
    print("========================= EXPERIMENT C: HIDDEN SIZE EFFECT =========================")
    seed_everything(SEED)
    device = "cuda" if torch.cuda.is_available() else "cpu"
    print("Using device:", device)

    train_loader, val_loader, char2idx, idx2char = get_dataloaders(sequence_length=50)
    vocab_size = len(char2idx)

    hidden_sizes = [32, 64]
    train_loss_dict = {}
    perplexity_dict = {}
    samples = {}

    for hidden_size in hidden_sizes:
        print("----------------------------------------------------------------------------------------")
        print(f"Training LSTM with hidden size = {hidden_size}")
        seed_everything(SEED)
        model = LSTMModel(vocab_size=vocab_size, hidden_size=hidden_size)
        print(f"LSTM parameters: {model.count_parameters()}")

        train_losses, _, val_perplexities = train_model(model, train_loader, val_loader, lr=1e-3, epochs=EPOCHS, device=device, debug=debug)
        sample = generate_text(model, SEED_TEXT, char2idx, idx2char, device=device)

        train_loss_dict[f"Hidden Size {hidden_size}"] = train_losses
        perplexity_dict[f"Hidden Size {hidden_size}"] = val_perplexities
        samples[f"LSTM Hidden Size {hidden_size} Generated Text"] = sample

    #------------------------- TRAINING LOSS -------------------------
    plot_curves(train_loss_dict, "Experiment C: Hidden Size Comparison - Training Loss", ylabel="Training Loss", filename="experiment_C_training_loss.png")

    #------------------------- VALIDATION PERPLEXITY -------------------------
    plot_curves(perplexity_dict, "Experiment C: Hidden Size Comparison - Validation Perplexity", ylabel="Validation Perplexity", filename="experiment_C_validation_perplexity.png")

    #------------------------- GENERATED TEXT -------------------------
    for label, sample in samples.items():
        print("\n" + label + ":")
        print(sample)

    save_generated_text("Experiment C: Hidden Size Effect", samples)

    print("====================================================================================")



#-------------------------------------------------- MAIN --------------------------------------------------
if __name__ == "__main__":
    seed_everything(SEED)
    debug = True

    if os.path.exists(os.path.join(OUTPUT_DIR, "generated_text_samples.txt")):
        os.remove(os.path.join(OUTPUT_DIR, "generated_text_samples.txt"))

    run_experiment_A(debug)
    run_experiment_B(debug)
    run_experiment_C(debug)

    plt.show()
