"""
LEAP - Atmospheric Physics using AI (ClimSim)
Row-Aligned Atmospheric ResNet Emulator & Submission Generator
Author: Matheus Bonjour (Oceanographer & Meteorologist)
"""

import os
import sys
import gc
import glob
import time
import pandas as pd
import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader

# ==============================================================================
# Hardware Setup
# ==============================================================================
DEVICE = torch.device("cpu")
if torch.cuda.is_available():
    try:
        t = torch.zeros(1).cuda()
        DEVICE = torch.device("cuda")
        print(f"=== GPU Initialization Successful: {torch.cuda.get_device_name(0)} ===")
    except Exception as e:
        print(f"=== CUDA GPU incompatible with PyTorch build ({e}). Falling back to CPU. ===")

print(f"Using compute device: {DEVICE}")

# ==============================================================================
# Dynamic File Path Discovery
# ==============================================================================
def find_file(filename):
    search_paths = ["/kaggle/input", "."]
    for sp in search_paths:
        matches = glob.glob(os.path.join(sp, "**", filename), recursive=True)
        if matches:
            print(f"Found {filename} at: {matches[0]}")
            return matches[0]
    raise FileNotFoundError(f"Could not locate {filename} in /kaggle/input or current directory.")

SAMPLE_SUB_PATH = find_file("sample_submission.csv")
TRAIN_PATH = find_file("train.csv")
TEST_PATH = find_file("test.csv")

BATCH_SIZE = 2048
EPOCHS = 10
LR = 3e-3
TRAIN_ROWS = 600000  # High performance training subset

# ==============================================================================
# 1. Load Sample Submission Header & Mask
# ==============================================================================
print("\n--- 1. Loading Sample Submission Header & Target Mask ---")
sub_sample = pd.read_csv(SAMPLE_SUB_PATH, nrows=1)
target_cols = [c for c in sub_sample.columns if c != "sample_id"]
weight_mask = sub_sample.iloc[0][target_cols].values.astype(np.float32)
active_indices = np.where(weight_mask > 0)[0]

print(f"Total Target Columns: {len(target_cols)}")
print(f"Active Target Columns: {len(active_indices)}")
print(f"Masked (Zeroed) Columns: {len(target_cols) - len(active_indices)}")

# ==============================================================================
# 2. Dataset Preprocessing & Standardization
# ==============================================================================
print(f"\n--- 2. Loading Training Data ({TRAIN_ROWS} rows) ---")
train_sample = pd.read_csv(TRAIN_PATH, nrows=2)
feature_cols = [c for c in train_sample.columns if c not in target_cols and c != "sample_id"]
print(f"Feature Columns: {len(feature_cols)}")

train_df = pd.read_csv(TRAIN_PATH, nrows=TRAIN_ROWS)
X_train_raw = train_df[feature_cols].values.astype(np.float32)
y_train_raw = train_df[target_cols].values.astype(np.float32)
del train_df
gc.collect()

# Standardize Inputs
X_mean = np.nanmean(X_train_raw, axis=0, keepdims=True)
X_std = np.nanstd(X_train_raw, axis=0, keepdims=True) + 1e-6
X_train = (X_train_raw - X_mean) / X_std

# Standardize Targets
y_mean = np.nanmean(y_train_raw, axis=0, keepdims=True)
y_std = np.nanstd(y_train_raw, axis=0, keepdims=True) + 1e-6

y_train = (y_train_raw - y_mean) / y_std
y_train[:, weight_mask == 0] = 0.0

del X_train_raw, y_train_raw
gc.collect()

class ClimDataset(Dataset):
    def __init__(self, X, y):
        self.X = torch.tensor(X, dtype=torch.float32)
        self.y = torch.tensor(y, dtype=torch.float32)
    def __len__(self):
        return len(self.X)
    def __getitem__(self, idx):
        return self.X[idx], self.y[idx]

train_dataset = ClimDataset(X_train, y_train)
train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, drop_last=True)

