#================================================
# Practical Image Processing and Natural Language Processing
# Dr. Hager Saleh
# Rami AbuLaban - 233000631
#================================================

import os
import zipfile
import torch
import torchvision.transforms as transforms
import torchvision.datasets as datasets
import matplotlib.pyplot as plt
from torch.utils.data import DataLoader, random_split

# 🔹 Step 1: Define Dataset Path
dataset_folder = r"C:\Users\ramia\Desktop\eye_diseases_classification\Project1"  # ✅ Update with your folder path
dataset_zip = os.path.join(dataset_folder, "eye_diseases.zip")  # If dataset is a ZIP file
extract_path = os.path.join(dataset_folder, "eye_diseases_dataset")  # Folder where it will be extracted

# 🔹 Step 2: Extract Dataset (Only If Needed)
if os.path.exists(dataset_zip) and not os.path.exists(extract_path):
    print("📂 Extracting dataset...")
    with zipfile.ZipFile(dataset_zip, "r") as zip_ref:
        zip_ref.extractall(extract_path)
    print("✅ Dataset extracted successfully!")

# 🔹 Step 3: Set Dataset Path
data_dir = extract_path if os.path.exists(extract_path) else dataset_folder  # ✅ Use correct dataset path

# 🔹 Step 4: Define Data Transformations
data_transform = transforms.Compose([
    transforms.Resize((224, 224)),  
    transforms.ToTensor(),  
    transforms.Normalize([0.5], [0.5])  
])

# 🔹 Step 5: Load Dataset
full_dataset = datasets.ImageFolder(data_dir, transform=data_transform)

# 🔹 Step 6: Print Class Names
print(f"📌 Classes found: {full_dataset.classes}")  # ✅ Displays ["cataract", "diabetic_retinopathy", "glaucoma", "normal"]

# 🔹 Step 7: Split into Train and Test Sets
train_size = int(0.8 * len(full_dataset))
test_size = len(full_dataset) - train_size
train_dataset, test_dataset = random_split(full_dataset, [train_size, test_size])

# 🔹 Step 8: Define Data Loaders
batch_size = 32
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)

# 🔹 Step 9: Show Sample Images with Correct Class Names
def show_images(dataset, num_images=5):
    fig, axes = plt.subplots(1, num_images, figsize=(15,5))
    for i in range(num_images):
        img, label = dataset[i]
        img = img.permute(1, 2, 0)  # ✅ Convert to (H, W, C) format
        
        class_name = full_dataset.classes[label]  # ✅ Get correct class name
        axes[i].imshow(img)
        axes[i].set_title(f"Class: {class_name}", fontsize=12)  # ✅ Show real class name
        axes[i].axis("off")
    
    plt.show()

# 🔹 Step 10: Display Sample Images
print("\n📸 Displaying sample images from dataset...")
show_images(full_dataset)
