{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.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":21154,"databundleVersionId":1243559,"sourceType":"competition"},{"sourceId":13209338,"sourceType":"datasetVersion","datasetId":8271423}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 1. Import packages and setup environment (CPU/GPU/TPU)","metadata":{}},{"cell_type":"code","source":"## For TPU environment (install missing packages / reinstall tensorflow to solve NaN topic during training / restart kernel)\n\nimport IPython\nimport tensorflow as tf\n\nif len(tf.config.experimental.list_logical_devices('TPU'))>0:\n    !pip install -q tensorflow-tpu -f https://storage.googleapis.com/libtpu-tf-releases/index.html --force-reinstall\n    !pip install -q pydot\n    !pip install -q -U keras-tuner\n    IPython.Application.instance().kernel.do_shutdown(True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-29T10:22:12.459376Z","iopub.execute_input":"2025-09-29T10:22:12.459623Z","iopub.status.idle":"2025-09-29T10:22:29.885993Z","shell.execute_reply.started":"2025-09-29T10:22:12.459590Z","shell.execute_reply":"2025-09-29T10:22:29.885074Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Import packages\n\n# General packages\nimport os\nimport re\nimport math\nfrom tqdm import tqdm\n\n# Main data handling and visualization packages\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom matplotlib.pyplot import imshow\n\n# Tensorflow modules\nimport tensorflow as tf\nimport keras_tuner as kt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-29T10:22:29.887486Z","iopub.execute_input":"2025-09-29T10:22:29.887927Z","iopub.status.idle":"2025-09-29T10:22:30.406741Z","shell.execute_reply.started":"2025-09-29T10:22:29.887908Z","shell.execute_reply":"2025-09-29T10:22:30.406105Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Detect hardware (CPU/GPU/TPU), setup environment and return appropriate distribution strategy\n\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect(tpu='local') # set tpu is local as it should be available in the VM\n    print('✅ Running on TPU ', tpu.master())\nexcept:\n    print('❌ Using CPU/GPU')\n    tpu = None\n\nif tpu:\n    strategy = tf.distribute.TPUStrategy(tpu)\nelse:\n    strategy = tf.distribute.get_strategy() # default distribution strategy in Tensorflow. Works on CPU and single GPU.\n\nprint(\"REPLICAS: \", strategy.num_replicas_in_sync)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-29T10:22:30.407525Z","iopub.execute_input":"2025-09-29T10:22:30.407770Z","iopub.status.idle":"2025-09-29T10:22:30.703053Z","shell.execute_reply.started":"2025-09-29T10:22:30.407744Z","shell.execute_reply":"2025-09-29T10:22:30.702398Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Load and Preprocess Data","metadata":{}},{"cell_type":"code","source":"## Image parameters and classes\n\nimg_size_source = 512 # 192/224/331/512\nimg_size_model = 528 #224\nIMAGE_SIZE = [img_size_model, img_size_model]\nCLASSES = ['pink primrose',    'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea',     'wild geranium',     'tiger lily',           'moon orchid',              'bird of paradise', 'monkshood',        'globe thistle',         # 00 - 09\n           'snapdragon',       \"colt's foot\",               'king protea',      'spear thistle', 'yellow iris',       'globe-flower',         'purple coneflower',        'peruvian lily',    'balloon flower',   'giant white arum lily', # 10 - 19\n           'fire lily',        'pincushion flower',         'fritillary',       'red ginger',    'grape hyacinth',    'corn poppy',           'prince of wales feathers', 'stemless gentian', 'artichoke',        'sweet william',         # 20 - 29\n           'carnation',        'garden phlox',              'love in the mist', 'cosmos',        'alpine sea holly',  'ruby-lipped cattleya', 'cape flower',              'great masterwort', 'siam tulip',       'lenten rose',           # 30 - 39\n           'barberton daisy',  'daffodil',                  'sword lily',       'poinsettia',    'bolero deep blue',  'wallflower',           'marigold',                 'buttercup',        'daisy',            'common dandelion',      # 40 - 49\n           'petunia',          'wild pansy',                'primula',          'sunflower',     'lilac hibiscus',    'bishop of llandaff',   'gaura',                    'geranium',         'orange dahlia',    'pink-yellow dahlia',    # 50 - 59\n           'cautleya spicata', 'japanese anemone',          'black-eyed susan', 'silverbush',    'californian poppy', 'osteospermum',         'spring crocus',            'iris',             'windflower',       'tree poppy',            # 60 - 69\n           'gazania',          'azalea',                    'water lily',       'rose',          'thorn apple',       'morning glory',        'passion flower',           'lotus',            'toad lily',        'anthurium',             # 70 - 79\n           'frangipani',       'clematis',                  'hibiscus',         'columbine',     'desert-rose',       'tree mallow',          'magnolia',                 'cyclamen ',        'watercress',       'canna lily',            # 80 - 89\n           'hippeastrum ',     'bee balm',                  'pink quill',       'foxglove',      'bougainvillea',     'camellia',             'mallow',                   'mexican petunia',  'bromelia',         'blanket flower',        # 90 - 99\n           'trumpet creeper',  'blackberry lily',           'common tulip',     'wild rose']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-29T10:22:30.703748Z","iopub.execute_input":"2025-09-29T10:22:30.703970Z","iopub.status.idle":"2025-09-29T10:22:30.711295Z","shell.execute_reply.started":"2025-09-29T10:22:30.703951Z","shell.execute_reply":"2025-09-29T10:22:30.710367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Helper functions for reading the data from TFRecords\n\ndef parse_data(raw_data):\n    feature_description = {\n        'image':tf.io.FixedLenFeature([],tf.string),\n        'id':tf.io.FixedLenFeature([],tf.string),\n        'class':tf.io.FixedLenFeature([],tf.int64),\n    }\n    parsed_data = tf.io.parse_single_example(raw_data,feature_description)\n    image = tf.io.decode_jpeg(parsed_data['image'],channels=3)\n    image = tf.image.resize(image,IMAGE_SIZE)\n    image = tf.cast(image,tf.float32)/255.0\n    label = parsed_data['class']\n    return image,label\n\ndef parse_data_without_label(raw_data):\n    feature_description = {\n        'image':tf.io.FixedLenFeature([],tf.string),\n        'id':tf.io.FixedLenFeature([],tf.string),\n    }\n    parsed_data = tf.io.parse_single_example(raw_data, feature_description) \n    image = tf.io.decode_jpeg(parsed_data['image'], channels=3)\n    image = tf.image.resize(image, IMAGE_SIZE)\n    image = tf.cast(image, tf.float32) / 255.0\n    return image\n\ndef parse_data_id(raw_data):\n    feature_description = {\n        'image':tf.io.FixedLenFeature([],tf.string),\n        'id':tf.io.FixedLenFeature([],tf.string),\n    }\n    parsed_data = tf.io.parse_single_example(raw_data, feature_description) \n    data_id = parsed_data['id']\n    return data_id\n\ndef data_augment(image, label):\n    # Augmentations that work well for flowers\n    image = tf.image.random_flip_left_right(image)\n    image = tf.image.random_brightness(image, 0.2)\n    image = tf.image.random_contrast(image, 0.8, 1.2)\n    image = tf.image.random_saturation(image, 0.7, 1.3)\n    image = tf.image.random_hue(image, 0.08)\n    \n    # Random rotation\n    if tf.random.uniform([]) > 0.5:\n        k = tf.random.uniform([], 0, 4, dtype=tf.int32)\n        image = tf.image.rot90(image, k=k)\n    \n    # Random zoom\n    if tf.random.uniform([]) > 0.5:\n        scale = tf.random.uniform([], 0.8, 1.2)\n        new_size = tf.cast(tf.cast(IMAGE_SIZE, tf.float32) * scale, tf.int32)\n        image = tf.image.resize(image, new_size)\n        image = tf.image.resize_with_crop_or_pad(image, IMAGE_SIZE[0], IMAGE_SIZE[1])\n    \n    image = tf.clip_by_value(image, 0.0, 1.0)\n    return image, label\n\ndef count_data_items(filenames):\n    n = [int(re.compile(r\"-([0-9]*)\\.\").search(filename).group(1)) \n         for filename in filenames]\n    return np.sum(n)\n\ndef imshow_title(image, title=None, color='black'):\n    '''displays an image with a corresponding label'''\n    if len(image.shape) > 3:\n        image = tf.squeeze(image, axis=0)\n    image = image#*255.0\n    #image = tf.reverse(image, axis=[-1]) # Convert back from bgr (vgg16 standard) to rgb (imshow standard)\n    plt.imshow(image)\n    if title:\n        plt.title(title, fontsize=10, color=color)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-29T10:22:30.712176Z","iopub.execute_input":"2025-09-29T10:22:30.712556Z","iopub.status.idle":"2025-09-29T10:22:30.733622Z","shell.execute_reply.started":"2025-09-29T10:22:30.712527Z","shell.execute_reply":"2025-09-29T10:22:30.732911Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Create Datasets from source TFRecords files\n\nBATCH_SIZE = 32\nVAL_BATCH_SIZE = 32\n\nfor source in ['train', 'val', 'test']:\n    globals()[f\"{source}_files\"] = tf.io.gfile.glob(f\"/kaggle/input/tpu-getting-started/tfrecords-jpeg-{img_size_source}x{img_size_source}/{source}/*.tfrec\")\n    globals()[f\"{source}_ds\"] = tf.data.TFRecordDataset(globals()[f\"{source}_files\"])\n    if source == 'train':\n        train_ds = train_ds.map(parse_data, num_parallel_calls=tf.data.AUTOTUNE)\n        train_ds = train_ds.map(data_augment, num_parallel_calls=tf.data.AUTOTUNE)\n        train_ds = train_ds.shuffle(2048).batch(BATCH_SIZE, drop_remainder=True).prefetch(tf.data.AUTOTUNE)\n    if source == 'val':\n        val_ds = val_ds.map(parse_data)\n        val_ds = val_ds.batch(VAL_BATCH_SIZE, drop_remainder=True)\n    if source == 'test':\n        test_images_ds = test_ds.map(parse_data_without_label).batch(1)\n        test_ids_ds = test_ds.map(parse_data_id).batch(1)\n\nnum_training = count_data_items(train_files)\nnum_validation = count_data_items(val_files)\nnum_test = count_data_items(test_files)\nsteps_per_epoch = num_training // BATCH_SIZE\nvalidation_steps = num_validation // VAL_BATCH_SIZE\n\nprint('Size of train dataset: '+ str(num_training) + ' images and labels in ' + str(steps_per_epoch) + ' batches')\nprint('Size of val dataset: '+ str(num_validation) + ' images and labels in ' + str(validation_steps) + ' batches')\nprint('Size of test dataset: '+ str(num_test) + ' images')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-29T10:22:30.734642Z","iopub.execute_input":"2025-09-29T10:22:30.734954Z","iopub.status.idle":"2025-09-29T10:22:31.512473Z","shell.execute_reply.started":"2025-09-29T10:22:30.734927Z","shell.execute_reply":"2025-09-29T10:22:31.511432Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Explore Data","metadata":{}},{"cell_type":"code","source":"## Check correctness of data\n\nfor i, image_label_batch in enumerate(tqdm(train_ds)):\n    if i == steps_per_epoch: break # Break needed as infinite dataset\n    for j in range(image_label_batch[0].shape[0]):\n        image = image_label_batch[0][j]\n        if tf.reduce_sum(tf.cast(tf.math.is_nan(image), tf.int32)) > 0:\n            print(f\"NaN value is present in image {j} of batch {i}\")\n        if tf.reduce_min(image) < 0.0:\n            print(f\"Max out-of-range value is present in label {j} of batch {i}\")\n        if tf.reduce_max(image) > 1.0:\n            print(f\"Max out-of-range value is present in label {j} of batch {i}\")\n\n    label_batch = image_label_batch[1]\n    if tf.reduce_sum(tf.cast(tf.math.is_nan(tf.cast(label_batch, np.float32)), tf.int32)) > 0:\n        print(f\"NaN value is present in label {j} of batch {i}\")\n    if tf.math.reduce_min(label_batch) < 0:\n        print(f\"Min out-of-range value is present in labels of batch {i}\")\n    if tf.math.reduce_max(label_batch) > 103:\n        print(f\"Max out-of-range value is present in labels of batch {i}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-29T10:22:31.514461Z","iopub.execute_input":"2025-09-29T10:22:31.514855Z","iopub.status.idle":"2025-09-29T10:24:39.720949Z","shell.execute_reply.started":"2025-09-29T10:22:31.514830Z","shell.execute_reply":"2025-09-29T10:24:39.720239Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Visualize data\n\ntrain_ds_vis = train_ds.unbatch().shuffle(2048,seed=43)\nnum_examples = 36\nnum_columns = 6\nnum_rows = math.ceil(num_examples/num_columns)\nplt.figure(figsize=(16, 16))\nfor i, (image, label) in enumerate(train_ds_vis.take(num_examples)):\n    if i == -1: # Set to 0 in case of interest\n        print(image.shape)\n        print('class id: '+str(label.numpy()))\n        print('class name: '+str(CLASSES[label.numpy()]))\n    class_id = str(label.numpy())\n    class_name = str(CLASSES[label.numpy()])\n    plt.subplot(num_rows, num_columns, i + 1)\n    imshow_title(image, title=f\"{class_name}({class_id})\")\n    plt.suptitle(\"Examples from train dataset\")\n    plt.xticks([])\n    plt.yticks([])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-29T10:24:39.721926Z","iopub.execute_input":"2025-09-29T10:24:39.722204Z","iopub.status.idle":"2025-09-29T10:25:19.969970Z","shell.execute_reply.started":"2025-09-29T10:24:39.722184Z","shell.execute_reply":"2025-09-29T10:25:19.969091Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Build and explore neural network","metadata":{}},{"cell_type":"code","source":"## Build network\n\n# Network configuration\nbase_network_type = 6 # Choose Pretrained Network\nbase_network = {0: \"mnv2\",  # MobileNetV2\n                1: \"vgg16\", # VGG16\n                2: \"vgg19\", # VGG19\n                3: \"rnv2\",  # ResNet50V2\n                4: \"enb0\",  # EfficientNetB0\n                5: \"enb0_v2\", # EfficientNetV2B0\n                6: \"enb6\"}  # EfficientNetB6\nTRAIN_BASE_LAYERS = True\n\n# Custom class for preprocessing layer\nclass PreProcess(tf.keras.layers.Layer):\n    def __init__(self, base_network_type, **kwargs):\n        super(PreProcess, self).__init__(**kwargs)\n        if base_network_type == 0: self.preprocess_input = tf.keras.applications.mobilenet_v2.preprocess_input\n        elif base_network_type == 1: self.preprocess_input = tf.keras.applications.vgg16.preprocess_input\n        elif base_network_type == 2: self.preprocess_input = tf.keras.applications.vgg19.preprocess_input\n        elif base_network_type == 3: self.preprocess_input = tf.keras.applications.resnet_v2.preprocess_input\n        elif base_network_type == 4: self.preprocess_input = tf.keras.applications.efficientnet.preprocess_input\n        elif base_network_type == 5: self.preprocess_input = tf.keras.applications.efficientnet_v2.preprocess_input\n        elif base_network_type == 6: self.preprocess_input = tf.keras.applications.efficientnet.preprocess_input\n        else: print('Wrong base network number have been choosen!!!')\n    def call(self, inputs):\n        return self.preprocess_input(inputs*255.0)\n\n# Build function for NN\ndef build_network(hp):\n    # Select base model and corresponding preprocessing\n    if base_network_type == 0: base_model = tf.keras.applications.MobileNetV2(include_top=False, input_shape=[*IMAGE_SIZE, 3], weights='imagenet')\n    elif base_network_type == 1: base_model = tf.keras.applications.VGG16(include_top=False, input_shape=[*IMAGE_SIZE, 3], weights='imagenet')\n    elif base_network_type == 2: base_model = tf.keras.applications.VGG19(include_top=False, input_shape=[*IMAGE_SIZE, 3], weights='imagenet')\n    elif base_network_type == 3: base_model = tf.keras.applications.ResNet50V2(include_top=False, input_shape=[*IMAGE_SIZE, 3], weights='imagenet')\n    elif base_network_type == 4: base_model = tf.keras.applications.EfficientNetB0(include_top=False, input_shape=[*IMAGE_SIZE, 3], weights='imagenet')\n    elif base_network_type == 5: base_model = tf.keras.applications.EfficientNetV2B0(include_top=False, input_shape=[*IMAGE_SIZE, 3], weights='imagenet')\n    elif base_network_type == 6: base_model = tf.keras.applications.EfficientNetB6(include_top=False, input_shape=[*IMAGE_SIZE, 3], weights='imagenet')\n    else: print('Wrong base network number have been choosen!!!')\n    prepocessing = PreProcess(base_network_type=base_network_type, name='preprocessing')  \n\n    # Choose base model layers to be trained\n    base_model.trainable = False # freeze base model layers\n    #layer_id = 118 # layer number from the network shall be trained (B0): all:238 / 2ab+:220 / 3ab+:191 / 4abc+:162 / 5abc+:118 / 6abcd+:75 / 7a+:16\n    layer_id = 328 # layer number from the network shall be trained (B6): all:667 / 2ab+:625 / 3ab+:536 / 4abc+:447 / 5abc+:328 / 6abcd+:210 / 7a+:46\n    print('Unfreeze base model layers from layer ' + str(base_model.layers[-layer_id]))\n    for layer in base_model.layers[-layer_id:]: # unfreeze choosen layers \n        if TRAIN_BASE_LAYERS:\n            layer.trainable = True\n    \n    # # Explore base model\n    # base_model.summary()\n    # tf.keras.utils.plot_model(base_model, to_file='basemodel_architecture.png', show_shapes=True, show_dtype=False,\n    #                       show_layer_names=True, show_layer_activations=True, show_trainable=False)\n    \n    # Define pretrained layers\n    raw_input = tf.keras.Input(shape=[*IMAGE_SIZE, 3], name='input')\n    x = prepocessing(raw_input)\n    x = base_model(x)\n\n    # Define final layers\n    reg = None #L2(1e-4)\n    do = 0.3\n    do_final = 0.3\n    layers_final = 0\n    units_final = 512\n    x = tf.keras.layers.GlobalAveragePooling2D(name=f'gap2d')(x)\n    x = tf.keras.layers.BatchNormalization(name=f'bn')(x)\n    x = tf.keras.layers.Dropout(do, name=f'do')(x)\n    for i, units in enumerate([units_final for j in range(layers_final)]):\n        x = tf.keras.layers.Dense(units, activation=\"relu\", kernel_regularizer=reg,\n                                  kernel_initializer='he_normal', name=f'final_fc{i+1}')(x)\n        x = tf.keras.layers.BatchNormalization(name=f'final_bn{i+1}')(x)\n        x = tf.keras.layers.Dropout(do_final, name=f'final_do{i+1}')(x)\n    outputs = tf.keras.layers.Dense(name='prediction', units=len(CLASSES), activation='softmax')(x)\n\n    # Define model\n    model = tf.keras.Model(inputs=raw_input, outputs=outputs, name='classification_model')\n\n    # Define optimizer/loss and compile model\n    lr_tune = hp.Float(name='learning_rate', min_value=5e-5, max_value=5e-3, sampling='log', default=5e-4)\n    optimizer = tf.keras.optimizers.Nadam(lr_tune)\n    loss_fn = tf.keras.losses.sparse_categorical_crossentropy\n    model.compile(optimizer=optimizer, loss=loss_fn, metrics=['accuracy'])\n    return model\n\n# Build model with choosen distribution strategy\nwith strategy.scope():\n    model = build_network(kt.HyperParameters())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-29T10:25:19.970935Z","iopub.execute_input":"2025-09-29T10:25:19.971192Z","iopub.status.idle":"2025-09-29T10:25:25.431013Z","shell.execute_reply.started":"2025-09-29T10:25:19.971174Z","shell.execute_reply":"2025-09-29T10:25:25.429999Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Explore model architecture\n\nmodel.summary()\n# tf.keras.utils.plot_model(model, to_file='model_architecture.png', show_shapes=True, show_dtype=False,\n#                           show_layer_names=True, show_layer_activations=True, show_trainable=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T18:15:45.046451Z","iopub.execute_input":"2025-09-25T18:15:45.047713Z","iopub.status.idle":"2025-09-25T18:15:45.086950Z","shell.execute_reply.started":"2025-09-25T18:15:45.047690Z","shell.execute_reply":"2025-09-25T18:15:45.082487Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Training","metadata":{}},{"cell_type":"code","source":"## Training parameters\n\nepochs = 200\nTUNING = False\nTRAINING = False\nFINETUNING = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T18:15:45.088406Z","iopub.execute_input":"2025-09-25T18:15:45.088574Z","iopub.status.idle":"2025-09-25T18:15:45.095755Z","shell.execute_reply.started":"2025-09-25T18:15:45.088558Z","shell.execute_reply":"2025-09-25T18:15:45.091827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Tuner configurations\n\ni_TunerTyp = 1 # Choose desired tuner type: {1: 'grid', 2: 'random', 3: 'hyper'}\nTunerStr = {1: 'grid', 2: 'random', 3: 'hyper'}\n\nif TUNING:\n    tuner_grid = kt.GridSearch(hypermodel=build_network, objective='val_accuracy', max_trials=10,\n                               max_consecutive_failed_trials=1, overwrite=True, directory=\"tuner\",\n                               project_name=\"Flower\", distribution_strategy = strategy)\n\n    tuner_random = kt.RandomSearch(hypermodel=build_network, objective='val_accuracy', max_trials=10,\n                                   executions_per_trial=1, overwrite=True, directory=\"tuner\",\n                                   project_name=\"Flower\", distribution_strategy = strategy)\n\n    tuner_hyper = kt.Hyperband(hypermodel=build_network, objective='val_accuracy', max_epochs=90, factor=5,\n                               hyperband_iterations=1, overwrite=True, directory=\"tuner\",\n                               project_name=\"Flower\", distribution_strategy = strategy)\n\n    tuner = globals()[f'tuner_{TunerStr[i_TunerTyp]}']\n    tuner.search_space_summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T18:21:07.837746Z","iopub.execute_input":"2025-09-25T18:21:07.838075Z","iopub.status.idle":"2025-09-25T18:21:07.849712Z","shell.execute_reply.started":"2025-09-25T18:21:07.838049Z","shell.execute_reply":"2025-09-25T18:21:07.844813Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Train or tune model\n\n# Callback functions\nlr_scheduler = tf.keras.callbacks.ReduceLROnPlateau(factor=0.2, patience=5, verbose=1, monitor='val_accuracy')\nearly_stopping_cb = tf.keras.callbacks.EarlyStopping(patience=10, verbose=1, monitor='val_accuracy', restore_best_weights=True)\nlr_schedule = tf.keras.callbacks.LearningRateScheduler(lambda epoch: 1e-5 * 10**(epoch / 20)) # Set the learning rate scheduler\n\nif TRAINING:\n    history = model.fit(train_ds, validation_data=val_ds, epochs=epochs, callbacks=[lr_scheduler, early_stopping_cb])\n\nif TUNING:\n    tuner.search(train_ds, validation_data=val_ds, epochs=epochs, callbacks=[lr_scheduler, early_stopping_cb])\n    best_models = tuner.get_best_models(num_models=2)\n    model = best_models[0]\n    model.summary()\n    tuner.results_summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T18:21:13.593045Z","iopub.execute_input":"2025-09-25T18:21:13.593291Z","execution_failed":"2025-09-25T19:01:16.586Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Save weights of model\n\nif TRAINING or TUNING:\n    model.save_weights('flower_' + base_network[base_network_type] +'_1_3_0.weights.h5')\n    print('Model weights have been saved!')","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-25T19:01:16.587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Load pre-trained weights\n\nif not TRAINING and not TUNING:\n    model.load_weights('/kaggle/input/flower-1xx/flower_enb6_1_3_0.weights.h5')\n    print('Model weights have been loaded!')","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-25T19:01:16.587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Finetuning of network on deeper layers (final + selected layers)\n\nif FINETUNING:\n    # Re-build model with updated configuration (TRAIN_BASE_LAYERS) and load trained final layer weight\n    TRAIN_BASE_LAYERS = True\n    with strategy.scope():\n        model = build_network(kt.HyperParameters())\n    model.load_weights('flower_' + base_network[base_network_type] +'_1_2_0.weights.h5')\n\n    # Fine tune deeper layers\n    history = model.fit(train_ds, validation_data=val_ds, epochs=epochs, callbacks=[lr_scheduler, early_stopping_cb])\n\n    # Save weights of model\n    model.save_weights('flower_' + base_network[base_network_type] +'_1_2_0.weights.h5')\n    print('Model weights have been saved!')","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-25T19:01:16.587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Plot learning curves\n\nif TRAINING or FINETUNING:\n    history_fil = {key: history.history[key] for key in ['accuracy', 'val_accuracy']}\n    history_fil2 = {key: history.history[key] for key in ['loss', 'val_loss']}\n    history_fil3 = {key: history.history[key] for key in ['learning_rate']}\n    \n    pd.DataFrame(history_fil).plot()\n    plt.ylabel(\"Accuracy\")\n    plt.xlabel(\"epochs\")\n    plt.axis([2, len(history_fil['val_accuracy']), history_fil['val_accuracy'][2], 1])\n    pd.DataFrame(history_fil2).plot()\n    plt.ylabel(\"Loss\")\n    plt.xlabel(\"epochs\")\n    plt.axis([2, len(history_fil2['val_loss']), 0, history_fil2['val_loss'][2]+0.1*history_fil2['val_loss'][2]])\n    pd.DataFrame(history_fil3).plot()\n    plt.ylabel(\"Learning rate\")\n    plt.xlabel(\"epochs\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 6. Evaluation","metadata":{}},{"cell_type":"code","source":"## Visualize predictions for val dataset\n\nval_ds_vis = val_ds.unbatch().shuffle(2048,seed=43)\nnum_examples = 36\nnum_columns = 6\nnum_rows = math.ceil(num_examples/num_columns)\nplt.figure(figsize=(16, 16))\nfor i, (image, label) in enumerate(val_ds_vis.take(num_examples)):\n    if i == -1: # Set to 0 in case of interest\n        print(image.shape)\n        print('class id: '+str(label.numpy()))\n        print('class name: '+str(CLASSES[label.numpy()]))\n    class_id = str(label.numpy())\n    class_name = str(CLASSES[label.numpy()])\n    image_ext = tf.expand_dims(image,axis=0)\n    probability = model.predict(image_ext, verbose=0)\n    prediction = probability.argmax(axis=1)\n    id_prediction = str(prediction[0])\n    class_prediction = str(CLASSES[prediction[0]])\n    plt.subplot(num_rows, num_columns, i + 1)\n    imshow_title(image, f\"{class_prediction}({class_name})\", 'red' if class_prediction != class_name else 'black')\n    plt.suptitle(\"Examples from val dataset\")\n    plt.xticks([])\n    plt.yticks([])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 7. Submission","metadata":{}},{"cell_type":"code","source":"## Make predictions and create submission file\n\ndata_ids = []\nfor i, data_id in enumerate(test_ids_ds):\n    data_ids.append(data_id.numpy().astype('U'))\n\nsubmission = pd.DataFrame(data_ids, columns=['id'])\nprobabilities = model.predict(test_images_ds)\npredictions = np.argmax(probabilities, axis=-1)\ntest_prediction = pd.DataFrame(predictions, columns=['label'])\nsubmission = submission.merge(test_prediction, how='left', left_index=True, right_index=True)\nsubmission.to_csv('submission.csv', index=False)\n\n# Look at the first few predictions\n!head submission.csv","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 8. Experimental code","metadata":{}},{"cell_type":"code","source":"# ##  Plot learning curves for definition of start learning rate\n\n# lrs = 1e-5 * (10 ** (np.arange(len(history.history[\"loss\"])) / 20)) # Define the learning rate array\n# plt.figure(figsize=(10, 6)) # Set the figure size\n# plt.grid(True) # Set the grid\n# plt.semilogx(lrs, history.history[\"loss\"]) # Plot the loss in log scale\n# plt.tick_params('both', length=10, width=1, which='both') # Increase the tickmarks size\n# plt.axis([1e-5, 1e-0, 0, 5]) # Set the plot boundaries","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}