{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":11340354,"sourceType":"datasetVersion","datasetId":7094759},{"sourceId":334579,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":280103,"modelId":301015}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras.layers import (\n    Input, Conv2D, MaxPooling2D, Dropout, Flatten, Dense, ReLU\n)\nfrom tensorflow.keras.models import Model\n\ndef SketchANet(input_shape=(225, 225, 1), num_classes=250):\n    inputs = Input(shape=input_shape)\n\n    x = Conv2D(64, (15, 15), strides=3, padding='valid')(inputs)  # conv1\n    x = ReLU()(x)\n    x = MaxPooling2D(pool_size=(3, 3), strides=2)(x)\n\n    x = Conv2D(128, (5, 5), strides=1, padding='valid')(x)        # conv2\n    x = ReLU()(x)\n    x = MaxPooling2D(pool_size=(3, 3), strides=2)(x)\n\n    x = Conv2D(256, (3, 3), strides=1, padding='same')(x)  # conv3\n    x = ReLU()(x)\n    x = Conv2D(256, (3, 3), strides=1, padding='same')(x)  # conv4\n    x = ReLU()(x)\n    x = Conv2D(256, (3, 3), strides=1, padding='same')(x)  # conv5\n    x = ReLU()(x)\n    x = MaxPooling2D(pool_size=(3, 3), strides=2)(x)\n\n    x = Conv2D(512, (7, 7), strides=1, padding='valid')(x)  # conv6\n    x = ReLU()(x)\n    x = Dropout(0.5)(x)\n\n    x = Conv2D(512, (1, 1), strides=1, padding='valid')(x) #conv7\n    x = ReLU()(x)\n    x = Dropout(0.5)(x)\n\n    x = Flatten()(x)   #conv8\n    outputs = Dense(num_classes, activation='softmax')(x)\n\n    return Model(inputs, outputs)\n\n\nmodel = SketchANet()\nmodel.summary()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:52:59.667151Z","iopub.execute_input":"2025-04-12T12:52:59.667537Z","iopub.status.idle":"2025-04-12T12:53:21.64582Z","shell.execute_reply.started":"2025-04-12T12:52:59.667487Z","shell.execute_reply":"2025-04-12T12:53:21.644794Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf\nimport os\nfrom tensorflow.keras.callbacks import EarlyStopping, ModelCheckpoint, ReduceLROnPlateau\n\ndata_dir = \"/kaggle/input/sketch/png\"\nimage_height, image_width = 256, 256\ncrop_height, crop_width = 225, 225\nbatch_size = 135\nseed = 123\ninitial_learning_rate = 1e-4\nepochs = 100\nmax_translation = 32.0 / crop_width\n\ndef binarize_and_invert(x):\n    x = tf.where(x < 0.9, 0.0, 1.0) \n    return 1.0 - x\n\ntrain_ds = tf.keras.preprocessing.image_dataset_from_directory(\n    data_dir,\n    validation_split=0.2,\n    subset=\"training\",\n    seed=seed,\n    image_size=(image_height, image_width),\n    batch_size=batch_size,\n    label_mode='categorical',\n    color_mode='grayscale'\n)\n\nval_ds = tf.keras.preprocessing.image_dataset_from_directory(\n    data_dir,\n    validation_split=0.2,\n    subset=\"validation\",\n    seed=seed,\n    image_size=(image_height, image_width),\n    batch_size=batch_size,\n    label_mode='categorical',\n    color_mode='grayscale'\n)\n\nclass_names = train_ds.class_names\nnum_classes = len(class_names)\nprint(f\"Classes: {num_classes}\")\n\ndata_augmentation = tf.keras.Sequential([\n    tf.keras.layers.Rescaling(1. / 255),\n    tf.keras.layers.RandomFlip(\"horizontal\"),\n    tf.keras.layers.RandomCrop(crop_height, crop_width),\n])\n\nval_preprocessing = tf.keras.Sequential([\n    tf.keras.layers.Rescaling(1. / 255),\n    tf.keras.layers.CenterCrop(crop_height, crop_width)\n])\n\nAUTOTUNE = tf.data.AUTOTUNE\n\ndef process_train(x, y):\n    x = data_augmentation(x)\n    x = binarize_and_invert(x)\n    return x, y\n\ndef process_val(x, y):\n    x = val_preprocessing(x)\n    x = binarize_and_invert(x)\n    return x, y\n\ntrain_ds = (\n    train_ds\n    .map(process_train, num_parallel_calls=AUTOTUNE)\n    .cache()\n    .shuffle(1000)\n    .prefetch(AUTOTUNE)\n)\n\nval_ds = (\n    val_ds\n    .map(process_val, num_parallel_calls=AUTOTUNE)\n    .cache()\n    .prefetch(AUTOTUNE)\n)\n\ninput_shape = (crop_height, crop_width, 1)\nmodel = SketchANet(input_shape=input_shape, num_classes=num_classes)\n\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(learning_rate=initial_learning_rate),\n    loss=tf.keras.losses.CategoricalCrossentropy(from_logits=False),\n    metrics=[\n        'accuracy',\n        tf.keras.metrics.TopKCategoricalAccuracy(k=5, name=\"top-5-accuracy\")\n    ]\n)\n\ncallbacks = [\n    EarlyStopping(monitor='val_accuracy', patience=8, restore_best_weights=True),\n    ReduceLROnPlateau(monitor='val_loss', factor=0.5, patience=4, verbose=1, min_lr=1e-6),\n    ModelCheckpoint(\n        filepath='outputs/sketchanet_best.keras',\n        monitor='val_accuracy',\n        mode='max',\n        save_best_only=True,\n        verbose=1\n    )\n]\nmodel.fit(\n    train_ds,\n    validation_data=val_ds,\n    epochs=epochs,\n    callbacks=callbacks\n)\n\nmodel.save('/kaggle/working/sketchanet_final.keras')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T14:22:29.296922Z","iopub.execute_input":"2025-04-12T14:22:29.297357Z","iopub.status.idle":"2025-04-12T14:34:47.025356Z","shell.execute_reply.started":"2025-04-12T14:22:29.29733Z","shell.execute_reply":"2025-04-12T14:34:47.024136Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf\nimport numpy as np\nimport os\nMODEL_PATHS = [\n    '/kaggle/input/sketch-a-net/keras/default/1/sketchanet_best1.keras', \n    '/kaggle/input/sketch-a-net/keras/default/1/sketchanet_best2.keras',\n    '/kaggle/input/sketch-a-net/keras/default/1/sketchanet_best3.keras',\n    '/kaggle/input/sketch-a-net/keras/default/1/sketchanet_best4.keras'\n]\nCROP_HEIGHT, CROP_WIDTH = 225, 225 \nNUM_CLASSES = 250\n\nprint(f\"Loading {len(MODEL_PATHS)} individual models...\")\nmodels = []\nfor path in MODEL_PATHS:\n    try:\n        model = tf.keras.models.load_model(path)\n        model.trainable = False # Freeze weights - crucial!\n        model._name = f\"ensemble_member_{len(models)+1}_{model.name}\" # Rename for clarity\n        models.append(model)\n    except Exception as e:\n        print(f\"Error loading model {path}: {e}\")\nif not models:\n    print(\"No models were loaded. Exiting.\")\n    exit()\nprint(f\"{len(models)} models loaded successfully.\")\n\nprint(\"Creating combined ensemble model...\")\n\nensemble_input = tf.keras.Input(shape=(CROP_HEIGHT, CROP_WIDTH, 1), name='ensemble_input')\n\noutputs = [model(ensemble_input) for model in models]\nensemble_output = tf.keras.layers.Average(name='ensemble_average')(outputs)\nensemble_model = tf.keras.Model(inputs=ensemble_input, outputs=ensemble_output, name='SketchANet_Ensemble')\n\nprint(\"Ensemble model created.\")\nensemble_model.summary()\n\nENSEMBLE_SAVE_PATH = '/kaggle/working/sketchanet_ensemble_model.keras'\nprint(f\"Saving combined ensemble model to {ENSEMBLE_SAVE_PATH}\")\nensemble_model.save(ENSEMBLE_SAVE_PATH)\nprint(\"Combined model saved.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T17:54:33.841718Z","iopub.execute_input":"2025-04-12T17:54:33.842074Z","iopub.status.idle":"2025-04-12T17:54:36.943198Z","shell.execute_reply.started":"2025-04-12T17:54:33.842042Z","shell.execute_reply":"2025-04-12T17:54:36.942195Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf\nimport numpy as np\nimport os\nimport time\nENSEMBLE_MODEL_PATH = '/kaggle/working/sketchanet_ensemble_model.keras'\n\nDATA_DIR = \"/kaggle/input/sketch/png\" \nIMAGE_HEIGHT, IMAGE_WIDTH = 256, 256\nCROP_HEIGHT, CROP_WIDTH = 225, 225\nNUM_CLASSES = 250\nBATCH_SIZE = 32 \nVALIDATION_SPLIT = 0.2\n\ndef binarize_and_invert(x):\n    x = tf.where(x < 0.9, 0.0, 1.0)\n    return 1.0 - x\n\nrescale_layer = tf.keras.layers.Rescaling(1. / 255)\n\n@tf.function\ndef preprocess_image(image):\n    image = rescale_layer(image)\n    image = binarize_and_invert(image)\n    return image\n    \n# --- Multi-view Cropping Function ---\n@tf.function\ndef get_10_crop_views(image):\n    image = tf.ensure_shape(image, (IMAGE_HEIGHT, IMAGE_WIDTH, 1))\n    views = []\n    target_h, target_w = CROP_HEIGHT, CROP_WIDTH\n    # Center Crop\n    views.append(tf.image.resize_with_crop_or_pad(image, target_h, target_w))\n    # Top-Left Crop\n    views.append(image[0:target_h, 0:target_w, :])\n    # Top-Right Crop\n    views.append(image[0:target_h, (IMAGE_WIDTH - target_w):IMAGE_WIDTH, :])\n    # Bottom-Left Crop\n    views.append(image[(IMAGE_HEIGHT - target_h):IMAGE_HEIGHT, 0:target_w, :])\n    # Bottom-Right Crop\n    views.append(image[(IMAGE_HEIGHT - target_h):IMAGE_HEIGHT, (IMAGE_WIDTH - target_w):IMAGE_WIDTH, :])\n\n    stacked_views = tf.stack(views) # Shape: (5, CROP_H, CROP_W, 1)\n    flipped_views = tf.image.flip_left_right(stacked_views)\n    all_views = tf.concat([stacked_views, flipped_views], axis=0) # Shape: (10, CROP_H, CROP_W, 1)\n    return all_views\n\n@tf.function\ndef batch_get_10_crop_views(images):\n    batch_views = tf.map_fn(get_10_crop_views, images, fn_output_signature=tf.TensorSpec((10, CROP_HEIGHT, CROP_WIDTH, 1), tf.float32))\n    return tf.reshape(batch_views, [-1, CROP_HEIGHT, CROP_WIDTH, 1])\n\nprint(f\"Loading combined ensemble model from {ENSEMBLE_MODEL_PATH}...\")\nstart_load = time.time()\nif not os.path.exists(ENSEMBLE_MODEL_PATH):\n    print(f\"ERROR: Ensemble model not found at {ENSEMBLE_MODEL_PATH}\")\n    exit()\ntry:\n    ensemble_model = tf.keras.models.load_model(ENSEMBLE_MODEL_PATH)\n    print(f\"Combined model loaded in {time.time() - start_load:.2f} seconds.\")\nexcept Exception as e:\n    print(f\"Error loading combined model: {e}\")\n    exit()\n\nprint(\"Loading validation data...\")\nstart_data = time.time()\nval_ds = tf.keras.utils.image_dataset_from_directory(\n    DATA_DIR,\n    labels='inferred', label_mode='int',\n    validation_split=VALIDATION_SPLIT, subset=\"validation\",\n    seed=123,\n    image_size=(IMAGE_HEIGHT, IMAGE_WIDTH),\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    color_mode='grayscale'\n)\nval_ds = val_ds.map(lambda x, y: (preprocess_image(x), y), num_parallel_calls=tf.data.AUTOTUNE)\nval_ds = val_ds.prefetch(tf.data.AUTOTUNE)\nprint(f\"Validation data loaded in {time.time() - start_data:.2f} seconds.\")\n\n@tf.function\ndef predict_batch_single_ensemble(images):\n    views_batch = batch_get_10_crop_views(images) \n    avg_model_preds = ensemble_model(views_batch, training=False) # Shape: (B*10, NUM_CLASSES)\n    B = tf.shape(images)[0] \n    reshaped_preds = tf.reshape(avg_model_preds, (B, 10, NUM_CLASSES))\n    final_preds = tf.reduce_mean(reshaped_preds, axis=1)\n    return final_preds\n\nall_preds = []\nall_labels = []\nnum_batches = len(val_ds)\n\nprint(f\"Starting prediction on {num_batches} batches using the combined ensemble model...\")\nstart_pred_loop = time.time()\n\nfor batch_idx, (processed_images, labels) in enumerate(val_ds):\n    start_batch_time = time.time()\n    probs = predict_batch_single_ensemble(processed_images)\n\n    all_preds.append(probs.numpy())\n    all_labels.extend(labels.numpy())\n\n    batch_time = time.time() - start_batch_time\n    print(f\"Batch {batch_idx+1}/{num_batches} processed in {batch_time:.2f} seconds.\")\n\nprint(f\"Prediction loop finished in {time.time() - start_pred_loop:.2f} seconds.\")\n\npredicted_class_indices = np.argmax(np.concatenate(all_preds, axis=0), axis=1)\n\nif all_labels:\n    correct_predictions = sum(1 for pred, true in zip(predicted_class_indices, all_labels) if pred == true)\n    total_images = len(all_labels)\n    accuracy = correct_predictions / total_images\n    print(f\"\\nCombined Ensemble Model Accuracy on Validation Set: {accuracy:.4f}\")\n\nprint(\"\\nExample Predictions (Class Index):\")\nfor i in range(min(5, len(predicted_class_indices))):\n    true_label = all_labels[i]\n    print(f\"Image {i+1}: Predicted Class = {predicted_class_indices[i]}, True Class = {true_label}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T17:59:43.037265Z","iopub.execute_input":"2025-04-12T17:59:43.037659Z","iopub.status.idle":"2025-04-12T18:00:29.555229Z","shell.execute_reply.started":"2025-04-12T17:59:43.037627Z","shell.execute_reply":"2025-04-12T18:00:29.55425Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}