{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":5897669,"sourceType":"datasetVersion","datasetId":3387673}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport tensorflow as tf\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.layers import Conv2D, MaxPooling2D, Flatten, Dense, Dropout, BatchNormalization\nfrom tensorflow.keras.optimizers import Adam\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import classification_report, confusion_matrix\nfrom PIL import Image\n\n# Define path to your dataset\ndataset_dir = '/kaggle/input/strawberry-dataset/strawberryDataset'\n\n# Define image size and batch size\nimg_size = (224, 224)\nbatch_size = 32\n\n# Function to check if an image file is valid\ndef is_valid_image(file_path):\n    try:\n        img = Image.open(file_path)\n        img.verify()  # Verify that it is, in fact, an image\n        img.close()\n        return True\n    except (IOError, SyntaxError):\n        return False\n\n# Create a list to store paths of valid images\nvalid_images = {\"Pickable\": [], \"UnPickable\": []}\n\n# Check and collect valid images\nfor category in [\"Pickable\", \"UnPickable\"]:\n    category_path = os.path.join(dataset_dir, category)\n    for filename in os.listdir(category_path):\n        file_path = os.path.join(category_path, filename)\n        if is_valid_image(file_path):\n            valid_images[category].append(file_path)\n        else:\n            print(f\"Skipping invalid image: {file_path}\")\n\n# Data augmentation for training set\ntrain_datagen = ImageDataGenerator(\n    rescale=1./255,\n    rotation_range=20,\n    width_shift_range=0.2,\n    height_shift_range=0.2,\n    shear_range=0.2,\n    zoom_range=0.2,\n    horizontal_flip=True,\n    fill_mode='nearest',\n    validation_split=0.2  # Split data into 80% training and 20% validation\n)\n\n# No data augmentation for validation set, only rescaling\ntest_datagen = ImageDataGenerator(rescale=1./255, validation_split=0.2)\n\n# Load and augment data\ntrain_generator = train_datagen.flow_from_directory(\n    directory=dataset_dir,\n    target_size=img_size,\n    batch_size=batch_size,\n    class_mode='binary',\n    subset='training'  # Set as training data\n)\n\nvalidation_generator = test_datagen.flow_from_directory(\n    directory=dataset_dir,\n    target_size=img_size,\n    batch_size=batch_size,\n    class_mode='binary',\n    subset='validation'  # Set as validation data\n)\n\n# Define the model\nmodel = Sequential()\n\n# Input layer\nmodel.add(Conv2D(32, (3, 3), activation='relu', input_shape=(img_size[0], img_size[1], 3)))\n\n# Convolutional block with increasing filters\nmodel.add(Conv2D(64, (3, 3), activation='relu'))\nmodel.add(BatchNormalization())\nmodel.add(MaxPooling2D(pool_size=(2, 2)))\nmodel.add(Dropout(0.2))\n\nmodel.add(Conv2D(128, (3, 3), activation='relu'))\nmodel.add(BatchNormalization())\nmodel.add(MaxPooling2D(pool_size=(2, 2)))\nmodel.add(Dropout(0.2))\n\n# Flatten layer\nmodel.add(Flatten())\n\n# Fully connected layers\nmodel.add(Dense(512, activation='relu'))\nmodel.add(Dropout(0.2))\n\n# Output layer\nmodel.add(Dense(1, activation='sigmoid'))\n\n# Compile the model\nmodel.compile(optimizer=Adam(learning_rate=0.0001), loss='binary_crossentropy', metrics=['accuracy'])\n\n# Train the model\nhistory = model.fit(\n    train_generator,\n    epochs=20,  # Adjust number of epochs as needed\n    validation_data=validation_generator\n)\n\n# Print training accuracy and loss\ntrain_loss = history.history['loss'][-1]\ntrain_accuracy = history.history['accuracy'][-1]\nprint(f'Training Loss: {train_loss}, Training Accuracy: {train_accuracy}')\n\n# Evaluate the model on validation set\nval_loss, val_accuracy = model.evaluate(validation_generator)\nprint(f'Validation Loss: {val_loss}, Validation Accuracy: {val_accuracy}')\n\n# Plot training & validation accuracy values\nplt.figure(figsize=(12, 4))\nplt.subplot(1, 2, 1)\nplt.plot(history.history['accuracy'])\nplt.plot(history.history['val_accuracy'])\nplt.title('Model accuracy')\nplt.ylabel('Accuracy')\nplt.xlabel('Epoch')\nplt.legend(['Train', 'Validation'], loc='upper left')\n\n# Plot training & validation loss values\nplt.subplot(1, 2, 2)\nplt.plot(history.history['loss'])\nplt.plot(history.history['val_loss'])\nplt.title('Model loss')\nplt.ylabel('Loss')\nplt.xlabel('Epoch')\nplt.legend(['Train', 'Validation'], loc='upper left')\nplt.show()\n\n# Predict the validation set\nvalidation_generator.reset()\ny_pred = model.predict(validation_generator)\ny_pred = np.round(y_pred).astype(int).reshape(-1)\n\n# Get true labels\nY_val = validation_generator.classes\n\n# Print classification report\nprint('Classification Report')\nprint(classification_report(Y_val, y_pred, target_names=['Pickable', 'UnPickable']))\n\n# Compute and plot confusion matrix\ncm = confusion_matrix(Y_val, y_pred)\nprint('Confusion Matrix')\nprint(cm)\n\n# Plot confusion matrix\nplt.figure(figsize=(6, 6))\nplt.imshow(cm, interpolation='nearest', cmap='Blues')\nplt.title('Confusion Matrix')\nplt.colorbar()\ntick_marks = np.arange(2)\nplt.xticks(tick_marks, ['Pickable', 'UnPickable'], rotation=45)\nplt.yticks(tick_marks, ['Pickable', 'UnPickable'])\nplt.ylabel('True label')\nplt.xlabel('Predicted label')\nplt.show()\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-14T11:29:24.659191Z","iopub.execute_input":"2024-05-14T11:29:24.659608Z","iopub.status.idle":"2024-05-14T12:13:18.054751Z","shell.execute_reply.started":"2024-05-14T11:29:24.659572Z","shell.execute_reply":"2024-05-14T12:13:18.053376Z"},"trusted":true},"execution_count":null,"outputs":[]}]}