import os
import glob
import random
import numpy as np
import pandas as pd
from pathlib import Path
from PIL import Image

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
from torch.optim import AdamW
from torch.optim.lr_scheduler import ReduceLROnPlateau
from torch.cuda.amp import GradScaler, autocast

import timm
import albumentations as A
from albumentations.pytorch import ToTensorV2
from sklearn.model_selection import StratifiedKFold
from sklearn.metrics import accuracy_score

# ============================================================
# 1. Config
# ============================================================
class CFG:
    # Essential Paths
    DATA_DIR = Path("/kaggle/input/competitions/cassava-leaf-disease-classification")
    TRAIN_IMG = DATA_DIR / "train_images"
    TRAIN_CSV = DATA_DIR / "train.csv"

    # Pretrained Weights Path
    MODEL_DIR = "/kaggle/input/models/ashishkubade/timm-tf-efficientnetv2-s-in21k/pytorch/default/1"

    # Model Parameters
    MODEL_NAME = "tf_efficientnetv2_s_in21k"
    IMG_SIZE = 384
    EPOCHS = 15
    LR = 1e-4
    SMOOTHING = 0.1

    NUM_CLASSES = 5
    SEED = 42
    FOLD = 0
    N_SPLITS = 5
    WEIGHT_DECAY = 1e-6
    PATIENCE = 3
    BATCH_SIZE = 16
    NUM_WORKERS = 4
    DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# ============================================================
# 2. Reproducibility
# ============================================================
def seed_everything(seed=42):
    random.seed(seed)
    np.random.seed(seed)
    os.environ["PYTHONHASHSEED"] = str(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    torch.backends.cudnn.deterministic = False
    torch.backends.cudnn.benchmark = True

seed_everything(CFG.SEED)

# ============================================================
# 3. Augmentation
# ============================================================
MEAN = [0.485, 0.456, 0.406]
STD = [0.229, 0.224, 0.225]

def get_transforms(mode="train"):
    if mode == "train":
        return A.Compose([
            A.RandomResizedCrop(size=(CFG.IMG_SIZE, CFG.IMG_SIZE), scale=(0.7, 1.0), p=1.0),
            A.HorizontalFlip(p=0.5),
            A.VerticalFlip(p=0.5),
            A.Affine(
                translate_percent={"x": (-0.1, 0.1), "y": (-0.1, 0.1)},
                scale=(0.8, 1.2),
                rotate=(-30, 30),
                p=0.5,
            ),
            A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1, p=0.5),
            A.GaussianBlur(blur_limit=(3, 7), p=0.2),
            A.Normalize(mean=MEAN, std=STD),
            ToTensorV2(),
        ])
    return A.Compose([
        A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),
        A.Normalize(mean=MEAN, std=STD),
        ToTensorV2(),
    ])

# ============================================================
# 4. Dataset
# ============================================================
class CassavaDataset(Dataset):
    def __init__(self, df, img_dir, transform=None, is_test=False):
        self.df = df.reset_index(drop=True)
        self.img_dir = img_dir
        self.transform = transform
        self.is_test = is_test

    def __len__(self):
        return len(self.df)

    def __getitem__(self, idx):
        row = self.df.iloc[idx]
        image_path = self.img_dir / row["image_id"]
        image = np.array(Image.open(image_path).convert("RGB"))
        if self.transform:
            image = self.transform(image=image)["image"]
        if self.is_test:
            return image
        label = torch.tensor(row["label"], dtype=torch.long)
        return image, label

# ============================================================
# 5. Loss
# ============================================================
class LabelSmoothingCrossEntropy(nn.Module):
    def __init__(self, smoothing=0.1):
        super(LabelSmoothingCrossEntropy, self).__init__()
        self.smoothing = smoothing

    def forward(self, pred, target):
        n = pred.size(-1)
        log_prob = F.log_softmax(pred, dim=-1)
        with torch.no_grad():
            smooth = torch.full_like(log_prob, self.smoothing / (n - 1))
            smooth.scatter_(1, target.unsqueeze(1), 1.0 - self.smoothing)
        return -(smooth * log_prob).sum(-1).mean()

