{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"}],"dockerImageVersionId":30558,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\ntrain = pd.read_csv(\"/kaggle/input/UBC-OCEAN/train.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-01-04T02:40:43.459608Z","iopub.execute_input":"2024-01-04T02:40:43.460307Z","iopub.status.idle":"2024-01-04T02:40:43.813561Z","shell.execute_reply.started":"2024-01-04T02:40:43.460273Z","shell.execute_reply":"2024-01-04T02:40:43.812758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"category_counts = train['label'].value_counts()\nprint(category_counts)","metadata":{"execution":{"iopub.status.busy":"2024-01-04T02:40:43.815083Z","iopub.execute_input":"2024-01-04T02:40:43.815387Z","iopub.status.idle":"2024-01-04T02:40:43.82921Z","shell.execute_reply.started":"2024-01-04T02:40:43.815361Z","shell.execute_reply":"2024-01-04T02:40:43.828381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport re\nimport glob\nimport numpy as np\nimport pandas as pd\nimport tensorflow as tf\nimport jax\nimport jax.numpy as jnp\nfrom jax import random\nimport optax\nimport flax.linen as nn\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom sklearn.preprocessing import LabelEncoder\nimport matplotlib.pyplot as plt\n\n# Constants\nIMAGE_DIR = '/kaggle/input/UBC-OCEAN/train_thumbnails/'\nTARGET_SIZE = (128, 128)\nBATCH_SIZE = 1\nNUM_EPOCHS = 10\nOUTPUT_SIZE = 5  # Number of classes\nL2_REGULARIZATION = 1e-4\nDROP_OUT_RATE = 0.5\nIMAGE_HEIGHT, IMAGE_WIDTH = TARGET_SIZE\nCHANNELS = 3\n\n# Load labels\nlabels_df = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\n\nlabel_dict = dict(zip(labels_df['image_id'], labels_df['label']))\n\n# Prepare Label Encoder\nlabel_encoder = LabelEncoder()\nlabel_encoder.fit(['CC', 'EC', 'HGSC', 'LGSC', 'MC'])\n\n# Custom batch normalization layer\nclass CustomBatchNorm(nn.Module):\n    epsilon: float = 1e-5\n    dtype: jnp.dtype = jnp.float32\n\n    @nn.compact\n    def __call__(self, x):\n        mean = jnp.mean(x, axis=(0, 1, 2), keepdims=True)\n        var = jnp.var(x, axis=(0, 1, 2), keepdims=True)\n        x = (x - mean) / jnp.sqrt(var + self.epsilon)\n        gamma = self.param('gamma', nn.initializers.ones, x.shape[-1])\n        beta = self.param('beta', nn.initializers.zeros, x.shape[-1])\n        return x * gamma + beta\n    \n# CNN model\nclass YourModel(nn.Module):\n    dropout_rate: float\n\n    @nn.compact\n    def __call__(self, x, deterministic=False):\n        x = nn.Conv(features=32, kernel_size=(3, 3))(x)\n        x = CustomBatchNorm()(x)\n        x = nn.relu(x)\n        x = nn.max_pool(x, window_shape=(2, 2), strides=(2, 2))\n        num_elements = x.shape[1] * x.shape[2] * x.shape[3]\n        x = x.reshape((x.shape[0], num_elements))\n\n        x = nn.Dropout(rate=self.dropout_rate, name='dropout')(x, deterministic=deterministic)\n        x = nn.Dense(features=OUTPUT_SIZE)(x)\n        \n        # Return both logits and the intermediate representation\n        return x, x\n    \ndef l2_regularization_loss(params, l2_factor=L2_REGULARIZATION):\n    l2_loss = 0.0\n    for param in jax.tree_util.tree_leaves(params):\n        l2_loss += jnp.sum(param**2)\n    return l2_factor * l2_loss\n\ndef load_and_preprocess_from_path_label(path, label):\n    # Load the image file\n    image = tf.io.read_file(path)\n    # Decode the image and set the channels to 3 (for color images)\n    image = tf.image.decode_png(image, channels=3)\n    # Resize the image to the desired dimensions\n    image = tf.image.resize(image, TARGET_SIZE)\n    # Normalize the pixel values\n    image = tf.cast(image, tf.float32) / 255.0\n    return image, label\n\n# Modify the dataset loading function to pass image paths as tensors\ndef load_and_create_dataset(image_dir, labels_df, batch_size=32):\n    # Get the image file paths and labels\n    image_paths = labels_df['image_id'].apply(lambda x: os.path.join(image_dir, str(x) + '_thumbnail.png'))\n    labels = labels_df['label'].values\n\n    # Filter out missing files\n    existing_image_paths = [path for path in image_paths if os.path.exists(path)]\n\n    # Make sure labels have the same number of elements as existing_image_paths\n    labels = labels[:len(existing_image_paths)]\n\n    # Create a dataset of image file paths and labels\n    image_labels = tf.data.Dataset.from_tensor_slices((existing_image_paths, labels))\n\n    # Load and preprocess images in parallel using num_parallel_calls\n    dataset = image_labels.map(load_and_preprocess_from_path_label, num_parallel_calls=tf.data.AUTOTUNE)\n\n    # Shuffle and batch the dataset\n    dataset = dataset.shuffle(buffer_size=10000)\n    dataset = dataset.batch(batch_size)\n    dataset = dataset.prefetch(buffer_size=tf.data.AUTOTUNE)\n\n    return dataset\n\nlabel_dict = dict(zip(labels_df['image_id'].astype(int), labels_df['label']))\n\ndef extract_id_from_path(path):\n    \"\"\"Extracts the image ID (all numbers before the underscore) from the file path.\"\"\"\n    filename = os.path.basename(path)\n    match = re.match(r'(\\d+)_thumbnail\\.png', filename)\n    if match:\n        return int(match.group(1))  # Convert to integer\n    else:\n        raise ValueError(f\"Can't extract an ID from the filename: {filename}\")\n        return None\n\ndefault_label = 'unknown'\n\ndef get_labels_from_paths(paths):\n    \"\"\"For each image path, extract the ID and lookup its label.\"\"\"\n    ids = [extract_id_from_path(path.decode('utf-8')) for path in paths.numpy()]\n    labels = []\n    for id in ids:\n        if id is not None:\n            labels.append(label_dict.get(id, default_label))\n        else:\n            labels.append(default_label)  # or handle this case as needed\n    return labels\n\n# Compute class weights for imbalanced classes\nclass_labels = label_encoder.transform(labels_df.label)\nclass_weights = compute_class_weight('balanced', classes=np.unique(class_labels), y=class_labels)\nclass_weights_dict = dict(enumerate(class_weights))\n\n# Instantiate the model\nmock_input = jnp.ones((1, IMAGE_HEIGHT, IMAGE_WIDTH, CHANNELS), dtype=jnp.float32)\nmodel = YourModel(dropout_rate=0.5)\nexample_input = jnp.ones((BATCH_SIZE, IMAGE_HEIGHT, IMAGE_WIDTH, CHANNELS), jnp.float32)\n\n# Initialize the RNGs and split for parameters and dropout\nrng = random.PRNGKey(0)\nparam_rng, dropout_rng = random.split(rng)\n\n# Initialize the model parameters\ninit_output = model.init({'params': param_rng, 'dropout': dropout_rng}, mock_input)\nparams = init_output['params']\n\n# Define optimizer\noptimizer = optax.adam(learning_rate=1e-3)\nopt_state = optimizer.init(params)\n\n# Loss function\ndef loss_fn(params, batch_images, batch_labels, rng):\n    logits, _ = model.apply({'params': params}, batch_images, rngs={'dropout': rng}, deterministic=False)\n    loss = -jnp.sum(jax.nn.log_softmax(logits) * jax.nn.one_hot(batch_labels, OUTPUT_SIZE), axis=-1)\n    mean_loss = loss.mean() + l2_regularization_loss(params)\n    return mean_loss\n\n# Training step\n@jax.jit\ndef train_step(params, opt_state, batch_images, batch_labels, rng):\n    grads = jax.grad(loss_fn)(params, batch_images, batch_labels, rng)\n    updates, new_opt_state = optimizer.update(grads, opt_state)\n    new_params = optax.apply_updates(params, updates)\n    return new_params, new_opt_state\n\n# Prepare the dataset\nall_image_paths = glob.glob(os.path.join(IMAGE_DIR, '*.png'))\ntrain_dataset = load_and_create_dataset(IMAGE_DIR, labels_df, batch_size=BATCH_SIZE)\n\n# Training loop\nfor epoch in range(NUM_EPOCHS):\n    for batch_images, batch_labels in train_dataset:\n        batch_images_np = jnp.array(batch_images.numpy())\n        encoded_labels = label_encoder.transform(batch_labels)\n        batch_labels_np = jnp.array(encoded_labels, dtype=jnp.int32)\n        \n        # Update the optimizer state using jax.grad and optax\n        grads = jax.grad(loss_fn)(params, batch_images_np, batch_labels_np, rng)\n        updates, new_opt_state = optimizer.update(grads, opt_state)\n        new_params = optax.apply_updates(params, updates)\n        \n        # Update params and opt_state\n        params = new_params\n        opt_state = new_opt_state\n    \n    print(f\"Epoch {epoch + 1} completed.\")\n\n# Save the trained model parameters if needed\nnp.save('trained_params.npy', params)\n\n# Simulate make_gradcam_heatmap function\ndef make_gradcam_heatmap(img_array, model, last_conv_layer_name):\n    heatmap = np.random.random((img_array.shape[1], img_array.shape[2]))  # Placeholder\n    return heatmap\n\n# Simulate display_gradcam function\ndef display_gradcam(img, heatmap):\n    plt.imshow(img)\n    plt.imshow(heatmap, cmap='jet', alpha=0.5)  # Overlay heatmap\n    plt.show()\n\n# Select a test image\ntest_image_path = all_image_paths[0]\n\n# Load and preprocess the test image\ntest_img = tf.keras.preprocessing.image.load_img(test_image_path, target_size=(128, 128))\ntest_img_array = tf.keras.preprocessing.image.img_to_array(test_img)\ntest_img_array = np.expand_dims(test_img_array, axis=0)\n\n# Predict and get the heatmap\nheatmap = make_gradcam_heatmap(test_img_array, model, 'last_conv_layer')\n\n# Display the heatmap\ndisplay_gradcam(test_img, heatmap)","metadata":{"execution":{"iopub.status.busy":"2024-01-04T02:40:43.830509Z","iopub.execute_input":"2024-01-04T02:40:43.830843Z","iopub.status.idle":"2024-01-04T02:54:22.363276Z","shell.execute_reply.started":"2024-01-04T02:40:43.830809Z","shell.execute_reply":"2024-01-04T02:54:22.362359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport os\nimport tensorflow as tf\nimport jax\nfrom jax import numpy as jnp\n\n# Define the path for the test images\ntest_image_dir = '/kaggle/input/UBC-OCEAN/test_thumbnails/'\n\n# Load test images\ntest_image_paths = glob.glob(os.path.join(test_image_dir, '*.png'))\n\n# Define function for loading and preprocessing test images\ndef load_and_preprocess_test_images(image_paths):\n    images = []\n    ids = []\n    for path in image_paths:\n        # Load and preprocess the image\n        img = tf.keras.preprocessing.image.load_img(path, target_size=TARGET_SIZE)\n        img_array = tf.keras.preprocessing.image.img_to_array(img)\n        img_array = img_array / 255.0  # Normalize to [0,1]\n        images.append(img_array)\n        \n        # Extract image ID\n        image_id = os.path.basename(path).split('_')[0]\n        ids.append(image_id)\n    \n    return np.array(images), ids\n\n# Load and preprocess test images\ntest_images, image_ids = load_and_preprocess_test_images(test_image_paths)\n\n# Convert images to JAX arrays\ntest_images_jax = jnp.array(test_images)\n\n# Predict classes using the trained model\nlogits, _ = model.apply({'params': params}, test_images_jax, rngs={'dropout': rng}, deterministic=True)\npredicted_classes = jnp.argmax(logits, axis=1)\n\n# Convert numeric predictions back to original class labels\npredicted_labels = label_encoder.inverse_transform(predicted_classes)\n\n# Create a DataFrame with image_id and predicted labels\nsubmission_df = pd.DataFrame({\n    'image_id': image_ids,\n    'label': predicted_labels\n})\n\nsubmission_df","metadata":{"execution":{"iopub.status.busy":"2024-01-04T02:54:22.365017Z","iopub.execute_input":"2024-01-04T02:54:22.365312Z","iopub.status.idle":"2024-01-04T02:54:22.763166Z","shell.execute_reply.started":"2024-01-04T02:54:22.365286Z","shell.execute_reply":"2024-01-04T02:54:22.762274Z"},"trusted":true},"execution_count":null,"outputs":[]}]}