{"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":"gpu","dataSources":[{"sourceId":25954,"databundleVersionId":2091745,"sourceType":"competition"},{"sourceId":33246,"databundleVersionId":3221581,"sourceType":"competition"},{"sourceId":44224,"databundleVersionId":5188730,"sourceType":"competition"},{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":8048860,"sourceType":"datasetVersion","datasetId":4745108},{"sourceId":171058746,"sourceType":"kernelVersion"}],"dockerImageVersionId":30674,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install audiomentations -qq\nimport os\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'\nos.environ['KERAS_BACKEND'] = 'tensorflow'\n\nfrom audiomentations import Compose, AddGaussianNoise, TimeStretch, PitchShift, Shift, AirAbsorption\nfrom keras import models, layers, losses, applications, optimizers, saving, ops\nfrom sklearn.model_selection import KFold, train_test_split, StratifiedKFold\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom matplotlib import pyplot as plt\nfrom dataclasses import dataclass\nfrom IPython import display\n\nimport tensorflow as tf\nimport seaborn as sns\nimport pandas as pd\nimport numpy as np\nimport joblib\nimport os\n\nsns.set_theme(style=\"darkgrid\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-09T01:14:33.959540Z","iopub.execute_input":"2024-04-09T01:14:33.959947Z","iopub.status.idle":"2024-04-09T01:15:01.519448Z","shell.execute_reply.started":"2024-04-09T01:14:33.959916Z","shell.execute_reply":"2024-04-09T01:15:01.518022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_strategy():\n    try:\n        resolver = tf.distribute.cluster_resolver.TPUClusterResolver(tpu='local')\n        tf.config.experimental_connect_to_cluster(resolver)\n        tf.tpu.experimental.initialize_tpu_system(resolver)\n        strategy = tf.distribute.TPUStrategy(resolver)\n        device = \"TPU\"\n    except:\n        gpus = tf.config.list_physical_devices('GPU')\n        if  len(gpus) == 1:\n            strategy = tf.distribute.OneDeviceStrategy(\"/gpu:0\")\n            device = \"GPU\"\n        elif len(gpus) > 1:\n            strategy = tf.distribute.MirroredStrategy()\n            device = \"GPU\"\n        else:\n            strategy = tf.distribute.get_strategy()\n            device = \"CPU\"\n    msg = f\"Notebook running on {device} | Number of accelerators {strategy.num_replicas_in_sync}\"\n    print(msg)\n    return strategy, device\n\nstrategy, device = get_strategy()","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:01.522094Z","iopub.execute_input":"2024-04-09T01:15:01.522727Z","iopub.status.idle":"2024-04-09T01:15:02.067549Z","shell.execute_reply.started":"2024-04-09T01:15:01.522692Z","shell.execute_reply":"2024-04-09T01:15:02.066471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@dataclass\nclass Config:\n    ebird_taxo       = \"/kaggle/input/birdclef-2024/eBird_Taxonomy_v2021.csv\"\n    # Addationnal dataset\n    meta_data2024    = \"/kaggle/input/birdclef-2024/train_metadata.csv\"\n    meta_data2023    = \"/kaggle/input/birdclef-2023/train_metadata.csv\"\n    meta_data2022    = \"/kaggle/input/birdclef-2022/train_metadata.csv\"\n    meta_data2021    = \"/kaggle/input/birdclef-2021/train_metadata.csv\"\n    \n    train_dir_2024   = \"/kaggle/input/birdclef-2024/train_audio\"\n    train_dir_2023   = \"/kaggle/input/birdclef-2023/train_audio\"\n    train_dir_2022   = \"/kaggle/input/birdclef-2022/train_audio\"\n    train_dir_2021   = \"/kaggle/input/birdclef-2021/train_short_audio\"\n    \n    \n    orig_sample_rate = 32_000\n    time_mask_n      = 5\n    freq_mask_n      = 5\n    frame_secondes   = 5\n    frame_step       = (orig_sample_rate*frame_secondes)//768\n    frame_length     = frame_step + 128\n    fft_length       = 1664\n    n_mel            = 128\n    fmin             = 20.0\n    fmax             = 16_000\n    eps              = 1e-10\n    n_folds          = 5\n    shuffle          = True\n    random_state     = 1\n    mixed_prec       = device == \"CPU\" or device == \"GPU\"\n    deterministic    = True\n    \n    epochs           = 3\n    batch_per_device = 32\n    batch            = batch_per_device*strategy.num_replicas_in_sync","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:02.068634Z","iopub.execute_input":"2024-04-09T01:15:02.068897Z","iopub.status.idle":"2024-04-09T01:15:02.077461Z","shell.execute_reply.started":"2024-04-09T01:15:02.068874Z","shell.execute_reply":"2024-04-09T01:15:02.076355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if Config.mixed_prec:\n    #tf.config.optimizer.set_experimental_options({\"auto_mixed_precision\": True})\n    tf.keras.mixed_precision.set_global_policy(\"mixed_float16\")\n    print(\"Using Mixed Precision !!!\")\nelse:\n    print(\"Using Full Precision\")","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:02.079909Z","iopub.execute_input":"2024-04-09T01:15:02.080226Z","iopub.status.idle":"2024-04-09T01:15:02.093228Z","shell.execute_reply.started":"2024-04-09T01:15:02.080201Z","shell.execute_reply":"2024-04-09T01:15:02.092380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_random_seed(seed: int = 42, deterministic: bool = False):\n    import random\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    tf.random.set_seed(seed)\n    if deterministic:\n        os.environ['TF_DETERMINISTIC_OPS'] = '1'\n    else:\n        os.environ.pop('TF_DETERMINISTIC_OPS', None)\n\n# Set a deterministic behavior\nset_random_seed(seed=Config.random_state, deterministic=Config.deterministic)","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:02.094263Z","iopub.execute_input":"2024-04-09T01:15:02.094544Z","iopub.status.idle":"2024-04-09T01:15:02.104847Z","shell.execute_reply.started":"2024-04-09T01:15:02.094522Z","shell.execute_reply":"2024-04-09T01:15:02.103941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df    = pd.read_csv(\"/kaggle/input/count-time-of-all-dataset/train_202122024_with_add.csv\")\nkfold = StratifiedKFold(n_splits=Config.n_folds, shuffle=Config.shuffle, random_state=Config.random_state)\n\nfor idx, (_, index) in enumerate(kfold.split(df, df.primary_label)):\n    df.loc[index, \"Fold\"] = idx\n    \ndf['Fold']       = df['Fold'].astype(\"int32\")\ndf['num_sample'] = df['duration'].apply(lambda x: max(x/Config.frame_secondes, 1))\ndf['num_sample'] = df['num_sample'].astype('int32')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:02.106031Z","iopub.execute_input":"2024-04-09T01:15:02.106419Z","iopub.status.idle":"2024-04-09T01:15:02.501887Z","shell.execute_reply.started":"2024-04-09T01:15:02.106386Z","shell.execute_reply":"2024-04-09T01:15:02.500865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"audio_augment = Compose([\n    AddGaussianNoise(min_amplitude=0.001, max_amplitude=0.015, p=0.5),\n    #TimeStretch(min_rate=0.8, max_rate=1.25, p=0.5),\n    #PitchShift(min_semitones=-4, max_semitones=4, p=1.0),\n    Shift(p=1.0),\n    AirAbsorption(min_distance=10.0,max_distance=50.0,p=1.0)\n])","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:02.503218Z","iopub.execute_input":"2024-04-09T01:15:02.503620Z","iopub.status.idle":"2024-04-09T01:15:02.510940Z","shell.execute_reply.started":"2024-04-09T01:15:02.503587Z","shell.execute_reply":"2024-04-09T01:15:02.509370Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import soundfile as sf\nlabel_encoder   = layers.StringLookup(num_oov_indices=0, output_mode=\"one_hot\", vocabulary=df[\"primary_label\"].unique())\n\ndef show_dataset(ds, sample=8, figsize=(15, 15), dpi=100):\n    if isinstance(ds.element_spec[0], dict):\n        n_cols = 2\n    else:\n        n_cols = 1\n    \n    fig = plt.figure(figsize=figsize, dpi=dpi)\n    for idx, (data, label) in zip(range(1, sample*n_cols+1, n_cols), ds.unbatch().take(sample)):\n        if isinstance(ds.element_spec[0], dict):\n            ax = fig.add_subplot(sample, n_cols, idx)\n            ax.imshow(data[\"spec\"], aspect=\"auto\", cmap='viridis')\n            ax.set_axis_off()            \n            ax.set_title(f\"Spec[label={label_encoder.get_vocabulary()[label.numpy().argmax()]}]\", fontsize=12)\n            ax = fig.add_subplot(sample, n_cols, idx+1)\n            ax.plot(data[\"audio\"])\n            ax.set_title(f\"Waveform[label={label_encoder.get_vocabulary()[label.numpy().argmax()]}]\")\n        else:\n            ax = fig.add_subplot(sample, n_cols, idx)\n            ax.plot(data)\n            ax.set_title(f\"Waveform[label={label_encoder.get_vocabulary()[label.numpy().argmax()]}]\")  \n    plt.tight_layout()\n    plt.show()\n    return None\n\ndef show_tf_dataset(ds, sample=8, figsize=(15, 15), dpi=100, process_layer=None):\n    def wrapper(audio, label):\n        spec = process_layer(audio)\n        data = {\n            \"audio\": audio,\n            \"spec\" : spec[..., 0]\n        }\n        return data, label\n    \n    if process_layer:\n        ds = ds.map(wrapper, tf.data.AUTOTUNE).prefetch(tf.data.AUTOTUNE)\n    return show_dataset(ds)\n\ndef random_crop_audio(audio, crop_size, seed=None):\n    audio_length = tf.shape(audio)[0]    \n    # Make sure crop_size is valid\n    if crop_size > audio_length:\n        audio = tf.pad(audio, paddings=[(0, crop_size-audio_length)])\n        return audio\n    random_offset = tf.random.uniform(\n          shape=(), minval=0, maxval=audio_length - crop_size + 1, dtype=tf.int32, seed=seed\n        )\n    cropped_audio = tf.slice(audio, begin=[random_offset], size=[crop_size])    \n    return cropped_audio\n\n@tf.function\n@tf.numpy_function(Tout=tf.float32)\ndef read_audio_tf(path):\n    if isinstance(path, bytes):\n        path = path.decode(\"utf-8\")\n    \n    audio, _ = sf.read(path)\n    if audio.ndim == 1:\n        audio = np.expand_dims(audio, axis=-1)\n    audio = audio.astype(\"float32\")\n    return audio\n\n@tf.numpy_function(Tout=tf.float32)\ndef transform(audio):\n    audio = audio_augment(samples=audio, sample_rate=Config.orig_sample_rate)\n    return audio\n\ndef audio_aug_fn(audio, label=None):\n    audio = tf.map_fn(transform, audio)\n    audio = tf.ensure_shape(audio, (audio.shape[0], Config.frame_secondes*Config.orig_sample_rate))\n    if not label is None:\n        return audio, label\n    return audio\n\ndef prepare_single_audio(path, label):\n    audio = read_audio_tf(path)\n    audio = tf.gather(audio, [0], axis=-1)\n    \n    if tf.shape(audio)[0] < Config.frame_secondes*Config.orig_sample_rate:\n        audio = tf.pad(audio, [(0, (Config.frame_secondes*Config.orig_sample_rate)-tf.shape(audio)[0]), (0, 0)])\n        \n    frames = tf.signal.frame(\n        audio, \n        frame_length=(Config.frame_secondes*Config.orig_sample_rate), \n        frame_step=(Config.frame_secondes*Config.orig_sample_rate),\n        axis=0,\n    )\n    frames = tf.ensure_shape(frames, (None, Config.frame_secondes*Config.orig_sample_rate, 1))\n    frames = tf.squeeze(frames, axis=-1)\n    labels = tf.repeat(label, tf.shape(frames)[0], axis=0)\n    labels = label_encoder(labels)\n    \n    return tf.data.Dataset.from_tensor_slices((frames, labels))\n\ndef create_dataset_from_pandas(dataframe, batch=16, shuffle=False, cache=False, is_train=False, cache_dir=\"\"):\n    dataframe = dataframe.sample(frac=1.0).copy()\n    auto      = tf.data.AUTOTUNE\n    opt       = tf.data.Options()\n    opt.experimental_deterministic = False\n    prefetch_multi = 64 if device == \"tpu\".upper() else 16\n    shuffle_multi = 8 if device == \"tpu\".upper() else 8\n    \n    ds = tf.data.Dataset.from_tensor_slices((dataframe['path'].values, dataframe[\"primary_label\"].values))\n    ds = ds.interleave(prepare_single_audio, cycle_length=auto, num_parallel_calls=auto)\n    ds = ds.with_options(opt)\n    ds = ds.prefetch(buffer_size=prefetch_multi*batch)\n    \n    if not cache_dir == \"\":\n        os.makedirs(cache_dir)\n        \n    ds = ds.cache(cache_dir) if cache else ds\n    ds = ds.repeat()\n    ds = ds.shuffle(shuffle_multi*batch)\n    ds = ds.batch(batch, drop_remainder=True)\n    ds = ds.map(audio_aug_fn, auto) if is_train else ds\n    return ds.prefetch(auto)","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:02.512395Z","iopub.execute_input":"2024-04-09T01:15:02.512759Z","iopub.status.idle":"2024-04-09T01:15:02.897143Z","shell.execute_reply.started":"2024-04-09T01:15:02.512726Z","shell.execute_reply":"2024-04-09T01:15:02.896351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tds = create_dataset_from_pandas(df, batch=8, is_train=True)\nvds = create_dataset_from_pandas(df, batch=8, is_train=False)","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:02.898455Z","iopub.execute_input":"2024-04-09T01:15:02.898801Z","iopub.status.idle":"2024-04-09T01:15:05.287737Z","shell.execute_reply.started":"2024-04-09T01:15:02.898771Z","shell.execute_reply":"2024-04-09T01:15:05.286875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tds","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:05.290646Z","iopub.execute_input":"2024-04-09T01:15:05.290910Z","iopub.status.idle":"2024-04-09T01:15:05.297139Z","shell.execute_reply.started":"2024-04-09T01:15:05.290888Z","shell.execute_reply":"2024-04-09T01:15:05.296275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vds","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:05.298165Z","iopub.execute_input":"2024-04-09T01:15:05.298447Z","iopub.status.idle":"2024-04-09T01:15:05.309371Z","shell.execute_reply.started":"2024-04-09T01:15:05.298424Z","shell.execute_reply":"2024-04-09T01:15:05.308450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\naudio, label = next(iter(tds))","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:05.310628Z","iopub.execute_input":"2024-04-09T01:15:05.310906Z","iopub.status.idle":"2024-04-09T01:15:16.087766Z","shell.execute_reply.started":"2024-04-09T01:15:05.310879Z","shell.execute_reply":"2024-04-09T01:15:16.086813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\naudio, label = next(iter(vds))","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:16.088882Z","iopub.execute_input":"2024-04-09T01:15:16.089399Z","iopub.status.idle":"2024-04-09T01:15:17.093989Z","shell.execute_reply.started":"2024-04-09T01:15:16.089373Z","shell.execute_reply":"2024-04-09T01:15:17.093033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display.Audio(np.asarray(audio[0]), rate=Config.orig_sample_rate)","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:17.095278Z","iopub.execute_input":"2024-04-09T01:15:17.096563Z","iopub.status.idle":"2024-04-09T01:15:17.127504Z","shell.execute_reply.started":"2024-04-09T01:15:17.096533Z","shell.execute_reply":"2024-04-09T01:15:17.126667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"saving.get_custom_objects().clear()\n\n@saving.register_keras_serializable(name=\"MelSpec\")\nclass MelSpec(layers.Layer):\n    def __init__(\n        self, \n        frame_length,\n        frame_step,\n        num_mels=80,\n        fft_length=1024,\n        sample_rate=32_000,\n        fmin=10.0,\n        fmax=15_000,\n        transform=None,\n        epsilon=1e-5,\n        **kwargs\n        ):\n        super().__init__(**kwargs)\n        self.frame_length = frame_length\n        self.frame_step   = frame_step\n        self.num_mels     = num_mels\n        self.fft_length   = fft_length\n        self.sample_rate  = sample_rate\n        self.fmin         = fmin\n        self.fmax         = fmax\n        self.transform    = transform\n        self.epsilon      = epsilon\n        self.built        = True\n    \n    def get_config(self):\n        config = super().get_config()\n        config.update({\n            \"frame_step\"   : self.frame_step,\n            \"num_mels\"     : self.num_mels,\n            \"fft_length\"   : self.fft_length,\n            \"frame_length\" : self.frame_length,\n            \"sample_rate\"  : self.sample_rate,\n            \"fmin\"         : self.fmin,\n            \"fmax\"         : self.fmax,\n            \"epsilon\"      : self.epsilon\n        })\n        return config\n    \n    def compute_output_shape(self, input_shape):\n        if input_shape[1]:\n            dim = input_shape[1]//self.frame_step+1\n        else:\n            dim = None\n        return (None, self.num_mels, dim, 3)\n\n    def _get_melspectogram(self, waveform):\n        # Convert audio waveform to spectrogram\n        stfts = tf.signal.stft(\n            waveform, frame_length=self.frame_length, frame_step=self.frame_step, fft_length=self.fft_length, pad_end=True\n        )\n        spectrograms = tf.square(tf.abs(stfts))\n        # Create mel filter bank\n        num_spectrogram_bins = tf.shape(stfts)[-1]\n        linear_to_mel_weight_matrix = tf.signal.linear_to_mel_weight_matrix(\n            self.num_mels, num_spectrogram_bins, self.sample_rate, lower_edge_hertz=self.fmin, upper_edge_hertz=self.fmax\n        )\n        # Apply mel filter bank to spectrogram\n        mel_spectrogram = tf.tensordot(spectrograms, linear_to_mel_weight_matrix, 1)\n        mel_spectrogram.set_shape(spectrograms.shape[:-1].concatenate(linear_to_mel_weight_matrix.shape[-1:]))\n        return mel_spectrogram\n    \n    def call(self, inputs, **kwargs):\n        inputs = tf.cast(inputs, \"float32\")\n        x      = self._get_melspectogram(inputs)\n        if self.transform:\n            x = self.transform(x)\n        x      = tf.math.log(x + self.epsilon)\n        x      = tf.transpose(x, (0, 2, 1))\n        mean   = tf.reduce_mean(x, axis=(1, 2), keepdims=True)\n        stddev = tf.math.reduce_std(x, axis=(1, 2), keepdims=True)\n        x     -= mean\n        x     /= (stddev + self.epsilon)\n        x      = tf.expand_dims(x, axis=-1)\n        x      = tf.tile(x, [1, 1, 1, 3])\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:17.128831Z","iopub.execute_input":"2024-04-09T01:15:17.129149Z","iopub.status.idle":"2024-04-09T01:15:17.149948Z","shell.execute_reply.started":"2024-04-09T01:15:17.129121Z","shell.execute_reply":"2024-04-09T01:15:17.148883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def MyLossFunction(y_true, y_pred, sample_weight=None):\n    lossv1 = tf.nn.compute_average_loss(\n        losses.categorical_crossentropy(y_true, y_pred, from_logits=True, label_smoothing=0.01),\n        sample_weight=sample_weight\n    )\n    lossv2 = tf.nn.compute_average_loss(\n        losses.kl_divergence(y_true, tf.nn.softmax(y_pred, axis=-1)),\n        sample_weight=sample_weight\n    )\n    loss = 0.2*lossv1 + 0.8*lossv2\n    return loss","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:17.151017Z","iopub.execute_input":"2024-04-09T01:15:17.151315Z","iopub.status.idle":"2024-04-09T01:15:17.162989Z","shell.execute_reply.started":"2024-04-09T01:15:17.151272Z","shell.execute_reply":"2024-04-09T01:15:17.162206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_model(input_shape=((Config.orig_sample_rate*Config.frame_secondes),), pretrained=False):\n    inputs = layers.Input(shape=input_shape, name=\"audio\")\n    \n    x = MelSpec(frame_length=Config.frame_length, frame_step=Config.frame_step, fft_length=Config.fft_length,\n                num_mels=Config.n_mel, fmin=Config.fmin, fmax=Config.fmax, epsilon=1e-10, name=\"spec\")(inputs)\n    base_model = applications.ResNet50V2(\n        include_top=False, weights=\"imagenet\" if pretrained else None, \n        input_shape=(Config.n_mel, input_shape[0]//Config.frame_step + 1, 3)\n    )\n    x = base_model(x, training=False)\n    x = layers.GlobalAveragePooling2D(name=\"pooling\")(x)\n    outputs = layers.Dense(label_encoder.vocabulary_size(), name=\"targets\")(x)\n    model   = models.Model(inputs, outputs, name=\"BirdClassifier\")\n    #loss_fn = losses.CategoricalCrossentropy(from_logits=False)\n    opt = optimizers.Adam(5e-5)\n    model.compile(\n        loss=MyLossFunction,\n        optimizer=opt,\n        metrics=[tf.keras.metrics.CategoricalAccuracy(name='accuracy'), tf.keras.metrics.AUC(from_logits=True)]\n    )\n    return model\n\ntf.keras.backend.clear_session()\nwith strategy.scope():\n    model = create_model(pretrained=True)\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2024-04-09T01:15:17.164064Z","iopub.execute_input":"2024-04-09T01:15:17.164385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fold = 0\nvalid_ds = df[df[\"Fold\"] == fold]\ntrain_ds = df[df[\"Fold\"] != fold]\nprint(f\"Train size {len(train_ds)}, validation size {len(valid_ds)}\")\n\ntrain_step = train_ds['num_sample'].sum()//Config.batch + 1\nvalid_step = valid_ds['num_sample'].sum()//Config.batch + 1\n\nprint(\"Train step\", train_step)\nprint(\"Valid step\", valid_step)\n\ntrain_ds = create_dataset_from_pandas(train_ds, batch=Config.batch, shuffle=True, is_train=True)\nvalid_ds = create_dataset_from_pandas(valid_ds, batch=Config.batch, shuffle=False, is_train=False)\n\nshow_tf_dataset(train_ds, process_layer=model.get_layer(\"spec\"))\nshow_tf_dataset(valid_ds, process_layer=model.get_layer(\"spec\"))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    import wandb\nexcept:\n    !pip install wandb -qq\n    import wandb\n    \nfrom kaggle_secrets import UserSecretsClient\nuser_secrets  = UserSecretsClient()\nwanbd_api_key = user_secrets.get_secret(\"WANDB_API_KEY\")\nwandb.login(key=wanbd_api_key)\n\nfrom datetime import datetime\nwandb.init(project=\"CNN[Train]\", name=f\"Run-time: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\", \n           config={name:getattr(Config, name) for name in dir(Config) if not \"_\" in name})","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nMC = tf.keras.callbacks.ModelCheckpoint(\n    \"Checpoint/resenet50_fold_0.weights.h5\", monitor=\"val_loss\", save_best_only=True, save_weights_only=True\n)\ndef get_lr_callback(batch_size=8, mode='cos', epochs=10, plot=False):\n    lr_start, lr_max, lr_min = 1.0e-7, 0.65e-6 * batch_size, 0.3e-6\n    lr_ramp_ep, lr_sus_ep, lr_decay = 0, 0, 0.75\n\n    def lrfn(epoch):  # Learning rate update function\n        if epoch < lr_ramp_ep: lr = (lr_max - lr_start) / lr_ramp_ep * epoch + lr_start\n        elif epoch < lr_ramp_ep + lr_sus_ep: lr = lr_max\n        elif mode == 'exp': lr = (lr_max - lr_min) * lr_decay**(epoch - lr_ramp_ep - lr_sus_ep) + lr_min\n        elif mode == 'step': lr = lr_max * lr_decay**((epoch - lr_ramp_ep - lr_sus_ep) // 2)\n        elif mode == 'cos':\n            decay_total_epochs, decay_epoch_index = epochs - lr_ramp_ep - lr_sus_ep + 3, epoch - lr_ramp_ep - lr_sus_ep\n            phase = math.pi * decay_epoch_index / decay_total_epochs\n            lr = (lr_max - lr_min) * 0.5 * (1 + math.cos(phase)) + lr_min\n        return lr\n\n    if plot:  # Plot lr curve if plot is True\n        plt.figure(figsize=(15, 5), dpi=100)\n        plt.plot(np.arange(epochs), [lrfn(epoch) for epoch in np.arange(epochs)], marker='o')\n        plt.xlabel('epoch'); plt.ylabel('lr')\n        plt.title('LR Scheduler')\n        plt.show()\n\n    return tf.keras.callbacks.LearningRateScheduler(lrfn, verbose=False)\n\ncallbacks = [\n    MC,\n    tf.keras.callbacks.TerminateOnNaN(),\n    tf.keras.callbacks.EarlyStopping(monitor=\"val_loss\", verbose=True, patience=3),\n    tf.keras.callbacks.ReduceLROnPlateau(patience=1, factor=0.80),\n    wandb.keras.WandbMetricsLogger(log_freq=\"batch\")\n]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.fit(train_ds, validation_data=valid_ds, epochs=Config.epochs, steps_per_epoch=train_step, validation_steps=valid_step,\n          callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}