# ============================================================
# 6. Model
# ============================================================
def find_weight_file(model_dir):
    if os.path.isfile(model_dir):
        return model_dir
    patterns = ["**/*.safetensors", "**/*.pth", "**/*.bin", "**/*.pt"]
    for pat in patterns:
        files = glob.glob(os.path.join(model_dir, pat), recursive=True)
        if files:
            files.sort(key=lambda x: os.path.getsize(x), reverse=True)
            return files[0]
    return None

def build_model(load_weights=True):
    weight_file = find_weight_file(CFG.MODEL_DIR) if load_weights else None
    
    # Pretrained=False to avoid internet requests during competition
    backbone = timm.create_model(CFG.MODEL_NAME, pretrained=False, num_classes=0)
    
    if weight_file:
        print("Loading local weights: {}".format(os.path.basename(weight_file)))
        try:
            state_dict = torch.load(weight_file, map_location="cpu")
            if "model" in state_dict:
                state_dict = state_dict["model"]
            elif "state_dict" in state_dict:
                state_dict = state_dict["state_dict"]
            
            # Clean prefixes (model., backbone., module.)
            new_state_dict = {}
            for k, v in state_dict.items():
                name = k
                if k.startswith("model."):
                    name = k[6:]
                elif k.startswith("backbone."):
                    name = k[9:]
                elif k.startswith("module."):
                    name = k[7:]
                new_state_dict[name] = v
            
            backbone.load_state_dict(new_state_dict, strict=False)
            print("Model weights initialized successfully.")
        except Exception as e:
            print("Warning: Weight loading issue: {}".format(e))
    else:
        if load_weights:
            print("No weight file found. Starting from random initialization.")

    model = nn.Sequential(
        backbone,
        nn.Sequential(
            nn.Dropout(p=0.3),
            nn.Linear(backbone.num_features, CFG.NUM_CLASSES),
        ),
    )
    return model

# ============================================================
# 7. EarlyStopping
# ============================================================
class EarlyStopping:
    def __init__(self, patience=3, path="best_model.pth"):
        self.patience = patience
        self.path = path
        self.counter = 0
        self.best_score = None
        self.early_stop = False

    def __call__(self, val_loss, model):
        score = -val_loss
        if self.best_score is None or score > self.best_score:
            self.best_score = score
            torch.save(model.state_dict(), self.path)
            print("  [SAVE] New best val_loss={:.4f}".format(val_loss))
            self.counter = 0
        else:
            self.counter += 1
            print("  [EarlyStop] {}/{}".format(self.counter, self.patience))
            if self.counter >= self.patience:
                self.early_stop = True

# ============================================================
# 8. Training Functions
# ============================================================
def train_one_epoch(model, loader, criterion, optimizer, scaler):
    model.train()
    total_loss = 0.0
    all_preds = []
    all_labels = []
    for images, labels in loader:
        images = images.to(CFG.DEVICE)
        labels = labels.to(CFG.DEVICE)
        optimizer.zero_grad()
        with autocast(enabled=True):
            outputs = model(images)
            loss = criterion(outputs, labels)
        scaler.scale(loss).backward()
        scaler.step(optimizer)
        scaler.update()
        total_loss += loss.item() * images.size(0)
        all_preds.extend(outputs.argmax(1).detach().cpu().numpy())
        all_labels.extend(labels.cpu().numpy())
    return (total_loss / len(loader.dataset)), accuracy_score(all_labels, all_preds)

@torch.no_grad()
def valid_one_epoch(model, loader, criterion):
    model.eval()
    total_loss = 0.0
    all_preds = []
    all_labels = []
    for images, labels in loader:
        images = images.to(CFG.DEVICE)
        labels = labels.to(CFG.DEVICE)
        outputs = model(images)
        loss = criterion(outputs, labels)
        total_loss += loss.item() * images.size(0)
        all_preds.extend(outputs.argmax(1).cpu().numpy())
        all_labels.extend(labels.cpu().numpy())
    return (total_loss / len(loader.dataset)), accuracy_score(all_labels, all_preds)

