{"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":"none","dataSources":[{"sourceId":7046335,"sourceType":"datasetVersion","datasetId":3628570}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# SWAV\n\nOFFICAL CODE: https://github.com/ayulockin/SwAV-TF\n\ndata prep: https://colab.research.google.com/drive/1pGQW7OnxxNaV23-88s6kgBLx_3Q7H5ud?usp=sharing","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\nimport tensorflow as tf\nfrom tensorflow.keras import layers\nfrom tensorflow.keras import models\nimport tensorflow_datasets as tfds\n\nfrom tqdm import trange\nfrom itertools import groupby\n\nimport random\nimport os\nimport gc\n\nimport statsmodels.api as sm\nimport matplotlib.cm as cm\nimport numpy as np\nimport scipy\n\nplot_sample_flag = 0 #qq and dist plots\nplot_all_flag = 0 #cluster assigments plots\ntrain_flag=1\n\nprototypes = 15\n\n#to not use proportions declare it as empty list\n# y_prop = [0.29546847,0.08448649,0.17759459,0.09261261,0.01972072,0.10914414,0.11337838,0.06413514,0.04345946,]\ny_prop = []\n\nif len(y_prop):\n    prototypes = len(y_prop)\n\nmodel_tag  = ''\n\ntf.random.set_seed(0)\nnp.random.seed(0)\ntfds.disable_progress_bar()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-11-28T15:36:21.310430Z","iopub.execute_input":"2023-11-28T15:36:21.311287Z","iopub.status.idle":"2023-11-28T15:36:21.329011Z","shell.execute_reply.started":"2023-11-28T15:36:21.311252Z","shell.execute_reply":"2023-11-28T15:36:21.327933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"options = tf.data.Options()\noptions.experimental_optimization.noop_elimination = True             # eliminate no-op transformations\ntf.compat.v1.data.experimental.OptimizationOptions.map_vectorization = True    # vectorize map transformations\noptions.experimental_optimization.apply_default_optimizations = True  # apply default graph optimizations\noptions.experimental_deterministic = True                            # False disable deterministic order\noptions.threading.max_intra_op_parallelism = 1           # overrides the maximum degree of intra-op parallelism\n\nBS = 256\nSIZE_CROPS = [50, 32]  #RESCALED IMAGE \nNUM_CROPS= [2,3]\n\n#SCALE FOR CROPPING\nMIN_SCALE = [0.7, 0.5]\nMAX_SCALE = [1., 0.7]","metadata":{"execution":{"iopub.status.busy":"2023-11-28T15:36:21.331259Z","iopub.execute_input":"2023-11-28T15:36:21.331633Z","iopub.status.idle":"2023-11-28T15:36:21.343497Z","shell.execute_reply.started":"2023-11-28T15:36:21.331604Z","shell.execute_reply":"2023-11-28T15:36:21.342632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"*normalized* =  min max normalization\n\ncomplete_data_half: 1988 to 2023 (complete) data with step* 2 (half)\n\ncomplete_data_limited: 1988 to 2023 (complete) data and image sum >= 0.1 (limited)\n\ncomplete_data_limited_half: 1988 to 2023 (complete) data with step* 2 (half) and image sum >= 0.1 (limited)\n\n*each step is a 3 hours block","metadata":{}},{"cell_type":"code","source":"# import shutil; shutil.rmtree('/kaggle/working/trainloaders_zipped')","metadata":{"execution":{"iopub.status.busy":"2023-11-28T15:36:21.344484Z","iopub.execute_input":"2023-11-28T15:36:21.344786Z","iopub.status.idle":"2023-11-28T15:36:21.354713Z","shell.execute_reply.started":"2023-11-28T15:36:21.344749Z","shell.execute_reply":"2023-11-28T15:36:21.353786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import numpy as np\n# import matplotlib.pyplot as plt\n\n# data = np.load('/kaggle/input/era5-data-smaller/era5_1940_2023.npy',mmap_mode='c')[5:-14] #data start at 1940-01-01 7h; will use since 1940-01-01 12h so its divisi\n\n# min_temp, max_temp = np.quantile(data,0.0001), np.quantile(data,0.9999)\n\n# #min max \n# norm_imgs = np.clip(data, min_temp, max_temp)\n# norm_imgs = (norm_imgs - min_temp)/(max_temp -min_temp)\n\n# views = 3\n# norm_imgs =norm_imgs.reshape(int(norm_imgs.shape[0]/views), views, 50, 67)\n# norm_imgs = np.transpose(norm_imgs, (0, 2, 3, 1))\n# norm_imgs.shape\n\n# np.save(r'/kaggle/working/normalized_precipitation_half',norm_imgs[::2],allow_pickle=False)\n\n# del data, norm_imgs\n# gc.collect()\n\n#download\n# from IPython.display import FileLink\n# FileLink(r'normalized_precipitation_half.npy')","metadata":{"execution":{"iopub.status.busy":"2023-11-28T15:36:21.356574Z","iopub.execute_input":"2023-11-28T15:36:21.356896Z","iopub.status.idle":"2023-11-28T15:36:21.367774Z","shell.execute_reply.started":"2023-11-28T15:36:21.356869Z","shell.execute_reply":"2023-11-28T15:36:21.366878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# data = np.load('/kaggle/input/era5-data-smaller/normalized_complete_data_limited_half.npy',mmap_mode='c') first tests version\ndata = np.load('/kaggle/working/normalized_precipitation_half.npy',mmap_mode='c')\n\nprint(data.shape)\nimages_timesteps = data.shape[-1]\n    \nif  plot_sample_flag:\n    test_ds = np.random.choice(data.shape[0],size=20000)\n    test_ds = data[test_ds]\n\nif not plot_all_flag  :\n    data = tf.data.Dataset.from_tensor_slices(data) \n\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-11-28T15:36:21.368905Z","iopub.execute_input":"2023-11-28T15:36:21.369163Z","iopub.status.idle":"2023-11-28T15:36:36.137160Z","shell.execute_reply.started":"2023-11-28T15:36:21.369139Z","shell.execute_reply":"2023-11-28T15:36:36.135987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train_flag == 1:\n\n    AUTO = tf.data.experimental.AUTOTUNE\n\n    @tf.function\n    def random_resize_crop(image, min_scale, max_scale, crop_size,brightness_shift=1):\n\n        # Conditional resizing\n        if crop_size >= 32:\n            image_shape = 50\n            image = tf.image.resize(image, (image_shape, image_shape))\n\n        else:\n            image_shape = 32\n            image = tf.image.resize(image, (image_shape, image_shape))\n\n        # Get the crop size for given min and max scale\n        size = tf.random.uniform(shape=(1,), minval=min_scale*image_shape,\n            maxval=max_scale*image_shape, dtype=tf.float32)\n\n        size = tf.cast(size, tf.int32)[0]\n        crop = tf.image.random_crop(image, (size, size, 3))\n\n        # Resize cropped\n        crop_resize = tf.image.resize(crop, (crop_size, crop_size))\n\n        return crop_resize\n\n\n    @tf.function\n    def random_apply(func, x, p=0.5):\n        return tf.cond(tf.less(tf.random.uniform([], minval=0, maxval=1, dtype=tf.float32),\n                               tf.cast(p, tf.float32)),lambda: func(x),lambda: x)\n\n    @tf.function\n    def scale_image(image):\n        image = tf.image.convert_image_dtype(image, tf.float32)\n        return image\n\n    def random_flip(image,random_i):\n\n        if random_i >= 0.5:\n            image = tf.image.flip_left_right(image)\n\n        elif random_i <= 0.3:\n             image = tf.image.flip_up_down(image)\n\n        return image\n\n    @tf.function\n    def tie_together(image, min_scale, max_scale, crop_size,random_i,id_):\n\n        # Scale the pixel values\n        image = scale_image(image)\n\n        image = random_flip(image,random_i)\n\n        # Random resized crops\n        image = random_resize_crop(image, min_scale, max_scale, crop_size)\n\n        # Random select one of the three images\n        image = image[:,:,id_:id_+1] \n\n        return image\n\n    def get_multires_dataset(dataset,size_crops,num_crops,min_scale,max_scale,options=None):\n\n        last_idx = 0 #control random number\n        last_image_id = 0\n\n        loaders = tuple()\n        for i, num_crop in enumerate(num_crops):\n\n            for _ in range(num_crop):\n\n                random_i = np.random.random()\n                image_random_id = np.random.randint(images_timesteps,size=1)[0]\n\n                while random_i == last_idx:\n                    random_i = np.random.random()\n\n                while image_random_id == last_image_id:\n                    image_random_id = np.random.randint(images_timesteps,size=1)[0]\n\n\n                loader = (dataset.shuffle(1024,seed=1).map(lambda x: tie_together(x, min_scale[i],\n                            max_scale[i], size_crops[i], random_i,image_random_id), num_parallel_calls=AUTO))\n\n                if options!=None:\n                    loader = loader.with_options(options)\n\n                loaders += (loader, )\n\n                last_idx = random_i\n                last_image_id = image_random_id\n\n        return loaders\n\n    def shuffle_zipped_output(a,b,c,d,e):\n\n        listify = [a,b,c,d,e]\n        random.shuffle(listify)\n\n        return listify[0], listify[1], listify[2], listify[3], listify[4]","metadata":{"execution":{"iopub.status.busy":"2023-11-28T15:36:36.139580Z","iopub.execute_input":"2023-11-28T15:36:36.139920Z","iopub.status.idle":"2023-11-28T15:36:36.161293Z","shell.execute_reply.started":"2023-11-28T15:36:36.139890Z","shell.execute_reply":"2023-11-28T15:36:36.160301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainloaders_save_name = 'trainloaders_zipped'\n\nif train_flag == 1:\n\n    try:\n        trainloaders_zipped = tf.data.Dataset.load(f'/kaggle/working/{trainloaders_save_name}')\n\n    except:\n        \n        print('Running exception')\n        \n        # Get multiple data loaders\n        trainloaders = get_multires_dataset(data,\n            size_crops=SIZE_CROPS,\n            num_crops=NUM_CROPS,\n            min_scale=MIN_SCALE,\n            max_scale=MAX_SCALE,\n            options=options)\n\n        del data\n        gc.collect()\n\n        # Zipping\n        trainloaders_zipped = tf.data.Dataset.zip(trainloaders)\n        trainloaders_zipped.save(trainloaders_save_name)\n        \n\n        del trainloaders\n        gc.collect()\n\n    try:\n        del data\n        gc.collect()\n        \n    except:\n        gc.collect()\n\n    # Final trainloader\n    trainloaders_zipped = (\n        trainloaders_zipped\n        .batch(BS)\n        .prefetch(AUTO)\n    )\n\n    gc.collect()\n\n\n    im1, im2, im3, im4, im5 = next(iter(trainloaders_zipped))\n    print(im1.shape, im2.shape, im3.shape, im4.shape, im5.shape)\n\n    \n    images = [i[3] for i in next(iter(trainloaders_zipped))]\n\n    fig, axes = plt.subplots(2, 3, figsize=(10, 7))\n\n    for ax, img in zip(axes.ravel(), images):\n        ax.imshow(img,cmap='Blues')\n        ax.axis('off')\n\n    # If there are more axes than images, turn off the remaining axes\n    for ax in axes.ravel()[len(images):]:\n        ax.axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-11-28T15:36:36.162306Z","iopub.execute_input":"2023-11-28T15:36:36.162585Z","iopub.status.idle":"2023-11-28T15:36:37.848364Z","shell.execute_reply.started":"2023-11-28T15:36:36.162558Z","shell.execute_reply":"2023-11-28T15:36:37.847378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"plot views:","metadata":{}},{"cell_type":"markdown","source":"### Model Architecture","metadata":{}},{"cell_type":"code","source":"from typing import List\nclass Sampling(tf.keras.layers.Layer):\n    \"\"\" Uses (z_mean, z_log_var) to sample z, the vector encoding a digit.\n        Extracted from: https://keras.io/examples/generative/vae/\n    \"\"\"\n\n    def call(self, inputs):\n        z_mean, z_log_var = inputs\n        batch = tf.shape(z_mean)[0]\n        dim = tf.shape(z_mean)[1]\n        epsilon = tf.random.normal(shape=(batch, dim), mean=0.0, stddev=1.0)\n\n        return z_mean + tf.exp(0.5 * z_log_var) * epsilon\n\ndef residual_block(x: tf.Tensor, filters: int, kernel_size: int, stride: int, bn: bool, name: str = None) -> tf.Tensor:\n    \"\"\" Residual block with identity or convolutional shortcut, depending if the number of channels of x is the same as\n        filters (identity) or not (convolutional).\n        Based on https://github.com/keras-team/keras/blob/master/keras/applications/resnet.py\n\n    Args:\n        x: input tensor.\n        filters: (int) filters of the conv layers.\n        kernel_size: (int) kernel size of the conv layers.\n        stride: (int) stride of the first layer.\n        bn: (bool) if True, use batch normalization after convolutional layers.\n        name: (str) block name.\n\n    Returns:\n        Output tensor for the residual block.\n    \"\"\"\n\n    if x.shape[-1] == filters:\n        shortcut = x\n    else:\n        shortcut = tf.keras.layers.Conv2D(filters, 1, strides=stride, padding='same', name=name + '_0_conv')(x)\n        if bn:\n            shortcut = tf.keras.layers.BatchNormalization(axis=-1, epsilon=1.001e-5, name=name + '_0_bn')(shortcut)\n\n    x = tf.keras.layers.Conv2D(filters, kernel_size, strides=stride, padding='same', name=name + '_1_conv')(x)\n    if bn:\n        x = tf.keras.layers.BatchNormalization(axis=-1, epsilon=1.001e-5, name=name + '_1_bn')(x)\n    x = tf.keras.layers.ReLU(name=name + '_1_relu')(x)\n\n    x = tf.keras.layers.Conv2D(filters, kernel_size, strides=1, padding='same', name=name + '_2_conv')(x)\n    if bn:\n        x = tf.keras.layers.BatchNormalization(axis=-1, epsilon=1.001e-5, name=name + '_2_bn')(x)\n\n    x = tf.keras.layers.Add(name=name + '_add')([shortcut, x])\n    x = tf.keras.layers.ReLU(name=name + '_out')(x)\n\n    return x\n\n\ndef resnet_encoder(input_shape: List[int], num_layers: List[int], filters: List[int], strides: List[int],\n                   latent_dim: int, bn: bool) -> tf.keras.Model:\n    \"\"\" Build an encoder based on ResNet. For ResNet18, for example, use:\n        - num_layers = [2, 2, 2, 2], filters = [64, 128, 256, 512], strides = [1, 2, 2, 2]\n\n        Adaptations for VAE use: after the final residual block, 2 dense layers produce z_mean and z_logvar.\n\n    Args:\n        input_shape: (list) input data shape.\n        num_layers: list containing the number of residual layers in each block.\n        filters: list with the number of filters for each residual block.\n        strides: list with the strides of the first layer of each residual block.\n        latent_dim: (int) dimensionality of the latent space.\n        bn: (bool) if True, use batch normalization after convolutions (including in residual units).\n\n    Returns:\n        tf.keras.Model encoder.\n    \"\"\"\n\n    inputs = tf.keras.Input(shape=input_shape)\n\n    x = tf.keras.layers.Conv2D(64, kernel_size=3, strides=1, padding=\"same\", name='conv1_conv')(inputs)\n    if bn:\n        x = tf.keras.layers.BatchNormalization(axis=-1, epsilon=1.001e-5, name='conv1_bn')(x)\n    x = tf.keras.layers.ReLU(name='conv1_relu')(x)\n\n    for block_idx, (n, f, stride) in enumerate(zip(num_layers, filters, strides)):\n\n        x = residual_block(x, filters=f, kernel_size=3, stride=stride, bn=bn, name=f'b{block_idx+1}_1')\n        for i in range(1, n):\n            x = residual_block(x, filters=f, kernel_size=3, stride=1, bn=bn, name=f'b{block_idx+1}_{i+1}')\n\n    # x = tf.keras.layers.Flatten()(x)\n\n    # z_mean = tf.keras.layers.Dense(latent_dim, name=\"z_mean\")(x)\n    # z_log_var = tf.keras.layers.Dense(latent_dim, name=\"z_log_var\")(x)\n    # z = Sampling()([z_mean, z_log_var])\n\n    encoder = tf.keras.Model(inputs, x, name=\"encoder\")\n\n    return encoder\n\ndef get_resnet_backbone():\n  # base_model = tf.keras.applications.ResNet50(include_top=False,\n                                              # weights='imagenet')\n  # base_model = models.Model(base_model.input, base_model.get_layer('conv3_block4_out').output)\n\n  base_model = resnet_encoder([None, None, 1], [2, 2, 2, 2],[64, 128, 256, 512], [1, 2, 2, 2] ,300, True)\n\n  #base_model.summary()\n  base_model.trainable = True\n\n  inputs = layers.Input((None, None, 1))\n\n  h = base_model(inputs, training=True)\n  h = layers.GlobalAveragePooling2D()(h)\n  backbone = models.Model(inputs, h)\n    \n  return backbone\n\ndef get_projection_prototype(dense_1=1024, dense_2=96, prototype_dimension=15):\n    inputs = layers.Input((512, ))\n    projection_1 = layers.Dense(dense_1)(inputs)\n    projection_1 = layers.BatchNormalization()(projection_1)\n    projection_1 = layers.Activation(\"relu\")(projection_1)\n\n    projection_2 = layers.Dense(dense_2)(projection_1)\n    projection_2_normalize = tf.math.l2_normalize(projection_2, axis=1, name='projection')\n\n    prototype = layers.Dense(prototype_dimension, use_bias=False, name='prototype')(projection_2_normalize)\n\n    return models.Model(inputs=inputs,outputs=[projection_2_normalize, prototype])\n\n\n\ndef sinkhorn(sample_prototype_batch,y_prop=[]):\n    \n    Q = tf.transpose(tf.exp(sample_prototype_batch/0.05))\n    Q /= tf.keras.backend.sum(Q)\n    K, B = Q.shape\n\n    u = tf.zeros_like(K, dtype=tf.float32)\n    c = tf.ones_like(B, dtype=tf.float32) / B\n    \n    if len(y_prop):\n        r = y_prop\n    else:\n        r = tf.ones_like(K, dtype=tf.float32) / K\n    \n    for _ in range(3):\n        u = tf.keras.backend.sum(Q, axis=1)\n        Q *= tf.expand_dims((r / u), axis=1)\n        Q *= tf.expand_dims(c / tf.keras.backend.sum(Q, axis=0), 0)\n\n    final_quantity = Q / tf.keras.backend.sum(Q, axis=0, keepdims=True)\n    final_quantity = tf.transpose(final_quantity)\n\n    return final_quantity\n","metadata":{"execution":{"iopub.status.busy":"2023-11-28T15:36:37.849994Z","iopub.execute_input":"2023-11-28T15:36:37.850348Z","iopub.status.idle":"2023-11-28T15:36:37.877689Z","shell.execute_reply.started":"2023-11-28T15:36:37.850314Z","shell.execute_reply":"2023-11-28T15:36:37.876771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Train","metadata":{}},{"cell_type":"code","source":"if train_flag == 1:\n    # @tf.function\n    # Reference: https://github.com/facebookresearch/swav/blob/master/main_swav.py\n    def train_step(input_views, feature_backbone, projection_prototype,\n                   optimizer, crops_for_assign, temperature):\n\n        # ============ retrieve input data ... ============\n        im1, im2, im3, im4, im5  = input_views\n        inputs = [im1, im2, im3, im4, im5]\n\n        batch_size = inputs[0].shape[0]\n\n        # ============ create crop entries with same shape ... ============\n        #A vector of indices to reorder as views with similar resolutions\n        crop_sizes = [inp.shape[1] for inp in inputs] # list of crop size of views\n        unique_consecutive_count = [len([elem for elem in g]) for _, g in groupby(crop_sizes)] # equivalent to torch.unique_consecutive\n\n        #(unique_consecutive_count)\n        idx_crops = tf.cumsum(unique_consecutive_count)\n        # print(idx_crops)\n\n        # ============ multi-res forward passes ... ============\n        # tf.stop_gradient have been placed carefully in order to exclude the computations from dependency tracing.\n        # This is useful any time you want to compute a value with TensorFlow but need to pretend that the value\n        #was a constant.\n        start_idx = 0\n\n        with tf.GradientTape() as tape:\n\n            for end_idx in idx_crops:\n\n                concat_input = tf.stop_gradient(tf.concat(inputs[start_idx:end_idx], axis=0))\n                _embedding = feature_backbone(concat_input) # get embedding of same dim views together\n                if start_idx == 0:\n                    embeddings = _embedding # for first iter\n                else:\n                    embeddings = tf.concat((embeddings, _embedding), axis=0) # concat all the embeddings from all the views\n                start_idx = end_idx\n\n            projection, prototype = projection_prototype(embeddings) # get normalized projection and prototype\n            projection = tf.stop_gradient(projection)\n\n            # ============ swav loss ... ============\n            # https://github.com/facebookresearch/swav/issues/19\n\n            loss = 0\n            for i, crop_id in enumerate(crops_for_assign): # crops_for_assign = [0,1] hold that we use to create these codes the views in these positions.\n                with tape.stop_recording():   # there to ensure the computations for cluster assignments do not get traced for gradient updates\n                    out = prototype[batch_size * crop_id: batch_size * (crop_id + 1)]\n\n                    # get assignments\n                    q = sinkhorn(out,y_prop) # sinkhorn is used for cluster assignment\n\n                # cluster assignment prediction\n                subloss = 0\n                for v in np.delete(np.arange(np.sum(NUM_CROPS)), crop_id): # (for rest of the portions compute p and take cross entropy with q)\n                    p = tf.nn.softmax(prototype[batch_size * v: batch_size * (v + 1)] / temperature)\n                    subloss -= tf.math.reduce_mean(tf.math.reduce_sum(q * tf.math.log(p), axis=1))\n                loss += subloss / tf.cast((tf.reduce_sum(NUM_CROPS) - 1), tf.float32)\n\n            loss /= len(crops_for_assign)\n\n        # ============ backprop ... ============\n        variables = feature_backbone.trainable_variables + projection_prototype.trainable_variables\n        gradients = tape.gradient(loss, variables)\n        optimizer.apply_gradients(zip(gradients, variables))\n\n        return loss\n\n    def train_swav(feature_backbone,\n                   projection_prototype,\n                   dataloader,\n                   optimizer,\n                   crops_for_assign,\n                   temperature,\n                   epochs=50):\n\n        step_wise_loss = []\n        epoch_wise_loss = []\n\n        for epoch in range(epochs):\n\n            # normalize the prototypes\n            w = projection_prototype.get_layer('prototype').get_weights()\n            w = tf.transpose(w)\n            w = tf.math.l2_normalize(w, axis=1)\n            projection_prototype.get_layer('prototype').set_weights(tf.transpose(w))\n\n            iter_data = iter(trainloaders_zipped)\n\n            t = trange(len(dataloader), position=0, leave=True)\n            for i in  t:\n\n                inputs = next(iter_data)\n                loss = train_step(inputs, feature_backbone, projection_prototype,\n                                  optimizer, crops_for_assign, temperature)\n                step_wise_loss.append(loss)\n                t.set_postfix(loss='{:05.3f}'.format(loss))\n\n            epoch_wise_loss.append(np.mean(step_wise_loss))\n\n            print(\"epoch: {} loss: {:.3f}\".format(epoch + 1, np.mean(step_wise_loss)))\n\n            print('Saving weights')\n            feature_backbone_weights = feature_backbone.get_weights()\n            feature_backbone.save(f'/kaggle/working/trained/feature2d_{prototypes}_{BS}_BS_proto{model_tag}.h5')\n            projection_prototype.save(f'/kaggle/working/trained/proj2d_{prototypes}_{BS}_BS_proto{model_tag}.h5')\n\n        return epoch_wise_loss, [feature_backbone, projection_prototype]","metadata":{"execution":{"iopub.status.busy":"2023-11-28T15:36:37.879062Z","iopub.execute_input":"2023-11-28T15:36:37.879333Z","iopub.status.idle":"2023-11-28T15:36:37.899235Z","shell.execute_reply.started":"2023-11-28T15:36:37.879308Z","shell.execute_reply":"2023-11-28T15:36:37.898335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ============ re-initialize the networks and the optimizer ... ============\nfeature_backbone = get_resnet_backbone()\nprojection_prototype = get_projection_prototype(prototype_dimension = prototypes)\n\nif train_flag == 1:\n\n    print(prototypes)\n\n    decay_steps = 1000\n    lr_decayed_fn = tf.keras.experimental.CosineDecay(\n        initial_learning_rate=0.1, decay_steps=decay_steps)\n\n    opt = tf.keras.optimizers.SGD(learning_rate=lr_decayed_fn)\n\n    # ======================= train  ===========================\n    try:\n        feature_backbone.load_weights(f'/kaggle/working/trained/feature2d_{prototypes}_{BS}_BS_proto{model_tag}.h5')\n        projection_prototype.load_weights(f'/kaggle/working/trained/proj2d_{prototypes}_{BS}_BS_proto{model_tag}.h5')\n    except:\n        print('weights not available')\n\n    epoch_wise_loss, models_tr = train_swav(feature_backbone,projection_prototype,\n                                            trainloaders_zipped,opt,\n                                            crops_for_assign=[0, 1],temperature=0.1,epochs=50)\n\n    feature_backbone_weights = feature_backbone.get_weights()\n    feature_backbone.save(f'/kaggle/working/trained/feature2d_{prototypes}_{BS}_BS_proto{model_tag}.h5')\n    projection_prototype.save(f'/kaggle/working/trained/proj2d_{prototypes}_{BS}_BS_proto{model_tag}.h5')\n    gc.collect()\n\nelse:\n    feature_backbone.load_weights(f'/kaggle/working/trained/feature2d_{prototypes}_{BS}_BS_proto{model_tag}.h5')\n    projection_prototype.load_weights(f'/kaggle/working/trained/proj2d_{prototypes}_{BS}_BS_proto{model_tag}.h5')","metadata":{"execution":{"iopub.status.busy":"2023-11-28T15:36:37.902576Z","iopub.execute_input":"2023-11-28T15:36:37.902901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Test train:","metadata":{}},{"cell_type":"markdown","source":"# Plots (with sample):","metadata":{}},{"cell_type":"code","source":"from scipy.special import softmax\nimport pandas as pd\n\nif plot_sample_flag:\n    \n\n    test_data = []\n    for sample in test_ds:\n\n        for v in range(sample.shape[-1]):\n            test_data.append(sample[:,:,v])\n\n    print(np.asarray(test_data).shape)\n\n    del test_ds\n    gc.collect()\n    \n    try:\n         assignments = pd.read_csv(f'/kaggle/working/sample_assignments_{prototypes}proto_{BS}batch{model_tag}.csv',index_col=0)[\"0\"].values\n    except:\n        \n        #brak in block so we dont run out of memory \n        blocks = 100\n        size = int(len(test_data)/blocks)\n\n        for i in range(blocks):\n\n            embeddings_ = feature_backbone(np.asarray(test_data)[i*size:(i+1)*size])\n            projection_, prototype_ = projection_prototype(embeddings_)\n\n            if i == 0:\n                prototype=np.asarray(prototype_)\n            else:\n                prototype=np.concatenate([prototype,np.asarray(prototype_)])\n\n        prototype = np.asarray(prototype)\n\n        print(prototype.shape)\n\n        assignments = np.argmax(softmax(prototype),axis=1)\n        print(assignments.shape)\n\n        pd.Series(assignments).to_csv(f'/kaggle/working/sample_assignments_{prototypes}proto_{BS}batch{model_tag}.csv') #save\n\n        del projection_,embeddings_\n        gc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## QQ PLOTS","metadata":{}},{"cell_type":"code","source":"if plot_sample_flag:\n    \n    min_samples= 100 #random samples to get from each cluster\n    max_samples= 25000\n\n    test_data_list = np.asarray(test_data)\n    random_samples_ids = np.random.choice(range(test_data_list.shape[0]),size=min_samples) #random samples to QQPLOT\n\n    colors_c = [plt.colormaps['tab20'](c) for c in np.linspace(0, 1, num=prototypes)]\n\n    #QQ PLOT TO COMPARE CLUSTERS\n    fig, ax = plt.subplots(figsize=(15, 15))\n\n    pplot  = sm.ProbPlot(data=test_data_list[random_samples_ids].ravel())\n    pplot.qqplot(other=pplot,ax=ax, marker='', linestyle='dashed',label= 'self') #45 line\n\n#     enumerate(range(prototype.shape[1]))\n    for num, cluster in enumerate(range(prototypes)): #iterates over each cluster (uses num to get new color)\n        \n        cur_cluster = test_data_list[assignments==cluster].copy()\n        cur_cluster_shape = cur_cluster.shape\n\n        if (cur_cluster_shape[0]>min_samples) and (cur_cluster_shape[0]<max_samples):\n#         if num in [3,11]:\n            #show cluster stats\n            print('Cluster:',cluster,'Shape:', cur_cluster_shape,  'Avg:',np.mean(cur_cluster.ravel()),  'Std:', np.std(cur_cluster.ravel()))\n\n            random_samples_ids = np.random.choice(range(cur_cluster_shape[0]),size=min_samples) #random select samples from current cluster\n\n            cur_color = colors_c[num] if num != 11 else 'black'\n            pplot.qqplot(other=cur_cluster[random_samples_ids].ravel(),ax=ax,marker='', linestyle='solid', color=cur_color,\n                          label= f'{cluster}: {cur_cluster_shape[0]}')\n            \n#         else:\n#             random_samples_ids = np.random.choice(range(cur_cluster_shape[0]),size=min_samples) #random select samples from current cluster\n\n#             pplot.qqplot(other=cur_cluster[random_samples_ids].ravel(),ax=ax,marker='', linestyle='solid', color='lightgrey',\n#                           label= f'{cluster}: {cur_cluster_shape[0]}')\n            \n\n    plt.legend()\n    plt.title('Cluster Q-Q Plot | 16% Training Set Sample')\n    plt.show()\n    \n    fig.savefig(f'qq_plot.png')\n        \n    del cur_cluster, pplot\n    gc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Hist","metadata":{}},{"cell_type":"code","source":"if plot_sample_flag:\n    n_bins = 50\n    fig, ax = plt.subplots(figsize=(15, 15))\n\n    plt.hist(test_data_list.ravel(),density=True,bins=n_bins,alpha = 0.5,label= 'self')\n\n    for num, cluster in enumerate(range(prototypes)):\n        \n\n        cur_cluster = test_data_list[assignments==cluster].copy()\n        cur_cluster_shape = cur_cluster.shape\n\n        if (cur_cluster_shape[0]>min_samples) and (cur_cluster_shape[0]<max_samples):\n#         if num in [12,3]:\n\n                cur_color = colors_c[num] if num != 11 else 'black'\n                plt.hist(cur_cluster.ravel(),density=True,bins=n_bins,histtype='step',color=cur_color,label= f'{cluster}: {cur_cluster_shape[0]}',linewidth=1.5)\n            \n#         else:\n#                  plt.hist(cur_cluster.ravel(),density=True,bins=n_bins,histtype='step',color='lightgrey',label= f'{cluster}: {cur_cluster_shape[0]}',linewidth=0.5)\n            \n\n    plt.legend()\n    plt.title('Cluster Histogram | 16% Training Set Sample')\n    fig.savefig(f'hist.png')\n    \n#     ax.set_xlim(0.0,.60)\n#     ax.set_ylim(0,0.6)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if plot_sample_flag:\n    \n    n_dis=5\n    min_samples_plot = 150\n    max_samples_plot = 800\n\n    for num, cluster in enumerate(range(prototypes)):\n\n        cur_cluster = test_data_list[assignments==cluster].copy()\n        cur_cluster_shape = cur_cluster.shape\n\n        if ((cur_cluster_shape[0]<min_samples_plot) or (cur_cluster_shape[0]>max_samples_plot)) and (cur_cluster_shape[0]>0):\n\n            fig, ax1 = plt.subplots()\n\n            images_ = np.zeros([cur_cluster_shape[1],cur_cluster_shape[2]*n_dis])\n            idxs=np.asarray(range(cur_cluster_shape[0]))\n            \n            \n            \n            random_samples_ids = np.random.choice(range(cur_cluster_shape[0]),size=n_dis) #random select samples from current cluster\n\n\n            for j in range(n_dis):\n                images_[:,j*cur_cluster_shape[2]:(j+1)*cur_cluster_shape[2]] = cur_cluster[random_samples_ids[j],:,:]\n\n            plt.title(f'Cluster: {cluster} with {cur_cluster_shape[0]} obs')\n            plt.imshow(images_,vmin=0, vmax=0.5)\n            plt.tick_params(left = False, right = False , labelleft = False ,\n                    labelbottom = False, bottom = False)\n\n            plt.show()\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Cluster change x time (All train samples)\n","metadata":{}},{"cell_type":"markdown","source":"get assigments:","metadata":{}},{"cell_type":"code","source":"from math import ceil, floor\nimport pandas as pd\n\nif plot_all_flag:\n\n    try:\n        assignments = pd.read_csv(f'/kaggle/working/assignments_{prototypes}proto_{BS}batch{model_tag}.csv',index_col=0)[\"0\"].values\n        print(len(assignments))\n\n\n    except:\n\n\n        #get 1d time series\n        data_plot = []\n        for sample in data:\n\n            for v in range(sample.shape[-1]):\n                data_plot.append(sample[:,:,v])\n\n        gc.collect()\n\n        #get assigments\n        blocks = 1000\n        size = ceil(len(data_plot)/blocks)\n\n        for i in range(blocks):\n\n            embeddings_ = feature_backbone(np.asarray(data_plot)[i*size:(i+1)*size])\n            projection_, prototype_ = projection_prototype(embeddings_)\n\n            if i == 0:\n                prototype=np.asarray(prototype_)\n            else:\n                prototype=np.concatenate([prototype,np.asarray(prototype_)])\n\n        prototype = np.asarray(prototype)\n\n        print(prototype.shape)\n\n        assignments = np.argmax(softmax(prototype),axis=1)\n        print(assignments.shape)\n\n        del projection_,embeddings_\n        gc.collect()\n\n        pd.Series(assignments).to_csv(f'/kaggle/working/assignments_{prototypes}proto_{BS}batch{model_tag}.csv') #save\n\n\n# #remove step 3 \ndata = data.transpose(0,3,1,2)\ndata = data.reshape(data.shape[0]*data.shape[1], data.shape[2], data.shape[3])\nprint(data.shape)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    \n    cluster_dist = pd.read_csv(f'/kaggle/working/cluster_dist{prototypes}proto_{BS}batch{model_tag}.csv',index_col=0) \n    cluster_dist.index = pd.to_datetime(cluster_dist.index)\n    time_ref= np.load('/kaggle/input/temperature-era5/time_1940_2023.npy')[::2]\n    \nexcept:\n    cluster_dist =  pd.DataFrame(index=range(len(assignments)),columns=['Cluster','Average','Median','Q75','Q25'])\n\n    for idx, cluster in enumerate(assignments):\n\n        cur_point = data[idx]\n\n        cluster_dist.loc[idx, 'Average'] = cur_point.mean()\n        cluster_dist.loc[idx, 'Median'] = np.median(cur_point)\n        cluster_dist.loc[idx, 'Q75'] =np.quantile(cur_point,q=0.75)\n        cluster_dist.loc[idx, 'Q25'] = np.quantile(cur_point,q=0.25)\n        cluster_dist.loc[idx, 'Cluster'] = cluster\n\n    cluster_dist.head()\n    \n    time_ref= np.load('/kaggle/input/temperature-era5/time_1940_2023.npy')[::2]\n    cluster_dist.index = pd.to_datetime(time_ref)\n\n    cluster_dist.to_csv(f'/kaggle/working/cluster_dist{prototypes}proto_{BS}batch{model_tag}.csv') #save","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Count Plots","metadata":{}},{"cell_type":"code","source":"new_block_size = 1531\nassigments_blocks = assignments.reshape(new_block_size,-1) \n\ntime_ref = time_ref.reshape(new_block_size,-1) \ntimeline = [pd.to_datetime(time_ref[i][0]).date() for i in range(time_ref.shape[0])]\n\nproxy_days = len(assigments_blocks[0])/12\nprint('~day in a group',proxy_days)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Count assigments per group every 30 days:  ","metadata":{}},{"cell_type":"code","source":"cluster_counter = pd.DataFrame(index=range(len(assigments_blocks)),columns=np.unique(assignments))\n\nfor step, values in enumerate(assigments_blocks):\n    \n    unique, counts = np.unique(values, return_counts=True)\n    \n    cluster_counter.loc[step,unique] = counts\n    \ncluster_counter = cluster_counter.fillna(0)\n\n#get COUNT stats\ncount_stats = cluster_counter.describe().round(2).drop('count')\ncount_stats.loc['sum',cluster_counter.columns]=cluster_counter.sum()\ncount_stats =count_stats.T","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"count_stats.head() #the index is the cluster ID","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"count_stats['sum'].plot.bar(figsize=(20,5),title='Cluster Assigments Total Count')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"(count_stats['sum']/(count_stats['sum'].sum()))*100","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cluster_counter.plot.box(figsize=(20,5),title=f'Assignment Distribution by Group Over {proxy_days}-Day Intervals')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#get top 8 cluster with higher STD\nselected_clusters = list(count_stats['std'].rank().sort_values().tail(8).index)\ncount_stats.loc[selected_clusters]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rolling_wind_stats = int(assigments_blocks.shape[1]) #hours unit\nrolling_wind_count = int(360/proxy_days)  #~one year\n\nfig, ax = plt.subplots(len(selected_clusters), ncols=2,figsize=(20,30))\n\nfor idx, cluster in enumerate(sorted(selected_clusters)):\n\n    #ver se o cluster é consistente ao longo do tempo\n    cluster_dist.loc[cluster_dist['Cluster']==cluster,'Average'].rolling(rolling_wind_stats).mean().plot(legend=True,ax=ax[idx,0],title=cluster)\n    cluster_dist.loc[cluster_dist['Cluster']==cluster,'Median'].rolling(rolling_wind_stats).mean().plot(legend=True,ax=ax[idx,0])\n    cluster_dist.loc[cluster_dist['Cluster']==cluster,'Q75'].rolling(rolling_wind_stats).mean().plot(legend=True,ax=ax[idx,0])\n    cluster_dist.loc[cluster_dist['Cluster']==cluster,'Q25'].rolling(rolling_wind_stats).mean().plot(legend=True,ax=ax[idx,0])\n    \n    cluster_dist.loc[cluster_dist['Cluster']==cluster,'Average'].plot.hist(ax=ax[idx,1])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Contagem dos clusters (sem sazonalidade):","metadata":{}},{"cell_type":"code","source":"#1-Year and 5-Year Simple Moving Averages (to remove the sazonality) of Count of Assignments per Group\nfig, ax = plt.subplots(ncols=2, nrows=cluster_counter.columns[-1]+1, figsize=(20,50))\n\nfor  cluster in cluster_counter:\n    \n    cluster_counter[cluster].rolling(rolling_wind_count).mean().plot(ax=ax[cluster,0],title=f'Cluster {cluster} | Assigments Counts Moving Avg.',label='1Y Moving Avg')\n    cluster_counter[cluster].rolling(rolling_wind_count*5).mean().plot(ax=ax[cluster,0],label='5Y Moving Avg')\n#     cluster_counter[cluster].plot.hist(ax=ax[cluster,1])\n\n    cluster_dist.loc[cluster_dist['Cluster']==cluster,'Average'].plot.hist(ax=ax[cluster,1],title=f'Cluster {cluster} | Average Precipitation on Samples Histogram')\n    \n    \nplt.legend()\n# fig.savefig(f'cluster_{prototypes}_dist.png')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cur_assig.index[0]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"assignments_series = pd.Series(assignments)\n\nfor cluster in np.unique(assignments):\n    \n    fig, ax = plt.subplots(2,2,figsize=(5,5))\n    fig.subplots_adjust(wspace=0.2, hspace=0.01)\n    \n    cur_assig = assignments_series[assignments_series==cluster].copy()\n    \n    if len(cur_assig):\n        cur_assig = cur_assig.sample(4).index\n\n        ax[0,0].imshow(data[cur_assig[0]])\n        ax[0,1].imshow(data[cur_assig[1]])\n        ax[1,0].imshow(data[cur_assig[2]])\n        ax[1,1].imshow(data[cur_assig[3]])\n#         ax[1,1].imshow(data[cur_assig[4]])\n#         ax[1,2].imshow(data[cur_assig[5]])\n        \n        fig.suptitle(f'Cluster {cluster} Samples' )\n        plt.show()\n        \n    ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cluster_counter.plot.bar(stacked=True,figsize=(25,10), colormap='tab20',title='Count of Cluster Assigmentss x Time')\nplt.legend(loc='center left', bbox_to_anchor=(1.0, 0.5))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cluster_counter_rolling = cluster_counter.copy()\nfor  cluster in cluster_counter_rolling:\n    cluster_counter_rolling[cluster] = cluster_counter_rolling[cluster].rolling(rolling_wind_count).mean()\n    \ncluster_counter_rolling.dropna(axis=0,how='all').plot.bar(stacked=True,figsize=(25,10), colormap='tab20',title='Count of Cluster Assigments 1Y Moving Average x Time')\nplt.legend(loc='center left', bbox_to_anchor=(1.0, 0.5))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cluster_counter_rolling = cluster_counter.copy()\nfor  cluster in cluster_counter_rolling:\n    cluster_counter_rolling[cluster] = cluster_counter_rolling[cluster].rolling(rolling_wind_count*5).mean()\n    \ncluster_counter_rolling.dropna(axis=0,how='all').plot.bar(stacked=True,figsize=(25,10), colormap='tab20',title='Count of Cluster Assigments 5Y Moving Average x Time')\nplt.legend(loc='center left', bbox_to_anchor=(1.0, 0.5))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Proportions:","metadata":{}},{"cell_type":"code","source":"if prototypes == 15:\n    total =count_stats['sum'].values\n    total = pd.Series(total)\n    \n    hist_groups = [[1, 4, 12, 8], [11], [3, 2], [7], [5], [14, 9], [0, 6], [13], [10]]\n    hist_groups = pd.Series(hist_groups).to_frame('Hist Clusters')\n\n    for idx, row in hist_groups.iterrows():\n\n        clusters = row['Hist Clusters']\n        hist_groups.loc[idx,'Sum'] = total.loc[clusters].sum()\n\n    hist_groups['Percent'] = (hist_groups['Sum']/total.sum())*100\n    hist_groups.set_index('Hist Clusters')\n    display(hist_groups)\n    \n    \n    qq_groups = [[8, 4, 12], [11], [2, 3, 7], [1], [9], [5], [6, 0, 13, 14], [10]]\n    qq_groups = pd.Series(qq_groups).to_frame('QQ Plot Clusters')\n\n    for idx, row in qq_groups.iterrows():\n\n        clusters = row['QQ Plot Clusters']\n        qq_groups.loc[idx,'Sum'] = total.loc[clusters].sum()\n\n    qq_groups['Percent'] = (qq_groups['Sum']/total.sum())*100\n    qq_groups.set_index('QQ Plot Clusters')\n    display(qq_groups)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}