import numpy as np
import tensorflow as tf
from tensorflow.keras.applications import ResNet50
from tensorflow.keras import layers, models
from tensorflow.keras.preprocessing.image import ImageDataGenerator
from sklearn.metrics import classification_report, confusion_matrix
import seaborn as sns
import matplotlib.pyplot as plt

# ------------------ Data Augmentation ------------------

train_datagen = ImageDataGenerator(
    rescale=1./255,
    rotation_range=20,
    zoom_range=0.2,
    width_shift_range=0.1,
    height_shift_range=0.1,
    shear_range=0.1,
    horizontal_flip=True,
    validation_split=0.1
)

val_datagen = ImageDataGenerator(rescale=1./255, validation_split=0.1)

# Load training data
train_data = train_datagen.flow_from_directory(
    'dataset/train',
    target_size=(224,224),
    batch_size=32,
    class_mode='categorical',
    subset='training'
)

# Load validation data
val_data = val_datagen.flow_from_directory(
    'dataset/train',
    target_size=(224,224),
    batch_size=32,
    class_mode='categorical',
    subset='validation'
)

# ------------------ Load ResNet50 ------------------

base_model = ResNet50(weights='imagenet',
                      include_top=False,
                      input_shape=(224,224,3))

# Freeze all layers
for layer in base_model.layers:
    layer.trainable = False

# ------------------ Custom Model ------------------

x = base_model.output
x = layers.GlobalAveragePooling2D()(x)
x = layers.BatchNormalization()(x)
x = layers.Dense(256, activation='relu')(x)
x = layers.Dropout(0.5)(x)
output = layers.Dense(train_data.num_classes, activation='softmax')(x)

model = models.Model(inputs=base_model.input, outputs=output)

# Compile
model.compile(optimizer='adam',
              loss='categorical_crossentropy',
              metrics=['accuracy'])

# Train (Feature Extraction)
model.fit(train_data, epochs=5, validation_data=val_data)

# ------------------ Fine Tuning ------------------

for layer in base_model.layers:
    if (layer.name.startswith('conv4') or
        layer.name.startswith('conv5')):
        layer.trainable = True

# Recompile
model.compile(optimizer=tf.keras.optimizers.Adam(1e-5),
              loss='categorical_crossentropy',
              metrics=['accuracy'])

# Train again
model.fit(train_data, epochs=3, validation_data=val_data)

# ------------------ Evaluation ------------------

val_data.reset()
y_pred_probs = model.predict(val_data)
y_pred = np.argmax(y_pred_probs, axis=1)
y_true = val_data.classes

# Accuracy
loss, acc = model.evaluate(val_data)
print("Validation Accuracy:", acc)

# Precision, Recall, F1-score
print("\nClassification Report:\n")
print(classification_report(y_true, y_pred))

# Confusion Matrix
cm = confusion_matrix(y_true, y_pred)

plt.figure(figsize=(8,6))
sns.heatmap(cm, annot=True, fmt='d')
plt.xlabel("Predicted")
plt.ylabel("Actual")
plt.title("Confusion Matrix")
plt.show()