# ============================================================
# 9. Pipeline Runner
# ============================================================
def run_pipeline():
    # 1. Load Data
    df = pd.read_csv(CFG.TRAIN_CSV)
    print("Device: {} | Model: {} | Resolution: {}px".format(CFG.DEVICE, CFG.MODEL_NAME, CFG.IMG_SIZE))
    
    skf = StratifiedKFold(n_splits=CFG.N_SPLITS, shuffle=True, random_state=CFG.SEED)
    folds = list(skf.split(df, df["label"]))
    train_idx, val_idx = folds[CFG.FOLD]
    train_df = df.iloc[train_idx]
    val_df = df.iloc[val_idx]
    
    # 2. DataLoaders (Batch Size 16 is safer for 384px)
    train_loader = DataLoader(
        CassavaDataset(train_df, CFG.TRAIN_IMG, get_transforms("train")),
        batch_size=CFG.BATCH_SIZE,
        shuffle=True,
        num_workers=CFG.NUM_WORKERS,
        pin_memory=True,
        drop_last=True
    )
    val_loader = DataLoader(
        CassavaDataset(val_df, CFG.TRAIN_IMG, get_transforms("valid")),
        batch_size=CFG.BATCH_SIZE * 2,
        shuffle=False,
        num_workers=CFG.NUM_WORKERS,
        pin_memory=True
    )

    # 3. Build Model & Tools
    model = build_model(load_weights=True).to(CFG.DEVICE)
    criterion = LabelSmoothingCrossEntropy(smoothing=CFG.SMOOTHING)
    optimizer = AdamW(model.parameters(), lr=CFG.LR, weight_decay=CFG.WEIGHT_DECAY)
    scheduler = ReduceLROnPlateau(optimizer, mode="min", factor=0.5, patience=2)
    scaler = GradScaler()
    checkpoint_path = "best_model_fold{}.pth".format(CFG.FOLD)
    es = EarlyStopping(patience=CFG.PATIENCE, path=checkpoint_path)

    # 4. Training Loop
    best_val_acc = 0.0
    for epoch in range(1, CFG.EPOCHS + 1):
        curr_lr = optimizer.param_groups[0]["lr"]
        print("\n[Epoch {}/{}] LR: {:.2e}".format(epoch, CFG.EPOCHS, curr_lr))
        
        tr_loss, tr_acc = train_one_epoch(model, train_loader, criterion, optimizer, scaler)
        vl_loss, vl_acc = valid_one_epoch(model, val_loader, criterion)
        
        is_best = ""
        if vl_acc > best_val_acc:
            best_val_acc = vl_acc
            is_best = " << BEST"
            
        print("  Train Loss: {:.4f} | Acc: {:.4f}".format(tr_loss, tr_acc))
        print("  Valid Loss: {:.4f} | Acc: {:.4f}{}".format(vl_loss, vl_acc, is_best))
        
        scheduler.step(vl_loss)
        es(vl_loss, model)
        if es.early_stop:
            print("[Info] Early stopping triggered.")
            break

    print("\nTraining complete. Highest Validation Accuracy: {:.4f}".format(best_val_acc))

    # 5. Inference & Submission
    print("\nRunning test inference...")
    inference_model = build_model(load_weights=False).to(CFG.DEVICE)
    inference_model.load_state_dict(torch.load(checkpoint_path, map_location=CFG.DEVICE))
    
    test_dir = CFG.DATA_DIR / "test_images"
    test_files = sorted(list(test_dir.glob("*.jpg")))
    test_df = pd.DataFrame({"image_id": [f.name for f in test_files]})
    
    test_loader = DataLoader(
        CassavaDataset(test_df, test_dir, get_transforms("valid"), is_test=True),
        batch_size=CFG.BATCH_SIZE * 2,
        shuffle=False,
        num_workers=CFG.NUM_WORKERS
    )
    
    inference_model.eval()
    preds = []
    with torch.no_grad():
        for imgs in test_loader:
            imgs = imgs.to(CFG.DEVICE)
            with autocast(enabled=True):
                batch_preds = inference_model(imgs).argmax(1).cpu().numpy()
                preds.extend(batch_preds)
                
    submission = pd.DataFrame({"image_id": test_df["image_id"], "label": preds})
    submission.to_csv("submission.csv", index=False)
    print("Success! submission.csv created.")

# ============================================================
# Execution
# ============================================================
run_pipeline()
