import wandb
import pandas as pd
import yaml
from pathlib import Path
from sklearn.model_selection import train_test_split
import os
from shutil import copyfile
from ultralytics import YOLO
import torch
import cv2
import numpy as np
from tqdm import tqdm
import multiprocessing as mp
import logging

# Thiết lập logging
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')

# Thiết lập API key cho wandb
wandb.login(key='4f88ff0bbc6e3485258bca7079d7da4c47798ccd')

# Khởi tạo một run mới với tên dự án và tên lượt chạy cụ thể
wandb.init(project='xuan2261_yolov8', name='train_aortic_specialist')

# Paths to data directories
ROOT = Path("/kaggle/input/amia-public-challenge-2024")
YOLO_DIR = Path('/kaggle/working/yolov8_data/')

def read_and_preprocess_data():
    logging.info("Reading and preprocessing data...")
    train_df = pd.read_csv(ROOT / "train.csv")
    img_size_df = pd.read_csv(ROOT / "img_size.csv")
    train_df = train_df.merge(img_size_df, on='image_id', how='left')
    return train_df

def split_data(df):
    logging.info("Splitting data into train, validation, and test sets...")
    train_data, temp_data = train_test_split(df, test_size=0.3, random_state=42, stratify=df['class_id'])
    valid_data, test_data = train_test_split(temp_data, test_size=2/3, random_state=42, stratify=temp_data['class_id'])
    return train_data, valid_data, test_data

def convert_to_yolo_format(row, data_type):
    try:
        img_path = ROOT / f"train/train/{row['image_id']}.png"
        img_dest = YOLO_DIR / f"{data_type}/images/{row['image_id']}.png"
        copyfile(img_path, img_dest)
        
        label_path = YOLO_DIR / f"{data_type}/labels/{row['image_id']}.txt"
        
        orig_width, orig_height = row['dim1'], row['dim0']
        
        if row['class_id'] == 14:  # No finding
            x_center, y_center, width, height = 0.5, 0.5, 1.0, 1.0
        else:
            x_center = (row['x_min'] + row['x_max']) / 2 / orig_width
            y_center = (row['y_min'] + row['y_max']) / 2 / orig_height
            width = (row['x_max'] - row['x_min']) / orig_width
            height = (row['y_max'] - row['y_min']) / orig_height

        if not (0 <= x_center <= 1 and 0 <= y_center <= 1 and 0 <= width <= 1 and 0 <= height <= 1):
            logging.warning(f"Skipping image {row['image_id']} due to out of bounds coordinates.")
            return

        label_content = f"{row['class_id']} {x_center} {y_center} {width} {height}\n"

        with open(label_path, 'w') as f:
            f.write(label_content)
    except Exception as e:
        logging.error(f"Error processing image {row['image_id']}: {e}")

def prepare_yolo_data(train_data, valid_data, test_data):
    logging.info("Preparing YOLO data...")
    YOLO_DIR.mkdir(parents=True, exist_ok=True)
    for folder in ['train/images', 'train/labels', 'val/images', 'val/labels', 'test/images', 'test/labels']:
        (YOLO_DIR / folder).mkdir(parents=True, exist_ok=True)

    with mp.Pool(mp.cpu_count()) as pool:
        for dataset, data_type in [(train_data, 'train'), (valid_data, 'val'), (test_data, 'test')]:
            pool.starmap(convert_to_yolo_format, [(row, data_type) for _, row in dataset.iterrows()])

    class_names = ['Aortic enlargement', 'Atelectasis', 'Calcification', 'Cardiomegaly', 'Consolidation',
                   'ILD', 'Infiltration', 'Lung Opacity', 'Nodule/Mass', 'Other lesion', 'Pleural effusion',
                   'Pleural thickening', 'Pneumothorax', 'Pulmonary fibrosis', 'No finding']

    data_yaml = dict(
        train=str(YOLO_DIR / 'train/images'),
        val=str(YOLO_DIR / 'val/images'),
        nc=15,
        names=class_names
    )

    with open(YOLO_DIR / 'data.yaml', 'w') as outfile:
        yaml.dump(data_yaml, outfile, default_flow_style=False)

def visualize_predictions(image_path, results):
    img = cv2.imread(image_path)
    for result in results:
        boxes = result.boxes.xyxy.cpu().numpy()
        for box in boxes:
            cv2.rectangle(img, (int(box[0]), int(box[1])), (int(box[2]), int(box[3])), (0, 255, 0), 2)
    return img

def post_process(results, iou_threshold=0.5, conf_threshold=0.25):
    processed_results = []
    for result in results:
        # Apply NMS
        keep = torchvision.ops.nms(result.boxes.xyxy, result.boxes.conf, iou_threshold)
        processed_result = result[keep]
        # Filter by confidence threshold
        processed_result = processed_result[processed_result.boxes.conf > conf_threshold]
        processed_results.append(processed_result)
    return processed_results

def main():
    # Data preparation
    train_df = read_and_preprocess_data()
    train_data, valid_data, test_data = split_data(train_df)
    prepare_yolo_data(train_data, valid_data, test_data)

    # Load YOLOv5 weights and convert to YOLOv8
    yolov5_weights = '/kaggle/input/best-aortic-fold4-pt'
    model = YOLO('yolov8m.pt')  # Start with a pre-trained YOLOv8 model
    model.model.load(yolov5_weights)  # Load YOLOv5 weights
    model.save('yolov8_aortic_specialist.pt')  # Save as YOLOv8 weights

    # Model fine-tuning
    results = model.train(
        data='/kaggle/working/yolov8_data/data.yaml', 
        epochs=20,
        batch=32, 
        imgsz=640,
        lr0=0.001,
        augment=True
    )

    # Logging training metrics
    for epoch, result in enumerate(results):
        wandb.log({
            "epoch": epoch,
            "train/box_loss": result.box_loss,
            "train/cls_loss": result.cls_loss,
            "train/dfl_loss": result.dfl_loss
        })

    # Model evaluation
    metrics = model.val()
    logging.info(f"Validation metrics: {metrics}")

    # Prediction
    results = model.predict(source='/kaggle/working/yolov8_data/test/images', save=False)
    
    # Post-processing
    processed_results = post_process(results)

    # Visualization
    for i, result in enumerate(tqdm(processed_results, desc="Visualizing predictions")):
        img_path = f'/kaggle/working/yolov8_data/test/images/{i}.png'
        vis_img = visualize_predictions(img_path, [result])
        cv2.imwrite(f'/kaggle/working/predictions/{i}_pred.png', vis_img)

    logging.info("Process completed successfully.")

if __name__ == "__main__":
    main()