#!/usr/bin/env python3

import os
import csv
import yaml
import argparse
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader
from sklearn.metrics import accuracy_score

from datasets.hmdb51 import HMDB51Dataset
from datasets.ucf101 import UCF101Dataset

from models.models import r3d_18, r2plus1d_18, mc3_18
from utils.lut import load_kernel_lut, estimate_model_cost
from utils.battery import VirtualBattery
from utils.primal_dual import PrimalDualOptimizer
from utils.controller import TemporalController


# ------------------------------------------------------------
# RUN DIRECTORY
# ------------------------------------------------------------

def make_run_dirs(root, dataset, model, split, budget, optimizer):

    run_root = os.path.join(
        root,
        dataset.upper(),
        model.upper(),
        f"split{split}",
        f"budget_{budget:.5f}",
        optimizer.lower()
    )

    os.makedirs(run_root, exist_ok=True)

    metrics_dir = os.path.join(run_root, "metrics")
    os.makedirs(metrics_dir, exist_ok=True)

    return run_root, metrics_dir


# ------------------------------------------------------------
# EVALUATION
# ------------------------------------------------------------

def evaluate(model, loader, device):

    model.eval()

    preds_all = []
    labels_all = []

    with torch.no_grad():

        for video, label in loader:

            video = video.to(device)
            label = label.to(device)

            logits = model(video)
            preds = torch.argmax(logits, dim=1)

            preds_all.extend(preds.cpu().tolist())
            labels_all.extend(label.cpu().tolist())

    return accuracy_score(labels_all, preds_all)


# ------------------------------------------------------------
# TRAIN ONE EPOCH
# ------------------------------------------------------------

def train_one_epoch(
    model, controller, optimizer, scaler,
    loader, lut, battery, dual_opt,
    device, cfg
):

    model.train()

    total_loss = 0
    total_correct = 0
    total_samples = 0

    sum_latency = 0
    sum_energy = 0

    # Action counters
    action_counts = [0,0,0]

    for video, label in loader:

        video = video.to(device)
        label = label.to(device)

        bs = video.size(0)
        total_samples += bs

        # RL STATE

        with torch.no_grad():

            state = torch.mean(video, dim=[2,3,4])

            state = F.adaptive_avg_pool1d(
                state.unsqueeze(1),
                cfg["controller_state_dim"]
            ).squeeze(1)

        action, logprob, _ = controller.select_action(
            state,
            battery.get_battery_level(),
            None
        )

        action_flat = action.view(-1)
        a = int(torch.mode(action_flat).values.item())

        action_counts[a] += 1

        # TEMPORAL REDUCTION

        if a == 1:
            video = video[:,:,::2]

        elif a == 2:
            video = video[:,:,::4]

        # COST ESTIMATION

        latency, energy = estimate_model_cost(model, lut, video)

        battery.update(energy)

        sum_latency += latency
        sum_energy += energy

        optimizer.zero_grad()

        with torch.amp.autocast("cuda"):

            logits = model(video)
            ce_loss = F.cross_entropy(logits, label)

        preds = torch.argmax(logits, dim=1)

        correct = (preds == label).sum().item()
        total_correct += correct

        acc = correct / bs

        lambda_L, lambda_E = dual_opt.get_lambdas()

        lambda_L = min(lambda_L,20)
        lambda_E = min(lambda_E,20)

        reward = acc \
               - lambda_L * (latency / cfg["latency_budget"]) \
               - lambda_E * (energy / battery.energy_budget)

        loss = ce_loss + cfg["rl_weight"] * (-logprob * reward) - 0.001 * logprob

        scaler.scale(loss).backward()
        scaler.step(optimizer)
        scaler.update()

        total_loss += loss.item()

    avg_latency = sum_latency / len(loader)
    avg_energy = sum_energy / len(loader)

    dual_opt.update_duals(
        avg_latency=avg_latency,
        latency_budget=cfg["latency_budget"],
        avg_energy=avg_energy,
        energy_budget=battery.energy_budget
    )

    epoch_loss = total_loss / len(loader)
    epoch_acc = total_correct / total_samples

    action_dist = [c/sum(action_counts) for c in action_counts]

    return epoch_loss, epoch_acc, avg_latency, avg_energy, action_dist


# ------------------------------------------------------------
# MAIN
# ------------------------------------------------------------

