import os
import torch
import torch.nn.functional as F
from inference_utils import device, GPTLITE_CKPT_PATH, GPTLITE_DISTILLED_CKPT_PATH  # also adds GPTlite to the path
from gptlite import GPTlite
from utils import get_batch, get_tiny_shakespeare_data, get_gptlite_model_parameters, get_gptlite_distilled_model_parameters


def distillation_loss(logits_student, logits_teacher, temperature):
  """ KL divergence between the softened teacher and student distributions, averaged per token, and
      scaled by temperature^2 so that the size of the gradients does not depend on the temperature """
  vocab_size = logits_student.size(-1)
  log_softmax_student = F.log_softmax(logits_student.reshape(-1, vocab_size)/temperature, dim=-1)   #log softmax of student model
  log_softmax_teacher = F.log_softmax(logits_teacher.reshape(-1, vocab_size)/temperature, dim=-1)   #log softmax of teacher model
  return F.kl_div(log_softmax_student, log_softmax_teacher, log_target=True, reduction='batchmean') * (temperature ** 2)


if __name__=='__main__':
  torch.manual_seed(42) # random seed, for reproducibility
  vocab_size, train_data, valid_data, _, decode_fn = get_tiny_shakespeare_data()

  # Train parameters
  eval_interval = 100  # evaluation interval
  train_iters = 50000  # number of training iterations
  temperature = 2 # temperature for distillation

  # Load pre-trained large model parameters
  torch.manual_seed(42) # random seed for model initialization, for reproducibility
  n_layers, d_model, n_heads, d_head, batch_size, lr, seqlen, dropout_p = get_gptlite_model_parameters()
  model_teacher = GPTlite(vocab_size, d_model, n_heads, d_head, n_layers, dropout_p, seqlen).to(device).eval()
  if os.path.exists(GPTLITE_CKPT_PATH):
    model_teacher.load_state_dict(torch.load(GPTLITE_CKPT_PATH, map_location=device))
    print(f"Loaded model from {GPTLITE_CKPT_PATH} into GPTlite")

  n_layers, d_model, n_heads, d_head, batch_size, lr, seqlen, dropout_p = get_gptlite_distilled_model_parameters()
  model_student = GPTlite(vocab_size, d_model, n_heads, d_head, n_layers, dropout_p, seqlen).to(device)
  optimizer = torch.optim.Adam(model_student.parameters(), lr=lr)

  # train the model
  for step in range(1, train_iters+1):

    # train step
    model_student.train()
    idx, _ = get_batch(train_data, batch_size=batch_size, seqlen=seqlen)   #get a batch of training data
    idx = idx.to(device) #move data to GPU
    logits_student = model_student(idx)   #forward pass
    with torch.no_grad():
      logits_teacher = model_teacher(idx)   #forward pass of the frozen teacher, without gradients
    loss = distillation_loss(logits_student, logits_teacher, temperature)  #compute KL divergence loss
    loss.backward()   #backward pass
    torch.nn.utils.clip_grad_norm_(model_student.parameters(), max_norm=1.0) # gradient clipping to avoid exploding gradients
    optimizer.step()   #update parameters
    optimizer.zero_grad(set_to_none=True)  #sets to None instead of 0, to save memory
    if step % 10 == 0:
        print(f"Train step {step}, loss {loss.item():.4f}")

    if step % eval_interval > 0:
        continue

    # evaluation step
    model_student.eval()
    with torch.inference_mode():

      # perform one eval step and compute loss
      idx, _ = get_batch(valid_data, batch_size=batch_size, seqlen=seqlen)
      idx = idx.to(device) #move data to GPU
      logits_student = model_student(idx)   #forward pass
      logits_teacher = model_teacher(idx)   
      loss = distillation_loss(logits_student, logits_teacher, temperature)
      print(f"Eval step {step}, eval loss {loss.item():.4f}")

      # Generate a sentence from the current state of the model
      # Begin of String (batch size 1): the token with id 0 (the \n character)
      idx = torch.zeros((1,1), dtype=torch.long, device=device)
      generated_seqlen = 100

      for _ in range(generated_seqlen):
        idx_cond = idx[:, -seqlen:] #crop the context to the last seqlen tokens
        logits = model_student(idx_cond) #call fwd without targets
        logits = logits[:, -1] # take last token. shape (B=1, C)
        probs = F.softmax(logits, dim=-1) # shape (B=1, C)
        # sample the next token (we could take instead the argmax, but that would be deterministic and boring)
        idx_next = torch.multinomial(probs, num_samples=1) # shape (B=1, 1)
        # append next token idx to the solution sequence so far
        idx = torch.cat([idx, idx_next], dim=-1) # shape (B=1, T+1)

      # print generated string
      idx_without_bos = idx[0, 1:].tolist() # remove batch dim and BOS token (\n)
      print("Generated text:", decode_fn(idx_without_bos))

      # Now save the model
      torch.save(model_student.state_dict(), GPTLITE_DISTILLED_CKPT_PATH)
      print(f"Model saved to {GPTLITE_DISTILLED_CKPT_PATH}")