# ==============================================================================
# 3. Model Architecture: Atmospheric ResNet
# ==============================================================================
class ResBlock(nn.Module):
    def __init__(self, hidden_dim, dropout=0.1):
        super().__init__()
        self.block = nn.Sequential(
            nn.Linear(hidden_dim, hidden_dim),
            nn.LayerNorm(hidden_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(hidden_dim, hidden_dim),
            nn.LayerNorm(hidden_dim)
        )
        self.act = nn.GELU()

    def forward(self, x):
        return self.act(x + self.block(x))

class AtmosphericResNet(nn.Module):
    def __init__(self, in_features, out_targets, hidden_dim=512, num_blocks=2):
        super().__init__()
        self.in_proj = nn.Sequential(
            nn.Linear(in_features, hidden_dim),
            nn.LayerNorm(hidden_dim),
            nn.GELU()
        )
        self.blocks = nn.ModuleList([ResBlock(hidden_dim) for _ in range(num_blocks)])
        self.out_head = nn.Sequential(
            nn.Linear(hidden_dim, hidden_dim // 2),
            nn.LayerNorm(hidden_dim // 2),
            nn.GELU(),
            nn.Linear(hidden_dim // 2, out_targets)
        )

    def forward(self, x):
        h = self.in_proj(x)
        for block in self.blocks:
            h = block(h)
        return self.out_head(h)

model = AtmosphericResNet(in_features=len(feature_cols), out_targets=len(target_cols)).to(DEVICE)
print("\n--- Model Architecture ---")
print(model)

# ==============================================================================
# 4. Training Loop
# ==============================================================================
criterion = nn.MSELoss()
optimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)
weight_mask_tensor = torch.tensor(weight_mask, dtype=torch.float32, device=DEVICE)

print(f"\n--- 4. Training Model for {EPOCHS} Epochs ---")
start_time = time.time()

for epoch in range(1, EPOCHS + 1):
    model.train()
    total_loss = 0.0
    for bx, by in train_loader:
        bx, by = bx.to(DEVICE), by.to(DEVICE)
        optimizer.zero_grad()

        preds = model(bx)
        masked_preds = preds * weight_mask_tensor
        masked_targets = by * weight_mask_tensor
        loss = criterion(masked_preds, masked_targets)

        loss.backward()
        optimizer.step()

        total_loss += loss.item() * len(bx)

    scheduler.step()
    avg_loss = total_loss / len(train_dataset)
    print(f"Epoch {epoch:02d}/{EPOCHS:02d} | Train MSE Loss: {avg_loss:.6f} | LR: {scheduler.get_last_lr()[0]:.6f}")

elapsed = time.time() - start_time
print(f"Training completed in {elapsed/60:.2f} minutes.")

del X_train, y_train, train_dataset, train_loader
gc.collect()

# ==============================================================================
# 5. Test Set Inference & Row-Order Alignment
# ==============================================================================
print("\n--- 5. Generating Test Predictions ---")
model.eval()

sub_rows = []
chunk_size = 50000

for chunk in pd.read_csv(TEST_PATH, chunksize=chunk_size):
    sample_ids = chunk["sample_id"].values
    X_test_raw = chunk[feature_cols].values.astype(np.float32)
    X_test = (X_test_raw - X_mean) / X_std

    X_test_tensor = torch.tensor(X_test, dtype=torch.float32, device=DEVICE)

    with torch.no_grad():
        test_preds_norm = model(X_test_tensor).cpu().numpy()

    test_preds_physical = test_preds_norm * y_std + y_mean
    test_preds_final = test_preds_physical * weight_mask

    chunk_sub = pd.DataFrame(test_preds_final, columns=target_cols)
    chunk_sub.insert(0, "sample_id", sample_ids)
    sub_rows.append(chunk_sub)

print("Concatenating raw prediction chunks...")
raw_preds_df = pd.concat(sub_rows, ignore_index=True)

print("Aligning prediction rows to sample_submission.csv exact alphabetical order...")
sub_order_df = pd.read_csv(SAMPLE_SUB_PATH, usecols=["sample_id"])
sub_df = sub_order_df.merge(raw_preds_df, on="sample_id", how="left")

submission_file = "submission.csv"
sub_df.to_csv(submission_file, index=False)
print(f"Row-aligned submission saved successfully to {submission_file}! Shape: {sub_df.shape}")
