from tensorflow.keras.preprocessing.image import ImageDataGenerator
from tensorflow.keras.applications import MobileNetV2
from tensorflow.keras import layers, models

img_size = (224, 224)
batch_size = 16
epochs = 10

datagen = ImageDataGenerator(
    rescale=1./255,
    validation_split=0.2
)

def train_binary_model(folder, model_name):
    train_data = datagen.flow_from_directory(
        folder,
        target_size=img_size,
        batch_size=batch_size,
        class_mode="binary",
        subset="training"
    )

    val_data = datagen.flow_from_directory(
        folder,
        target_size=img_size,
        batch_size=batch_size,
        class_mode="binary",
        subset="validation"
    )

    base = MobileNetV2(weights="imagenet", include_top=False, input_shape=(224,224,3))
    base.trainable = False

    model = models.Sequential([
        base,
        layers.GlobalAveragePooling2D(),
        layers.Dense(128, activation="relu"),
        layers.Dense(1, activation="sigmoid")
    ])

    model.compile(
        optimizer="adam",
        loss="binary_crossentropy",
        metrics=["accuracy"]
    )

    model.fit(train_data, validation_data=val_data, epochs=epochs)
    model.save(model_name)

def train_skin_model():
    train_data = datagen.flow_from_directory(
        "dataset/skin_type",
        target_size=img_size,
        batch_size=batch_size,
        class_mode="categorical",
        subset="training"
    )

    val_data = datagen.flow_from_directory(
        "dataset/skin_type",
        target_size=img_size,
        batch_size=batch_size,
        class_mode="categorical",
        subset="validation"
    )

    base = MobileNetV2(weights="imagenet", include_top=False, input_shape=(224,224,3))
    base.trainable = False

    model = models.Sequential([
        base,
        layers.GlobalAveragePooling2D(),
        layers.Dense(128, activation="relu"),
        layers.Dense(3, activation="softmax")
    ])

    model.compile(
        optimizer="adam",
        loss="categorical_crossentropy",
        metrics=["accuracy"]
    )

    model.fit(train_data, validation_data=val_data, epochs=epochs)
    model.save("skin_type_model.h5")

print("Training pimples model...")
train_binary_model("dataset/pimples", "pimples_model.h5")

print("Training dark circles model...")
train_binary_model("dataset/dark_circles", "dark_circles_model.h5")

print("Training skin type model...")
train_skin_model()
print("ALL MODELS TRAINED SUCCESSFULLY")