def main():

    parser = argparse.ArgumentParser()

    parser.add_argument("--dataset", required=True)
    parser.add_argument("--model", required=True)
    parser.add_argument("--split", type=int, required=True)
    parser.add_argument("--optimizer", required=True)
    parser.add_argument("--energy_budget", type=float, required=True)

    parser.add_argument("--epochs", type=int, default=100)
    parser.add_argument("--batch_size", type=int, default=4)
    parser.add_argument("--runs_root", default="./runs")

    args = parser.parse_args()

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    with open("configs/default.yaml") as f:
        cfg = yaml.safe_load(f)

    run_root, metrics_dir = make_run_dirs(
        args.runs_root,
        args.dataset,
        args.model,
        args.split,
        args.energy_budget,
        args.optimizer
    )

    epoch_csv = os.path.join(metrics_dir,"epoch_metrics.csv")

    if not os.path.exists(epoch_csv):

        with open(epoch_csv,"w",newline="") as f:

            writer = csv.writer(f)

            writer.writerow([
                "epoch",
                "train_loss",
                "train_acc",
                "val_acc",
                "latency",
                "energy",
                "lambda_L",
                "lambda_E",
                "action_full",
                "action_half",
                "action_quarter"
            ])

    # DATASET

    if args.dataset.lower() == "hmdb51":

        root = os.path.expanduser("~/hmdb51")

        train_set = HMDB51Dataset(root,args.split,True,16,112)
        val_set = HMDB51Dataset(root,args.split,False,16,112)

        num_classes = 51

    else:

        root = os.path.expanduser("~/ucf101")

        train_set = UCF101Dataset(root,args.split,True,16,112)
        val_set = UCF101Dataset(root,args.split,False,16,112)

        num_classes = 101

    train_loader = DataLoader(train_set,args.batch_size,shuffle=True,num_workers=8)
    val_loader = DataLoader(val_set,args.batch_size,shuffle=False,num_workers=8)

    # MODEL

    if args.model=="r3d18":
        model = r3d_18(num_classes)

    elif args.model=="r2plus1d18":
        model = r2plus1d_18(num_classes)

    else:
        model = mc3_18(num_classes)

    model = model.to(device)

    # OPTIMIZER

    if args.optimizer=="sgd":

        optimizer = torch.optim.SGD(model.parameters(),lr=0.01,momentum=0.9)

    elif args.optimizer=="adam":

        optimizer = torch.optim.Adam(model.parameters(),lr=5e-4)

    else:

        optimizer = torch.optim.AdamW(model.parameters(),lr=5e-4)

    scaler = torch.amp.GradScaler("cuda")

    # CONTROLLER

    controller = TemporalController(
        cfg["controller_state_dim"],
        cfg["controller_hidden_dim"],
        cfg["action_dim"],
        cfg["controller_lr"]
    ).to(device)

    # COST MODELS

    lut = load_kernel_lut(cfg["lut_path"])

    battery = VirtualBattery(args.energy_budget,len(train_loader))

    dual_opt = PrimalDualOptimizer(lr_lambda=0.02)

    # TRAIN LOOP

    for epoch in range(1,args.epochs+1):

        battery.reset()

        train_loss,train_acc,avg_lat,avg_eng,action_dist = train_one_epoch(
            model,controller,optimizer,scaler,
            train_loader,lut,battery,dual_opt,
            device,cfg
        )

        val_acc = evaluate(model,val_loader,device)

        lambda_L,lambda_E = dual_opt.get_lambdas()

        print(f"\nEpoch {epoch}")
        print(f"Train Loss: {train_loss:.4f}")
        print(f"Train Acc:  {train_acc:.4f}")
        print(f"Val Acc:    {val_acc:.4f}")
        print(f"Latency:    {avg_lat:.4f}")
        print(f"Energy:     {avg_eng:.6f}")
        print(f"Lambda_L:   {lambda_L:.4f}")
        print(f"Lambda_E:   {lambda_E:.4f}")
        print(f"Actions: full={action_dist[0]:.2f} half={action_dist[1]:.2f} quarter={action_dist[2]:.2f}")

        with open(epoch_csv,"a",newline="") as f:

            writer = csv.writer(f)

            writer.writerow([
                epoch,
                train_loss,
                train_acc,
                val_acc,
                avg_lat,
                avg_eng,
                lambda_L,
                lambda_E,
                action_dist[0],
                action_dist[1],
                action_dist[2]
            ])


if __name__ == "__main__":
    main()
