{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"accelerator":"TPU","colab":{"authorship_tag":"ABX9TyPqKkAtrWIPshBfibnqjQkk","gpuType":"V28","machine_shape":"hm","name":"","version":""},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"},{"sourceId":11672208,"sourceType":"datasetVersion","datasetId":4402985}],"dockerImageVersionId":30685,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"1000 epochs","metadata":{"id":"mIkmEapk_ku6"}},{"cell_type":"markdown","source":"# config","metadata":{"id":"sDY_LjXw7_wT"}},{"cell_type":"code","source":"env = \"Kaggle\"\nDEBUG = False\nseed_num = 1\nmodel_name = f'Denoiser-{seed_num}'\nTPU = True","metadata":{"executionInfo":{"elapsed":6,"status":"ok","timestamp":1750539243598,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"d79cwdoz5Ryp","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:40:05.258512Z","iopub.execute_input":"2025-06-27T17:40:05.258748Z","iopub.status.idle":"2025-06-27T17:40:05.274169Z","shell.execute_reply.started":"2025-06-27T17:40:05.258715Z","shell.execute_reply":"2025-06-27T17:40:05.273556Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Imports","metadata":{"id":"CmH4B6Le8F9r"}},{"cell_type":"code","source":"import os\nif TPU:\n    os.environ[\"KERAS_BACKEND\"] = \"tensorflow\"\n    !pip install keras==2.15.0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:40:05.275300Z","iopub.execute_input":"2025-06-27T17:40:05.275529Z","iopub.status.idle":"2025-06-27T17:40:12.442920Z","shell.execute_reply.started":"2025-06-27T17:40:05.275506Z","shell.execute_reply":"2025-06-27T17:40:12.441777Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if TPU and env=='Colab':\n    !pip install tensorflow==2.18.0\n    !pip install tensorflow-tpu==2.18.0 --find-links=https://storage.googleapis.com/libtpu-tf-releases/index.html\n\nimport tensorflow as tf","metadata":{"executionInfo":{"elapsed":112789,"status":"ok","timestamp":1750539356390,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"eRs2SQ_9-wVe","outputId":"a40a08fb-8463-40ab-ee72-b7268ab82eef","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:40:12.444354Z","iopub.execute_input":"2025-06-27T17:40:12.444741Z","iopub.status.idle":"2025-06-27T17:40:29.088272Z","shell.execute_reply.started":"2025-06-27T17:40:12.444700Z","shell.execute_reply":"2025-06-27T17:40:29.087526Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, gc\nimport pandas as pd\nimport numpy as np\nfrom sklearn.model_selection import KFold\nimport sklearn\nimport matplotlib.pyplot as plt\nimport pickle\nimport shutil\n\nimport time\n\nimport scipy.stats as stats\nimport math","metadata":{"executionInfo":{"elapsed":230,"status":"ok","timestamp":1750539356622,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"8vYYu_pt8Ivl","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:40:29.089205Z","iopub.execute_input":"2025-06-27T17:40:29.089593Z","iopub.status.idle":"2025-06-27T17:40:31.061862Z","shell.execute_reply.started":"2025-06-27T17:40:29.089567Z","shell.execute_reply":"2025-06-27T17:40:31.061040Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Connect to drive/Save folder","metadata":{"id":"1AginKlE8MCv"}},{"cell_type":"code","source":"import os\nimport json\n\nsave_folder_name = f'models/model {model_name}'\n\nif env == 'Kaggle':\n    base_folder = '/kaggle/working'\n    save_folder = '/kaggle/working/'\n    try:\n        os.mkdir(save_folder)\n    except Exception as e:\n        print('exception error:')\n        print(e)\n    print(os.listdir('/kaggle'))\n    f = open('/kaggle/input/kaggle-json/kaggle.json')\n    kaggle_json = json.load(f)\n    KAGGLE_USERNAME = kaggle_json['username']\n    KAGGLE_KEY = kaggle_json['key']\n    os.environ[\"KAGGLE_USERNAME\"] = KAGGLE_USERNAME\n    os.environ[\"KAGGLE_KEY\"] = KAGGLE_KEY\nelif env == 'Colab':\n    from google.colab import drive\n    drive.mount('/content/drive')\n    save_folder = '/content/save_folder'\n    try:\n        os.mkdir(save_folder)\n    except Exception as e:\n        print('exception error:')\n        print(e)\n\n    f = open('/content/drive/MyDrive/kaggle/kaggle_auth/kaggle.json')\n    kaggle_json = json.load(f)\n\n    KAGGLE_USERNAME = kaggle_json['username']\n    KAGGLE_KEY = kaggle_json['key']\n\n    os.environ[\"KAGGLE_USERNAME\"] = KAGGLE_USERNAME\n    os.environ[\"KAGGLE_KEY\"] = KAGGLE_KEY","metadata":{"executionInfo":{"elapsed":27950,"status":"ok","timestamp":1750539384575,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"JccT2_Yt7y3x","outputId":"e6b35166-81e8-46d9-f80a-2c87537c9d70","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:40:31.064047Z","iopub.execute_input":"2025-06-27T17:40:31.064479Z","iopub.status.idle":"2025-06-27T17:40:31.077351Z","shell.execute_reply.started":"2025-06-27T17:40:31.064452Z","shell.execute_reply":"2025-06-27T17:40:31.076670Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Connect to TPU","metadata":{"id":"bgLhyBuF8Pge"}},{"cell_type":"code","source":"try:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver(tpu='local')\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.TPUStrategy(tpu)\n    print(\"TPU is running:\", tpu.master())\n    print(\"on TPU\")\n    print(\"REPLICAS: \", strategy.num_replicas_in_sync)\n    shards_num = 8\nexcept:\n    print(\"TPU NO\")\n    strategy = tf.distribute.get_strategy()\n    shards_num = 1\nprint(strategy)","metadata":{"executionInfo":{"elapsed":26135,"status":"ok","timestamp":1750539410896,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"lf8i8ZNF84Kp","outputId":"d75443ef-0e2a-4aa5-b031-f836493ccd9d","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:40:31.078202Z","iopub.execute_input":"2025-06-27T17:40:31.078445Z","iopub.status.idle":"2025-06-27T17:40:39.625601Z","shell.execute_reply.started":"2025-06-27T17:40:31.078421Z","shell.execute_reply":"2025-06-27T17:40:39.624809Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Download data/buckets/TFRecords path","metadata":{"id":"w5-Mo7VKDk6s"}},{"cell_type":"code","source":"!pip install kaggle","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:40:39.626482Z","iopub.execute_input":"2025-06-27T17:40:39.626714Z","iopub.status.idle":"2025-06-27T17:40:43.563179Z","shell.execute_reply.started":"2025-06-27T17:40:39.626690Z","shell.execute_reply":"2025-06-27T17:40:43.561947Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if env=='Colab':\n    !kaggle datasets download shlomoron/gwi-cleaned-models-ds\n    shutil.unpack_archive(f'/content/gwi-cleaned-models-ds.zip', 'cleaned_models')\n    !kaggle datasets download shlomoron/gwi-val-ensemble-ds\n    shutil.unpack_archive(f'/content/gwi-val-ensemble-ds.zip', 'val_ensemble')\nelse:\n    !kaggle datasets download shlomoron/gwi-cleaned-models-ds\n    shutil.unpack_archive(f'/kaggle/working/gwi-cleaned-models-ds.zip', 'cleaned_models')\n    !kaggle datasets download shlomoron/gwi-val-ensemble-ds\n    shutil.unpack_archive(f'/kaggle/working/gwi-val-ensemble-ds.zip', 'val_ensemble')","metadata":{"executionInfo":{"elapsed":1832,"status":"ok","timestamp":1750539412696,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"HvTPFnEXDjvs","outputId":"976efa54-4d88-4de7-9713-e7584d979fc0","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:40:43.564787Z","iopub.execute_input":"2025-06-27T17:40:43.565158Z","iopub.status.idle":"2025-06-27T17:42:40.125401Z","shell.execute_reply.started":"2025-06-27T17:40:43.565125Z","shell.execute_reply":"2025-06-27T17:42:40.124296Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if env=='Kaggle':\n    models_list_all = pickle.load(open('/kaggle/working/cleaned_models/models_list.p', 'br'))\n    val_ensemble_preds = pickle.load(open('/kaggle/working/val_ensemble/preds.p', 'br'))\n    val_ensemble_labels = pickle.load(open('/kaggle/working/val_ensemble/val_labels.p', 'br'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:42:40.126814Z","iopub.execute_input":"2025-06-27T17:42:40.127120Z","iopub.status.idle":"2025-06-27T17:42:46.510335Z","shell.execute_reply.started":"2025-06-27T17:42:40.127091Z","shell.execute_reply":"2025-06-27T17:42:46.509509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"shuffle_rng = np.random.default_rng(seed=4224)\nfor x in models_list_all:\n    shuffle_rng.shuffle(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:42:46.511378Z","iopub.execute_input":"2025-06-27T17:42:46.511637Z","iopub.status.idle":"2025-06-27T17:42:49.433224Z","shell.execute_reply.started":"2025-06-27T17:42:46.511613Z","shell.execute_reply":"2025-06-27T17:42:49.432351Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_models_list = [x[:500] for x in models_list_all]\ntrain_models_list = [x[500:] for x in models_list_all]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:42:49.434257Z","iopub.execute_input":"2025-06-27T17:42:49.434542Z","iopub.status.idle":"2025-06-27T17:43:15.774812Z","shell.execute_reply.started":"2025-06-27T17:42:49.434514Z","shell.execute_reply":"2025-06-27T17:43:15.773987Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DEBUG:\n    val_models_list = [x[:100] for x in val_models_list]\n    train_models_list = [x[:100] for x in models_list_all]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:15.775809Z","iopub.execute_input":"2025-06-27T17:43:15.776141Z","iopub.status.idle":"2025-06-27T17:43:16.472971Z","shell.execute_reply.started":"2025-06-27T17:43:15.776113Z","shell.execute_reply":"2025-06-27T17:43:16.472205Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(train_models_list))\nprint(len(train_models_list[-1]))\nplt.imshow(train_models_list[-1][43][0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:16.473991Z","iopub.execute_input":"2025-06-27T17:43:16.474293Z","iopub.status.idle":"2025-06-27T17:43:16.811153Z","shell.execute_reply.started":"2025-06-27T17:43:16.474267Z","shell.execute_reply":"2025-06-27T17:43:16.810390Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(val_ensemble_preds[9634])\nplt.show()\nplt.imshow(val_ensemble_labels[9634])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:16.815085Z","iopub.execute_input":"2025-06-27T17:43:16.815398Z","iopub.status.idle":"2025-06-27T17:43:17.117924Z","shell.execute_reply.started":"2025-06-27T17:43:16.815370Z","shell.execute_reply":"2025-06-27T17:43:17.117065Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Configs","metadata":{"id":"5g_fwN1RFSuA"}},{"cell_type":"code","source":"if DEBUG:\n    N_EPOCHS = 5\n    batch_size = 2\n    val_batch_size = 2\nelse:\n    N_EPOCHS = 500\n    batch_size = 512\n    val_batch_size = 512\nN_WARMUP_EPOCHS = 0\nLR_MAX = 1e-3\nLR_MIN = 8.4e-4\nWD_RATIO = 4.0\nWARMUP_METHOD = \"exp\"\nPAD = 0.0\nPAD_16 = tf.cast(PAD, tf.bfloat16)\n\nsteps_per_epoch = 24","metadata":{"executionInfo":{"elapsed":4,"status":"ok","timestamp":1750539413031,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"1HRAXQesFTtS","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:17.118937Z","iopub.execute_input":"2025-06-27T17:43:17.119250Z","iopub.status.idle":"2025-06-27T17:43:17.128614Z","shell.execute_reply.started":"2025-06-27T17:43:17.119219Z","shell.execute_reply":"2025-06-27T17:43:17.127810Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"training_epochs = N_EPOCHS","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:17.129604Z","iopub.execute_input":"2025-06-27T17:43:17.130131Z","iopub.status.idle":"2025-06-27T17:43:17.173704Z","shell.execute_reply.started":"2025-06-27T17:43:17.130098Z","shell.execute_reply":"2025-06-27T17:43:17.172878Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# TF data pipeline","metadata":{"id":"Q3pWxBJBQeUo"}},{"cell_type":"markdown","source":"# Test set","metadata":{}},{"cell_type":"code","source":"def val_ds_to_dict(features, labels, class_labels):\n    x = {}\n    x['features'] = features\n    x['labels'] = labels\n    x['class_labels'] = class_labels\n    return x\n\ndef get_output(x):\n    labels = x['labels']\n    class_labels = x['class_labels']\n    class_labels = labels*0.0+tf.cast(class_labels, tf.float32)\n    labels = labels[None]\n    class_labels = class_labels[None]\n    labels_concat = tf.concat([labels, class_labels], axis = 0)\n    return x['features'], labels_concat","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:17.174591Z","iopub.execute_input":"2025-06-27T17:43:17.174848Z","iopub.status.idle":"2025-06-27T17:43:17.267551Z","shell.execute_reply.started":"2025-06-27T17:43:17.174822Z","shell.execute_reply":"2025-06-27T17:43:17.266697Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_labels = sum([[x for i in range(1000)] for x in range(10)], [])\nclass_labels = np.asarray(class_labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:17.268575Z","iopub.execute_input":"2025-06-27T17:43:17.268848Z","iopub.status.idle":"2025-06-27T17:43:17.292166Z","shell.execute_reply.started":"2025-06-27T17:43:17.268822Z","shell.execute_reply":"2025-06-27T17:43:17.291336Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_ensemble_preds = np.concatenate([val_ensemble_preds, val_ensemble_preds[-1000:]], axis = 0)\nval_ensemble_labels = np.concatenate([val_ensemble_labels, val_ensemble_labels[-1000:]], axis = 0)\nclass_labels = np.concatenate([class_labels, class_labels[-1000:]], axis = 0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:17.293094Z","iopub.execute_input":"2025-06-27T17:43:17.293329Z","iopub.status.idle":"2025-06-27T17:43:17.495574Z","shell.execute_reply.started":"2025-06-27T17:43:17.293306Z","shell.execute_reply":"2025-06-27T17:43:17.494569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PAD = 0.0\nplt.plot(class_labels)\nval_ds = tf.data.Dataset.from_tensor_slices((val_ensemble_preds, val_ensemble_labels, class_labels))\nval_ds = val_ds.map(val_ds_to_dict, tf.data.AUTOTUNE)\nval_ds = val_ds.map(get_output, tf.data.AUTOTUNE)\nval_ds = val_ds.cache()\nif DEBUG:\n    val_ds = val_ds.take(64)\nsamples_num = val_ds.reduce(0, lambda x,_: x+1).numpy()\nval_ds = val_ds.padded_batch(\n            batch_size, padding_values=(PAD,PAD),\n    padded_shapes=([70,70],[2,70,70]), drop_remainder=True)\nval_ds = val_ds.prefetch(tf.data.AUTOTUNE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:17.496602Z","iopub.execute_input":"2025-06-27T17:43:17.496896Z","iopub.status.idle":"2025-06-27T17:43:19.304501Z","shell.execute_reply.started":"2025-06-27T17:43:17.496869Z","shell.execute_reply":"2025-06-27T17:43:19.303732Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"laa = [x for x in val_ds.skip(3000).take(2)]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:19.305453Z","iopub.execute_input":"2025-06-27T17:43:19.305720Z","iopub.status.idle":"2025-06-27T17:43:20.155750Z","shell.execute_reply.started":"2025-06-27T17:43:19.305694Z","shell.execute_reply":"2025-06-27T17:43:20.154795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"models_list_class_labels = []\nfor i, x in enumerate(models_list_all):\n    models_list_class_labels = models_list_class_labels+[i for k in range(len(x))]\n\nplt.plot(models_list_class_labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:20.156912Z","iopub.execute_input":"2025-06-27T17:43:20.157315Z","iopub.status.idle":"2025-06-27T17:43:20.392689Z","shell.execute_reply.started":"2025-06-27T17:43:20.157284Z","shell.execute_reply":"2025-06-27T17:43:20.391929Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_ds_to_dict(labels, class_labels):\n    x = {}\n    x['features'] = (labels-3000)/100\n    x['labels'] = (labels-3000)/100\n    x['class_labels'] = class_labels\n    return x\n\ndef add_gaussian_noise(data):\n    stddev = tf.random.uniform(shape=[], minval=0.0, maxval=0.5, dtype=tf.float32)\n    noise = tf.random.normal(shape=[70, 70], mean=0.0, stddev=stddev, dtype=tf.float32)\n    noisy_labels = data['features'] + noise\n    new_data = {\n        'labels': data['labels'],\n        'class_labels': data['class_labels'],\n        'features': noisy_labels\n    }\n    return new_data\n\ndef apply_gaussian_blur(data):\n    def gaussian_kernel(size=5, sigma=1.0):\n        x = tf.range(-size // 2 + 1, size // 2 + 1, dtype=tf.float32)\n        y = tf.range(-size // 2 + 1, size // 2 + 1, dtype=tf.float32)\n        x, y = tf.meshgrid(x, y)\n        kernel = tf.exp(-(x**2 + y**2) / (2.0 * sigma**2))\n        kernel = kernel / tf.reduce_sum(kernel)  # Normalize kernel\n        return tf.reshape(kernel, [size, size, 1, 1])\n\n    features = data['features']\n    \n    size_candidates = tf.constant([3, 5, 7], dtype=tf.int32)\n    size_idx = tf.random.uniform(shape=[], minval=0, maxval=3, dtype=tf.int32)\n    size = size_candidates[size_idx]\n    # Generate random sigma (float between 0 and 2)\n    sigma = tf.random.uniform(shape=[], minval=0.0, maxval=2.0, dtype=tf.float32)\n\n    # Apply Gaussian blur using convolution\n    kernel = gaussian_kernel(size, sigma)\n    \n    # Ensure features is in the correct shape: [1, 70, 70, 1]\n    features_expanded = tf.expand_dims(tf.expand_dims(features, 0), -1)\n    \n    # Apply padding and convolution\n    half_size = size//2  # For a 5x5 kernel\n    padded_features = tf.pad(features_expanded, [[0, 0], [half_size, half_size], [half_size, half_size], [0, 0]], mode='SYMMETRIC')\n    blurred_features = tf.nn.conv2d(\n        padded_features,\n        kernel,\n        strides=[1, 1, 1, 1],\n        padding='VALID'\n    )\n    \n    # Squeeze back to [70, 70]\n    blurred_features = tf.squeeze(blurred_features, [0, -1])\n    \n    return {\n        'labels': data['labels'],\n        'class_labels': data['class_labels'],\n        'features': blurred_features\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:20.393726Z","iopub.execute_input":"2025-06-27T17:43:20.393989Z","iopub.status.idle":"2025-06-27T17:43:20.417837Z","shell.execute_reply.started":"2025-06-27T17:43:20.393963Z","shell.execute_reply":"2025-06-27T17:43:20.416956Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train/val set","metadata":{}},{"cell_type":"code","source":"rng = np.random.default_rng(seed=seed_num)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:20.418744Z","iopub.execute_input":"2025-06-27T17:43:20.418976Z","iopub.status.idle":"2025-06-27T17:43:20.432033Z","shell.execute_reply.started":"2025-06-27T17:43:20.418953Z","shell.execute_reply":"2025-06-27T17:43:20.431311Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_dataset(models_list, batch_size, shuffle, split = 'valid', to_repeat = False, drop_reminder = True, cache = False):\n    models_list_class_labels = []\n    for i, x in enumerate(models_list):\n        models_list_class_labels = models_list_class_labels+[i for k in range(len(x))]\n        \n    models_list = np.concatenate([x[:,0,:,:] for x in models_list], axis = 0)\n    models_list_class_labels = np.asarray(models_list_class_labels)\n\n    if split == 'valid':\n        models_list = np.concatenate([models_list, models_list[-1000:]], axis = 0)\n        models_list_class_labels = np.concatenate([models_list_class_labels, models_list_class_labels[-1000:]], axis = 0)\n    \n    if split=='train':\n        random_indices = np.asarray(range(len(models_list_class_labels)))\n        rng.shuffle(random_indices)\n        models_list = models_list[random_indices]\n        models_list_class_labels = models_list_class_labels[random_indices]\n\n\n    ds = tf.data.Dataset.from_tensor_slices((models_list, models_list_class_labels))\n    ds = ds.map(train_ds_to_dict, tf.data.AUTOTUNE)\n\n    if shuffle:\n        ds = ds.shuffle(shuffle, reshuffle_each_iteration = True)\n    if to_repeat:\n        ds = ds.repeat()\n        \n    ds = ds.map(apply_gaussian_blur, num_parallel_calls=tf.data.AUTOTUNE)\n    ds = ds.map(add_gaussian_noise, num_parallel_calls=tf.data.AUTOTUNE)\n    ds = ds.map(get_output, tf.data.AUTOTUNE)\n\n    if cache:\n        ds = ds.cache()\n        samples_num = ds.reduce(0, lambda x,_: x+1).numpy()\n\n    if DEBUG:\n        ds = ds.take(64)\n        \n    ds = ds.padded_batch(\n                batch_size, padding_values=(PAD,PAD),\n        padded_shapes=([70,70],[2,70,70]), drop_remainder=True)\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n\n    return ds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:20.432820Z","iopub.execute_input":"2025-06-27T17:43:20.433058Z","iopub.status.idle":"2025-06-27T17:43:20.443309Z","shell.execute_reply.started":"2025-06-27T17:43:20.433035Z","shell.execute_reply":"2025-06-27T17:43:20.442665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ds = get_dataset(train_models_list, batch_size, shuffle=20000, split = 'train', to_repeat = True,\n                       drop_reminder = True, cache = False)\nval_ds_2 = get_dataset(val_models_list, batch_size, shuffle=False, split = 'valid', to_repeat = False,\n                       drop_reminder = True, cache = True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:20.444074Z","iopub.execute_input":"2025-06-27T17:43:20.444308Z","iopub.status.idle":"2025-06-27T17:43:38.785428Z","shell.execute_reply.started":"2025-06-27T17:43:20.444285Z","shell.execute_reply":"2025-06-27T17:43:38.784229Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"laa = [x for x in val_ds]\nlaa = [x[1][:,0] for x in laa]\nlaa = tf.concat(laa, axis = 0)\ngaa = (np.concatenate([x[:,0,:,:] for x in val_models_list], axis = 0)-3000)/100\nprint(np.mean((gaa==laa)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:38.786723Z","iopub.execute_input":"2025-06-27T17:43:38.787080Z","iopub.status.idle":"2025-06-27T17:43:40.019401Z","shell.execute_reply.started":"2025-06-27T17:43:38.787038Z","shell.execute_reply":"2025-06-27T17:43:40.018274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"laa = [x for x in train_ds.take(2)]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:40.020714Z","iopub.execute_input":"2025-06-27T17:43:40.021045Z","iopub.status.idle":"2025-06-27T17:43:55.970313Z","shell.execute_reply.started":"2025-06-27T17:43:40.020993Z","shell.execute_reply":"2025-06-27T17:43:55.969184Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(laa[0][1][0,0])\nplt.show()\nplt.imshow(laa[0][0][0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:55.972070Z","iopub.execute_input":"2025-06-27T17:43:55.972395Z","iopub.status.idle":"2025-06-27T17:43:56.252712Z","shell.execute_reply.started":"2025-06-27T17:43:55.972366Z","shell.execute_reply":"2025-06-27T17:43:56.251765Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch = [x for x in val_ds.take(1)][0]\nprint(batch[0].shape)\nprint(batch[1].shape)","metadata":{"executionInfo":{"elapsed":21,"status":"ok","timestamp":1750539452119,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"Jszy-eD_QbV7","outputId":"a29e1859-06f5-4382-9cc7-8ed7e282a895","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:43:56.253810Z","iopub.execute_input":"2025-06-27T17:43:56.254130Z","iopub.status.idle":"2025-06-27T17:44:00.397528Z","shell.execute_reply.started":"2025-06-27T17:43:56.254098Z","shell.execute_reply":"2025-06-27T17:44:00.396434Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(np.asarray(batch[0][0]).astype(np.float32))","metadata":{"executionInfo":{"elapsed":107,"status":"ok","timestamp":1750539452231,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"4df5LIaNOW_p","outputId":"edb7f97e-d516-472d-877f-c7066321e680","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:44:00.398819Z","iopub.execute_input":"2025-06-27T17:44:00.399283Z","iopub.status.idle":"2025-06-27T17:44:01.824486Z","shell.execute_reply.started":"2025-06-27T17:44:00.399250Z","shell.execute_reply":"2025-06-27T17:44:01.823445Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Ops","metadata":{"id":"JLBF2fZbN1xy"}},{"cell_type":"code","source":"import functools\nimport math\nfrom typing import Optional\n\nfrom absl import logging\nimport numpy as np\nimport tensorflow as tf\nimport tensorflow.keras as tf_keras\n\n\ndef activation_fn(features: tf.Tensor, act_fn: str):\n  \"\"\"Customized non-linear activation type.\"\"\"\n  if act_fn in ('silu', 'swish'):\n    return tf.nn.swish(features)\n  elif act_fn == 'silu_native':\n    return features * tf.sigmoid(features)\n  elif act_fn == 'hswish':\n    return features * tf.nn.relu6(features + 3) / 6\n  elif act_fn == 'relu':\n    return tf.nn.relu(features)\n  elif act_fn == 'relu6':\n    return tf.nn.relu6(features)\n  elif act_fn == 'elu':\n    return tf.nn.elu(features)\n  elif act_fn == 'leaky_relu':\n    return tf.nn.leaky_relu(features)\n  elif act_fn == 'selu':\n    return tf.nn.selu(features)\n  elif act_fn == 'mish':\n    return features * tf.math.tanh(tf.math.softplus(features))\n  elif act_fn == 'gelu':\n    return (\n        0.5\n        * features\n        * (\n            1\n            + tf.tanh(\n                np.sqrt(2 / np.pi) * (features + 0.044715 * tf.pow(features, 3))\n            )\n        )\n    )\n  else:\n    raise ValueError('Unsupported act_fn {}'.format(act_fn))\n\n\ndef get_act_fn(act_fn):\n  if act_fn is None:\n    act_fn = 'gelu'\n  if isinstance(act_fn, str):\n    return functools.partial(activation_fn, act_fn=act_fn)\n  elif callable(act_fn):\n    return act_fn\n  else:\n    raise ValueError('Unsupported act_fn %s.' % act_fn)\n\n\ndef pooling_2d(inputs, pool_type, stride, **kwargs):\n  \"\"\"Perform 2D pooling.\"\"\"\n  if stride > 1:\n    if pool_type == 'max':\n      pool_op = tf_keras.layers.MaxPool2D\n    elif pool_type == 'avg':\n      pool_op = tf_keras.layers.AveragePooling2D\n    else:\n      raise ValueError('Unsurpported pool_type %s' % pool_type)\n    output = pool_op(\n        pool_size=(stride, stride), strides=(stride, stride), **kwargs\n    )(inputs)\n  else:\n    output = inputs\n  return output\n\n\ndef drop_connect(inputs, training, survival_prob):\n  \"\"\"Drop the entire conv with given survival probability.\"\"\"\n  # \"Deep Networks with Stochastic Depth\", https://arxiv.org/pdf/1603.09382.pdf\n  if not training:\n    return inputs\n\n  # Compute tensor.\n  batch_size = tf.shape(inputs)[0]\n  random_tensor = survival_prob\n  random_tensor += tf.random.uniform([batch_size], dtype=inputs.dtype)\n  for _ in range(inputs.shape.rank - 1):\n    random_tensor = tf.expand_dims(random_tensor, axis=-1)\n  binary_tensor = tf.floor(random_tensor)\n  # Unlike conventional way that multiply survival_prob at test time, here we\n  # divide survival_prob at training time, such that no addition compute is\n  # needed at test time.\n  output = inputs / survival_prob * binary_tensor\n  return output\n\n\ndef residual_add(residual, shortcut, survival_prob, training):\n  \"\"\"Combine residual and shortcut.\"\"\"\n  if survival_prob is not None and 0 < survival_prob < 1:\n    residual = drop_connect(residual, training, survival_prob)\n  return shortcut + residual\n\n\ndef maybe_reshape_to_2d(x, height=None):\n  \"\"\"Reshape tensor to 2d if not already 2d.\"\"\"\n  if x.shape.rank == 3:\n    _, length, num_channel = x.shape.as_list()\n    if height is None:\n      height = int(np.sqrt(length))\n    else:\n      assert length % height == 0\n    width = length // height\n    logging.debug(\n        'Reshape %s -> %s', [length, num_channel], [height, width, num_channel]\n    )\n    return tf.reshape(x, [-1, height, width, num_channel])\n  elif x.shape.rank == 4:\n    return x\n  else:\n    raise ValueError('Unsupport shape {}'.format(x.shape))\n\n\ndef maybe_reshape_to_1d(x):\n  \"\"\"Reshape tensor to 1d if not already 1d.\"\"\"\n  if x.shape.rank == 4:\n    _, h, w, num_channel = x.shape.as_list()\n    logging.debug('Reshape %s -> %s', [h, w, num_channel], [h * w, num_channel])\n    return tf.reshape(x, [-1, h * w, num_channel])\n  elif x.shape.rank == 3:\n    return x\n  else:\n    raise ValueError('Unsupport shape {}'.format(x.shape))\n\n\ndef generate_lookup_tensor(\n    length: int,\n    max_relative_position: Optional[int] = None,\n    clamp_out_of_range: bool = False,\n    dtype: tf.DType = tf.float32) -> tf.Tensor:\n  \"\"\"Generate a one_hot lookup tensor to reindex embeddings along one dimension.\n\n  Args:\n    length: the length to reindex to.\n    max_relative_position: the maximum relative position to consider.\n      Relative position embeddings for distances above this threshold\n      are zeroed out.\n    clamp_out_of_range: bool. Whether to clamp out of range locations to the\n      maximum relative distance. If False, the out of range locations will be\n      filled with all-zero vectors.\n    dtype: dtype for the returned lookup tensor.\n  Returns:\n    ret: [length, length, vocab_size] lookup tensor that satisfies\n      ret[n,m,v] = 1{m - n + max_relative_position = v}.\n  \"\"\"\n  if max_relative_position is None:\n    max_relative_position = length - 1\n  vocab_size = 2 * max_relative_position + 1\n  ret = np.zeros((length, length, vocab_size))\n  for i in range(length):\n    for x in range(length):\n      v = x - i + max_relative_position\n      if abs(x - i) > max_relative_position:\n        if clamp_out_of_range:\n          v = np.clip(v, 0, vocab_size - 1)\n        else:\n          continue\n      ret[i, x, v] = 1\n  return tf.constant(ret, dtype)\n\n\ndef reindex_2d_einsum_lookup(\n    relative_position_tensor: tf.Tensor,\n    height: int,\n    width: int,\n    max_relative_height: Optional[int] = None,\n    max_relative_width: Optional[int] = None,\n    h_axis=None) -> tf.Tensor:\n  \"\"\"Reindex 2d relative position bias with 2 independent einsum lookups.\n\n  Args:\n    relative_position_tensor: tensor of shape\n      [..., vocab_height, vocab_width, ...].\n    height: height to reindex to.\n    width: width to reindex to.\n    max_relative_height: maximum relative height.\n      Position embeddings corresponding to vertical distances larger\n      than max_relative_height are zeroed out. None to disable.\n    max_relative_width: maximum relative width.\n      Position embeddings corresponding to horizontal distances larger\n      than max_relative_width are zeroed out. None to disable.\n    h_axis: Axis corresponding to vocab_height. Default to 0 if None.\n\n  Returns:\n    reindexed_bias: a Tensor of shape\n      [..., height * width, height * width, ...]\n  \"\"\"\n  height_lookup = generate_lookup_tensor(\n      height, max_relative_position=max_relative_height,\n      dtype=relative_position_tensor.dtype)\n  width_lookup = generate_lookup_tensor(\n      width, max_relative_position=max_relative_width,\n      dtype=relative_position_tensor.dtype)\n\n  if h_axis is None:\n    h_axis = 0\n\n  non_spatial_rank = relative_position_tensor.shape.rank - 2\n  non_spatial_expr = ''.join(chr(ord('n') + i) for i in range(non_spatial_rank))\n  prefix = non_spatial_expr[:h_axis]\n  suffix = non_spatial_expr[h_axis:]\n\n  reindexed_tensor = tf.einsum(\n      '{0}hw{1},ixh->{0}ixw{1}'.format(prefix, suffix),\n      relative_position_tensor, height_lookup, name='height_lookup')\n  reindexed_tensor = tf.einsum(\n      '{0}ixw{1},jyw->{0}ijxy{1}'.format(prefix, suffix),\n      reindexed_tensor, width_lookup, name='width_lookup')\n\n  ret_shape = relative_position_tensor.shape.as_list()\n  ret_shape[h_axis] = height * width\n  ret_shape[h_axis + 1] = height * width\n  reindexed_tensor = tf.reshape(reindexed_tensor, ret_shape)\n\n  return reindexed_tensor\n\n\ndef float32_softmax(x: tf.Tensor, *args, **kwargs) -> tf.Tensor:\n  y = tf.cast(tf.nn.softmax(tf.cast(x, tf.float32), *args, **kwargs), x.dtype)\n  return y\n\n\ndef get_shape_from_length(length: int, height: int = 1, width: int = 1):\n  \"\"\"Gets input 2D shape from 1D sequence length.\"\"\"\n  input_height = int(math.sqrt(length * height // width))\n  input_width = input_height * width // height\n  if input_height * input_width != length:\n    raise ValueError(\n        f'Invalid sequence length: {length} or shape: ({height, width}).'\n    )\n  return (input_height, input_width)\n\n\ndef absolute_position_encoding(\n    position: tf.Tensor, hidden_size: int, dtype=tf.float32) -> tf.Tensor:\n  \"\"\"Create absoulte position encoding.\"\"\"\n  position = tf.cast(position, dtype)\n  half_hid = hidden_size // 2\n  freq_seq = tf.cast(tf.range(half_hid), dtype=dtype)\n  inv_freq = 1 / (10000 ** (freq_seq / half_hid))\n  sinusoid = tf.einsum('S,D->SD', position, inv_freq)\n  sin = tf.sin(sinusoid)\n  cos = tf.cos(sinusoid)\n  return tf.concat([sin, cos], axis=-1)","metadata":{"executionInfo":{"elapsed":137,"status":"ok","timestamp":1750539452371,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"Sr81PLPgN23v","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:44:01.825806Z","iopub.execute_input":"2025-06-27T17:44:01.826111Z","iopub.status.idle":"2025-06-27T17:44:01.872893Z","shell.execute_reply.started":"2025-06-27T17:44:01.826083Z","shell.execute_reply":"2025-06-27T17:44:01.871897Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Layers","metadata":{"id":"J8e-G13jN3Yx"}},{"cell_type":"code","source":"!pip install einops\n\nimport math\n\nimport six\nfrom einops.layers.tensorflow import Rearrange\nimport tensorflow as tf\nfrom tensorflow.keras.callbacks import TensorBoard","metadata":{"executionInfo":{"elapsed":1876,"status":"ok","timestamp":1750539454252,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"wJiD9FKZPxBC","outputId":"084b4879-8c08-4b64-dcc2-1b045b46b541","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:44:01.874093Z","iopub.execute_input":"2025-06-27T17:44:01.874417Z","iopub.status.idle":"2025-06-27T17:44:06.028205Z","shell.execute_reply.started":"2025-06-27T17:44:01.874388Z","shell.execute_reply":"2025-06-27T17:44:06.026845Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class GLU(tf.keras.layers.Layer):\n    def __init__(self, **kwargs):\n        super().__init__(**kwargs)\n    def call(self, x, mask=None):\n        x,gate = tf.split(x, 2, axis = -1)\n        x = x*tf.keras.activations.swish(gate)\n        return x\n\nclass GLUMlp(tf.keras.layers.Layer):\n    def __init__(self, dim_expand, dim, **kwargs):\n        super().__init__(**kwargs)\n        self.dim_expand = dim_expand\n        self.dim = dim\n        self.dense_1 = tf.keras.layers.EinsumDense(\"abc,cd->abd\",output_shape=(None, self.dim_expand), activation = 'linear', bias_axes = 'd')\n        self.glu_1 = GLU()\n        self.dense_2 = tf.keras.layers.EinsumDense(\"abc,cd->abd\",output_shape=(None, self.dim), activation = 'linear', bias_axes = 'd')\n    def call(self, x, training = False):\n        x = self.dense_1(x)\n        x = self.glu_1(x)\n        x = self.dense_2(x)\n        return x\n\nclass Attention_X(tf.keras.Model):\n    def __init__(self, dim, heads = 8):\n        super().__init__()\n        self.heads = heads\n        self.scale = dim ** -0.5\n\n        self.to_qkv = tf.keras.layers.Dense(dim * 3, use_bias=False)\n        self.to_out = tf.keras.layers.Dense(dim)\n\n        self.rearrange_qkv = Rearrange('b n (qkv h d) -> qkv b h n d', qkv = 3, h = self.heads)\n        self.rearrange_out = Rearrange('b h n d -> b n (h d)')\n\n    def call(self, x):\n        qkv = self.to_qkv(x)\n        qkv = self.rearrange_qkv(qkv)\n        q = qkv[0]\n        k = qkv[1]\n        v = qkv[2]\n\n        dots = tf.einsum('bhid,bhjd->bhij', q, k) * self.scale\n        attn = tf.keras.activations.softmax(dots,axis=-1)\n\n        out = tf.einsum('bhij,bhjd->bhid', attn, v)\n        out = self.rearrange_out(out)\n        out =  self.to_out(out)\n        return out\n\nclass Transformer_X(tf.keras.layers.Layer):\n    def __init__(self, dim, heads, mlp_dim):\n        super().__init__()\n        self.att =  Attention_X(dim, heads = heads)\n        self.ffn = GLUMlp(mlp_dim, dim)\n        self.layer_norm_1 = tf.keras.layers.LayerNormalization(epsilon=1e-5)\n        self.layer_norm_2 = tf.keras.layers.LayerNormalization(epsilon=1e-5)\n    def call(self, x):\n        residual = x\n        x = self.layer_norm_1(x)\n        x = self.att(x)\n        x = x+residual\n        residual = x\n        x = self.layer_norm_2(x)\n        x = self.ffn(x)\n        x = x+residual\n        return x\n\nclass Pos_embedding_layer(tf.keras.layers.Layer):\n    def __init__(self, num_patches, dim, **kwargs):\n        super().__init__(**kwargs)\n        self.pos_embedding = self.add_weight(name = \"position_embeddings\",\n                                             shape=(num_patches,dim),\n                                             initializer=tf.keras.initializers.RandomNormal(),\n                                             dtype=tf.float32)\n    def call(self, x):\n        x += self.pos_embedding\n        return x","metadata":{"executionInfo":{"elapsed":4,"status":"ok","timestamp":1750539454259,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"SGAVLfUJPZJ4","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:44:06.029714Z","iopub.execute_input":"2025-06-27T17:44:06.030041Z","iopub.status.idle":"2025-06-27T17:44:06.046035Z","shell.execute_reply.started":"2025-06-27T17:44:06.029997Z","shell.execute_reply":"2025-06-27T17:44:06.045198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import functools\nimport string\nfrom typing import Any, Callable, Optional, Tuple, Union\n\nfrom absl import logging\n\n\nclass TrailDense(tf_keras.layers.Layer):\n  \"\"\"Dense module that projects multiple trailing dimensions.\"\"\"\n\n  def __init__(\n      self,\n      output_trailing_dims: Union[int, Tuple[int, ...]],\n      begin_axis: int = -1,\n      use_bias: bool = True,\n      kernel_initializer: Optional[str] = 'glorot_uniform',\n      bias_initializer: Optional[str] = 'zeros',\n      name: str = 'dense',\n  ):\n    super().__init__(name=name)\n\n    if isinstance(output_trailing_dims, int):\n      self._output_trailing_dims = [output_trailing_dims]\n    else:\n      assert isinstance(output_trailing_dims, (list, tuple)) and all(\n          isinstance(i, int) for i in output_trailing_dims\n      ), f'Invalid output shape: {output_trailing_dims}.'\n      self._output_trailing_dims = list(output_trailing_dims)\n    self.begin_axis = begin_axis\n    self.use_bias = use_bias\n\n    self.kernel_initializer = kernel_initializer\n    self.bias_initializer = bias_initializer\n\n  def build(self, input_shape: tf.TensorShape) -> None:\n    \"\"\"Create variables and einsum expression based on input shape.\"\"\"\n    # Create variables\n    weight_shape = input_shape[self.begin_axis :] + self._output_trailing_dims\n    self.weight = self.add_weight(\n        name='weight',\n        shape=weight_shape,\n        initializer=self.kernel_initializer,\n        trainable=True,\n    )\n    if self.use_bias:\n      self.bias = self.add_weight(\n          name='bias',\n          shape=self._output_trailing_dims,\n          initializer=self.bias_initializer,\n          trainable=True,\n      )\n\n    # Create einsum expression\n    input_rank = input_shape.rank\n    shared_size = self.begin_axis % input_rank\n    i_only_size = input_rank - shared_size\n    o_only_size = len(self._output_trailing_dims)\n\n    assert input_rank + o_only_size < len(\n        string.ascii_uppercase\n    ), 'Cannot use einsum as input rank + output rank > 26.'\n    einsum_str = string.ascii_uppercase[: input_rank + o_only_size]\n\n    offset = 0\n    shared_str = einsum_str[offset : offset + shared_size]\n    offset += shared_size\n    i_only_str = einsum_str[offset : offset + i_only_size]\n    offset += i_only_size\n    o_only_str = einsum_str[offset : offset + o_only_size]\n\n    input_str = f'{shared_str}{i_only_str}'\n    output_str = f'{shared_str}{o_only_str}'\n    weight_str = f'{i_only_str}{o_only_str}'\n    # Examples\n    # - For 4D tensors in conv, a common expr would be 'ABCD,DE->ABCE'.\n    # - For `q/k/v` head projection in multi-head attention with two output\n    #   trailing dims, the expr is 'ABC,CDE->ABDE'\n    # - For `o` output projection in multi-head attention with begin_axis = -2,\n    #   the expr is 'ABCD,CDE->ABE'\n    self.einsum_expr = f'{input_str},{weight_str}->{output_str}'\n\n  def call(self, inputs: tf.Tensor) -> tf.Tensor:\n    output = tf.einsum(self.einsum_expr, inputs, self.weight)\n    if self.use_bias:\n      output += self.bias\n    return output\n\n\nclass Attention(tf_keras.layers.Layer):\n  \"\"\"Multi-headed attention module.\"\"\"\n\n  def __init__(\n      self,\n      hidden_size: int,\n      head_size: int,\n      input_origin_height: int = 1,\n      input_origin_width: int = 1,\n      num_heads: Optional[int] = None,\n      dropatt: float = 0.0,\n      attn_axis: int = 0,\n      rel_attn_type: Optional[str] = None,\n      scale_ratio: Optional[float] = None,\n      kernel_initializer: Optional[str] = 'glorot_uniform',\n      bias_initializer: Optional[str] = 'zeros',\n      name: str = 'attention',\n  ):\n    super().__init__(name=name)\n\n    self.hidden_size = hidden_size\n    self.head_size = head_size\n    self.input_origin_height = input_origin_height\n    self.input_origin_width = input_origin_width\n    self.num_heads = num_heads or hidden_size // head_size\n    self.dropatt = dropatt\n    self.attn_axis = attn_axis\n    self.rel_attn_type = rel_attn_type\n    self.scale_ratio = scale_ratio\n\n    self.kernel_initializer = kernel_initializer\n    self.bias_initializer = bias_initializer\n\n    self._q_proj = TrailDense(\n        output_trailing_dims=(self.num_heads, self.head_size),\n        kernel_initializer=kernel_initializer,\n        bias_initializer=bias_initializer,\n        name='q',\n    )\n    self._k_proj = TrailDense(\n        output_trailing_dims=(self.num_heads, self.head_size),\n        kernel_initializer=kernel_initializer,\n        bias_initializer=bias_initializer,\n        name='k',\n    )\n    self._v_proj = TrailDense(\n        output_trailing_dims=(self.num_heads, self.head_size),\n        kernel_initializer=kernel_initializer,\n        bias_initializer=bias_initializer,\n        name='v',\n    )\n    self._o_proj = TrailDense(\n        output_trailing_dims=self.hidden_size,\n        begin_axis=-2,\n        kernel_initializer=kernel_initializer,\n        bias_initializer=bias_initializer,\n        name='o',\n    )\n\n    self.q_scale = self.head_size**-0.5\n    self.relative_bias = None\n\n  def build(self, query_shape: Any) -> None:\n    ##### Content attention\n    # Einsum expression:\n    #   B = batch_size\n    #   N = num_heads\n    #   K = head_size\n    #   S = query_len (of the given attn_axis)\n    #   T = key/value_len (of the given attn_axis)\n    #   [U-Z] = length of other attension axes\n    # Example for 5D query_heads, (e.g. images [B x H x W x N x K])\n    # - when attn_axis = 0 (H axis):\n    #     symbols = 'U'  => num_attn_dims = 2\n    #     q_expr = 'BSUNK' => 'S' is inserted, prefix = 'B', suffix = 'NK'\n    #     k_expr = 'BTUNK' => 'T' is inserted, prefix = 'B', suffix = 'NK'\n    #     v_expr = 'BTUNK' => 'T' is inserted, prefix = 'B', suffix = 'NK'\n    #     a_expr = 'BUNST' => 'N x S x T' attention map\n    num_attn_dims = query_shape.rank - 2  # -2 to account for bsz, hidden size\n    assert num_attn_dims < 6, 'Only support at most 6 attention dims.'\n    symbols = ''.join([chr(ord('U') + i) for i in range(num_attn_dims - 1)])\n    insert = lambda s, i, c: s[:i] + c + s[i:]\n    create_expr = lambda s, prefix='B', suffix='NK': prefix + s + suffix\n    self.q_expr = create_expr(insert(symbols, self.attn_axis, 'S'))\n    self.k_expr = create_expr(insert(symbols, self.attn_axis, 'T'))\n    self.v_expr = create_expr(insert(symbols, self.attn_axis, 'T'))\n    self.a_expr = create_expr(symbols, suffix='NST')\n\n    ##### Relative attention\n    if self.rel_attn_type in ['2d_multi_head', '2d_single_head']:\n      query_shape_list = query_shape.as_list()\n      if query_shape.rank == 4:\n        height, width = query_shape_list[1:3]\n      elif query_shape.rank == 3:\n        seq_len = query_shape_list[1]\n        height, width = get_shape_from_length(\n            seq_len, self.input_origin_height, self.input_origin_width\n        )\n        if height * width != seq_len:\n          raise ValueError(\n              'Sequence length: %s violates input size: (%s, %s).'\n              % (seq_len, height, width)\n          )\n      else:\n        raise ValueError(\n            'Does not support relative attention for query shape: %s.'\n            % query_shape_list\n        )\n\n      if self.scale_ratio is not None:\n        scale_ratio = eval(self.scale_ratio)  # pylint:disable=eval-used\n        vocab_height = 2 * int(height / scale_ratio) - 1\n        vocab_width = 2 * int(width / scale_ratio) - 1\n      else:\n        vocab_height = 2 * height - 1\n        vocab_width = 2 * width - 1\n\n      if self.rel_attn_type == '2d_multi_head':\n        rel_bias_shape = [self.num_heads, vocab_height, vocab_width]\n      elif self.rel_attn_type == '2d_single_head':\n        rel_bias_shape = [vocab_height, vocab_width]\n      else:\n        raise NotImplementedError(\n            f'rel_attn_type {self.rel_attn_type} not implemented yet.'\n        )\n\n      self._feat_height = height\n      self._feat_width = width\n      self.relative_bias = self.add_weight(\n          'relative_bias',\n          rel_bias_shape,\n          initializer=self.kernel_initializer,\n          trainable=True,\n      )\n\n  def call(\n      self,\n      query: tf.Tensor,\n      training: bool,\n      context: Optional[tf.Tensor] = None,\n      attn_mask: Optional[tf.Tensor] = None,\n  ) -> tf.Tensor:\n    if context is None:\n      context = query\n\n    q_heads = self._q_proj(query)\n    k_heads = self._k_proj(context)\n    v_heads = self._v_proj(context)\n    q_heads *= self.q_scale\n\n    # attention\n    attn_logits = tf.einsum(\n        f'{self.q_expr},{self.k_expr}->{self.a_expr}', q_heads, k_heads\n    )\n\n    if self.relative_bias is not None:\n      if self.rel_attn_type == '2d_multi_head':\n        h_axis = 1\n      else:\n        h_axis = 0\n\n      if self.scale_ratio is not None:\n        src_shape = self.relative_bias.shape.as_list()\n        relative_bias = tf.expand_dims(self.relative_bias, axis=-1)\n        relative_bias = tf.image.resize(\n            relative_bias, [2 * self._feat_height - 1, 2 * self._feat_width - 1]\n        )\n        relative_bias = tf.cast(\n            tf.squeeze(relative_bias, axis=-1), self.compute_dtype\n        )\n        tgt_shape = relative_bias.shape.as_list()\n        logging.info(\n            'Bilinear resize relative position bias %s -> %s.',\n            src_shape,\n            tgt_shape,\n        )\n      else:\n        relative_bias = tf.cast(self.relative_bias, self.compute_dtype)\n\n      reindexed_bias = reindex_2d_einsum_lookup(\n          relative_position_tensor=relative_bias,\n          height=self._feat_height,\n          width=self._feat_width,\n          max_relative_height=self._feat_height - 1,\n          max_relative_width=self._feat_width - 1,\n          h_axis=h_axis,\n      )\n      attn_logits += reindexed_bias\n\n    if attn_mask is not None:\n      # attn_mask: 1.0 means CAN attend, 0.0 means CANNOT attend\n      attn_logits += (1.0 - attn_mask) * attn_logits.dtype.min\n\n    attn_probs = float32_softmax(attn_logits, axis=-1)\n    if self.dropatt:\n      attn_probs = tf_keras.layers.Dropout(self.dropatt, name='attn_prob_drop')(\n          attn_probs, training=training\n      )\n\n    attn_out = tf.einsum(\n        f'{self.a_expr},{self.v_expr}->{self.q_expr}', attn_probs, v_heads\n    )\n    output = self._o_proj(attn_out)\n\n    return output\n\n\nclass FFN(tf_keras.layers.Layer):\n  \"\"\"Positionwise feed-forward network.\"\"\"\n\n  def __init__(\n      self,\n      hidden_size: int,\n      dropout: float = 0.0,\n      expansion_rate: int = 4,\n      activation: str = 'gelu',\n      kernel_initializer: Optional[str] = 'glorot_uniform',\n      bias_initializer: Optional[str] = 'zeros',\n      name: str = 'ffn',\n  ):\n    super().__init__(name=name)\n\n    self.hidden_size = hidden_size\n    self.expansion_rate = expansion_rate\n    self.expanded_size = self.hidden_size * self.expansion_rate\n    self.dropout = dropout\n    self.activation = activation\n\n    self._expand_dense = TrailDense(\n        output_trailing_dims=self.expanded_size,\n        kernel_initializer=kernel_initializer,\n        bias_initializer=bias_initializer,\n        name='expand_dense',\n    )\n    self._shrink_dense = TrailDense(\n        output_trailing_dims=self.hidden_size,\n        kernel_initializer=kernel_initializer,\n        bias_initializer=bias_initializer,\n        name='shrink_dense',\n    )\n    self._activation_fn = get_act_fn(self.activation)\n\n  def call(self, inputs: tf.Tensor, training: bool) -> tf.Tensor:\n    output = inputs\n    output = self._expand_dense(output)\n    output = self._activation_fn(output)\n    if self.dropout:\n      output = tf_keras.layers.Dropout(self.dropout, name='nonlinearity_drop')(\n          output, training=training\n      )\n    output = self._shrink_dense(output)\n\n    return output\n\n\nclass TransformerBlock(tf_keras.layers.Layer):\n  \"\"\"Transformer block = Attention + FFN.\"\"\"\n\n  def __init__(\n      self,\n      hidden_size: int,\n      head_size: int,\n      input_origin_height: int = 1,\n      input_origin_width: int = 1,\n      num_heads: Optional[int] = None,\n      expansion_rate: int = 4,\n      activation: str = 'gelu',\n      pool_type: str = '2d:avg',\n      pool_stride: int = 1,\n      pool_query_only: bool = False,\n      dropatt: Optional[Union[float, tf.Tensor]] = None,\n      dropout: Optional[Union[float, tf.Tensor]] = None,\n      rel_attn_type: Optional[str] = None,\n      scale_ratio: Optional[str] = None,\n      survival_prob: Optional[Union[float, tf.Tensor]] = None,\n      ln_epsilon: float = 1e-5,\n      ln_dtype: Optional[tf.DType] = None,\n      kernel_initializer: Optional[str] = 'glorot_uniform',\n      bias_initializer: Optional[str] = 'zeros',\n      name: str = 'transformer',\n  ) -> None:\n    super().__init__(name=name)\n\n    self._hidden_size = hidden_size\n    self._head_size = head_size\n    self._input_origin_height = input_origin_height\n    self._input_origin_width = input_origin_width\n    self._num_heads = num_heads\n    self._expansion_rate = expansion_rate\n    self._activation = activation\n    self._pool_type = pool_type\n    self._pool_stride = pool_stride\n    self._pool_query_only = pool_query_only\n    self._dropatt = dropatt\n    self._dropout = dropout\n    self._rel_attn_type = rel_attn_type\n    self._scale_ratio = scale_ratio\n    self._survival_prob = survival_prob\n    self._ln_epsilon = ln_epsilon\n    self._ln_dtype = ln_dtype\n    self._kernel_initializer = kernel_initializer\n    self._bias_initializer = bias_initializer\n\n  def build(self, input_shape: tf.TensorShape) -> None:\n    if len(input_shape.as_list()) == 4:\n      _, height, width, _ = input_shape.as_list()\n    elif len(input_shape.as_list()) == 3:\n      _, seq_len, _ = input_shape.as_list()\n      height, width = get_shape_from_length(\n          seq_len, self._input_origin_height, self._input_origin_width\n      )\n    else:\n      raise ValueError(f'Unsupported input shape: {input_shape.as_list()}.')\n\n    self.height, self.width = height, width\n    input_size = input_shape.as_list()[-1]\n\n    if input_size != self._hidden_size:\n      self._shortcut_proj = TrailDense(\n          self._hidden_size,\n          kernel_initializer=self._kernel_initializer,\n          bias_initializer=self._bias_initializer,\n          name='shortcut_proj',\n      )\n    else:\n      self._shortcut_proj = None\n\n    self._attn_layer_norm = tf_keras.layers.LayerNormalization(\n        axis=-1,\n        epsilon=self._ln_epsilon,\n        dtype=self._ln_dtype,\n        name='attn_layer_norm',\n    )\n\n    self._attention = Attention(\n        self._hidden_size,\n        self._head_size,\n        height // self._pool_stride,\n        width // self._pool_stride,\n        num_heads=self._num_heads,\n        dropatt=self._dropatt,\n        rel_attn_type=self._rel_attn_type,\n        scale_ratio=self._scale_ratio,\n        kernel_initializer=self._kernel_initializer,\n        bias_initializer=self._bias_initializer,\n    )\n\n    self._ffn_layer_norm = tf_keras.layers.LayerNormalization(\n        axis=-1,\n        epsilon=self._ln_epsilon,\n        dtype=self._ln_dtype,\n        name='ffn_layer_norm',\n    )\n\n    self._ffn = FFN(\n        self._hidden_size,\n        dropout=self._dropout,\n        expansion_rate=self._expansion_rate,\n        activation=self._activation,\n        kernel_initializer=self._kernel_initializer,\n        bias_initializer=self._bias_initializer,\n    )\n\n  def downsample(self, inputs: tf.Tensor, name: str) -> tf.Tensor:\n    output = inputs\n    if self._pool_stride > 1:\n      assert self._pool_type in [\n          '2d:avg',\n          '2d:max',\n          '1d:avg',\n          '1d:max',\n      ], f'Invalid pool_type {self._pool_type}'\n      if self._pool_type.startswith('2d'):\n        output = maybe_reshape_to_2d(output, height=self.height)\n        output = pooling_2d(\n            output,\n            self._pool_type.split(':')[-1],\n            self._pool_stride,\n            padding='same',\n            data_format='channels_last',\n            name=name,\n        )\n      else:\n        output = pooling_1d(\n            output,\n            self._pool_type.split(':')[-1],\n            self._pool_stride,\n            padding='same',\n            data_format='channels_last',\n            name=name,\n        )\n    return output\n\n  def shortcut_branch(self, shortcut: tf.Tensor) -> tf.Tensor:\n    shortcut = self.downsample(shortcut, 'shortcut_pool')\n    shortcut = maybe_reshape_to_1d(shortcut)\n    if self._shortcut_proj:\n      shortcut = self._shortcut_proj(shortcut)\n\n    return shortcut\n\n  def attn_branch(\n      self,\n      inputs: tf.Tensor,\n      training: bool,\n      attn_mask: Optional[tf.Tensor] = None,\n  ) -> tf.Tensor:\n    output = self._attn_layer_norm(inputs)\n    if self._pool_query_only:\n      query = self.downsample(output, 'query_pool')\n      query = maybe_reshape_to_1d(query)\n      output = maybe_reshape_to_1d(output)\n      output = self._attention(\n          query, training, context=output, attn_mask=attn_mask\n      )\n    else:\n      output = self.downsample(output, 'residual_pool')\n      output = maybe_reshape_to_1d(output)\n      output = self._attention(output, training, attn_mask=attn_mask)\n    return output\n\n  def ffn_branch(self, inputs: tf.Tensor, training: bool) -> tf.Tensor:\n    output = self._ffn_layer_norm(inputs)\n    output = self._ffn(output, training)\n    return output\n\n  def call(\n      self,\n      inputs: tf.Tensor,\n      training: bool,\n      attn_mask: Optional[tf.Tensor] = None,\n  ) -> tf.Tensor:\n    logging.info(\n        'Block %s input shape: %s, (%s).', self.name, inputs.shape, inputs.dtype\n    )\n\n    shortcut = self.shortcut_branch(inputs)\n    output = self.attn_branch(inputs, training, attn_mask)\n    if self._dropout:\n      output = tf_keras.layers.Dropout(self._dropout, name='after_attn_drop')(\n          output, training=training\n      )\n    output = residual_add(\n        output, shortcut, self._survival_prob, training\n    )\n\n    shortcut = output\n    output = self.ffn_branch(output, training)\n    if self._dropout:\n      output = tf_keras.layers.Dropout(self._dropout, name='after_ffn_drop')(\n          output, training=training\n      )\n    output = residual_add(\n        output, shortcut, self._survival_prob, training\n    )\n\n    return output\n\n\nclass SqueezeAndExcitation(tf_keras.layers.Layer):\n  \"\"\"Squeeze-and-excitation layer.\"\"\"\n\n  def __init__(\n      self,\n      se_filters: int,\n      output_filters: int,\n      local_pooling: bool = False,\n      data_format: str = 'channels_last',\n      activation: str = 'swish',\n      kernel_initializer: Optional[str] = 'glorot_uniform',\n      bias_initializer: Optional[str] = 'zeros',\n      name: str = 'se',\n  ):\n    super().__init__(name=name)\n\n    self._local_pooling = local_pooling\n    self._data_format = data_format\n    self._activation_fn = get_act_fn(activation)\n\n    # Squeeze and Excitation layer.\n    self._se_reduce = tf_keras.layers.Conv2D(\n        se_filters,\n        kernel_size=[1, 1],\n        strides=[1, 1],\n        padding='same',\n        data_format=self._data_format,\n        use_bias=True,\n        kernel_initializer=kernel_initializer,\n        bias_initializer=bias_initializer,\n        name='reduce_conv2d',\n    )\n    self._se_expand = tf_keras.layers.Conv2D(\n        output_filters,\n        kernel_size=[1, 1],\n        strides=[1, 1],\n        padding='same',\n        data_format=self._data_format,\n        use_bias=True,\n        kernel_initializer=kernel_initializer,\n        bias_initializer=bias_initializer,\n        name='expand_conv2d',\n    )\n\n  def call(self, inputs: tf.Tensor) -> tf.Tensor:\n    h_axis, w_axis = [2, 3] if self._data_format == 'channels_first' else [1, 2]\n    if self._local_pooling:\n      se_tensor = tf.nn.avg_pool(\n          inputs,\n          ksize=[1, inputs.shape[h_axis], inputs.shape[w_axis], 1],\n          strides=[1, 1, 1, 1],\n          padding='VALID',\n      )\n    else:\n      se_tensor = tf.reduce_mean(inputs, [h_axis, w_axis], keepdims=True)\n    se_tensor = self._se_expand(self._activation_fn(self._se_reduce(se_tensor)))\n    return tf.sigmoid(se_tensor) * inputs\n\n\ndef _config_batch_norm(\n    norm_type: str,\n    ln_epsilon: float = 1e-6,\n    bn_momentum: float = 0.99,\n    bn_epsilon: float = 1e-6,\n) -> Callable[..., Any]:\n  \"\"\"Defines the normalization class for MbConv based on `norm_type`.\"\"\"\n\n  if norm_type == 'layer_norm':\n    return functools.partial(\n        tf_keras.layers.LayerNormalization, epsilon=ln_epsilon\n    )\n  elif norm_type == 'batch_norm':\n    return functools.partial(\n        tf_keras.layers.BatchNormalization,\n        momentum=bn_momentum,\n        epsilon=bn_epsilon,\n    )\n  elif norm_type == 'sync_batch_norm':\n    return functools.partial(\n        tf_keras.layers.BatchNormalization,\n        momentum=bn_momentum,\n        epsilon=bn_epsilon,\n        synchronized=True,\n    )\n  else:\n    raise ValueError(f'Unsupported norm_type {norm_type}.')\n\nclass BatchNormLayerMaybeSynched(tf_keras.layers.Layer):\n  \"\"\"Squeeze-and-excitation layer.\"\"\"\n  def __init__(\n    self,\n    norm_type: str,\n    ln_epsilon: float = 1e-6,\n    bn_momentum: float = 0.99,\n    bn_epsilon: float = 1e-6,\n  ):\n    super().__init__()\n    if norm_type == 'layer_norm':\n        self._norm_layer = tf_keras.layers.LayerNormalization(epsilon=ln_epsilon)\n    elif norm_type == 'batch_norm':\n        self._norm_layer = tf_keras.layers.BatchNormalization(momentum=bn_momentum, epsilon=bn_epsilon)\n    elif norm_type == 'sync_batch_norm':\n        self._norm_layer = tf_keras.layers.BatchNormalization(momentum=bn_momentum, epsilon=bn_epsilon, synchronized=True)\n  def call(self, x, training = None):\n    return self._norm_layer(x, training = training)\n\ndef _build_downsample_layer(\n    pool_type: str, pool_stride: int, data_format: str = 'channels_last'\n) -> tf_keras.layers.Layer:\n  \"\"\"Builds a downsample layer for MbConv based on pool type.\"\"\"\n  if pool_type == 'max':\n    return tf_keras.layers.MaxPooling2D(\n        pool_size=(pool_stride, pool_stride),\n        strides=(pool_stride, pool_stride),\n        padding='same',\n        data_format=data_format,\n    )\n  elif pool_type == 'avg':\n    return tf_keras.layers.AveragePooling2D(\n        pool_size=(pool_stride, pool_stride),\n        strides=(pool_stride, pool_stride),\n        padding='same',\n        data_format=data_format,\n    )\n  else:\n    raise ValueError(f'Unsurpported pool_type {pool_type}')\n\n\nclass MBConvBlock(tf_keras.layers.Layer):\n  \"\"\"Mobile Inverted Residual Bottleneck (https://arxiv.org/abs/1905.02244).\"\"\"\n\n  def __init__(\n      self,\n      hidden_size: int,\n      data_format: str = 'channels_last',\n      kernel_size: int = 5,\n      expansion_rate: int = 4,\n      se_ratio: float = 0.25,\n      activation: str = 'gelu',\n      norm_type: str = 'sync_batch_norm',\n      bn_epsilon: float = 1e-3,\n      bn_momentum: float = 0.99,\n      kernel_initializer: Optional[str] = 'glorot_uniform',\n      bias_initializer: Optional[str] = 'zeros',\n      name: str = 'mbconv',\n  ):\n    super().__init__(name=name)\n\n    self._hidden_size = hidden_size\n    self._data_format = data_format\n    self._kernel_size = kernel_size\n    self._expansion_rate = expansion_rate\n    self._se_ratio = se_ratio\n    self._activation = activation\n    self._norm_type = norm_type\n    self._bn_epsilon = bn_epsilon\n    self._bn_momentum = bn_momentum\n    self._kernel_initializer = kernel_initializer\n    self._bias_initializer = bias_initializer\n    self._activation_fn = get_act_fn(self._activation)\n\n  def build(self, input_shape: tf.TensorShape) -> None:\n    inner_size = self._hidden_size * self._expansion_rate\n    self._pre_norm = BatchNormLayerMaybeSynched(self._norm_type)\n\n    self._expand_conv = tf_keras.layers.Conv2D(\n        filters=inner_size,\n        kernel_size=1,\n        strides=1,\n        kernel_initializer=self._kernel_initializer,\n        padding='same',\n        data_format=self._data_format,\n        use_bias=False,\n        name='expand_conv',\n    )\n    self._expand_norm = BatchNormLayerMaybeSynched(self._norm_type)\n\n    self._depthwise_conv = tf_keras.layers.DepthwiseConv2D(\n        kernel_size=self._kernel_size,\n        strides=1,\n        depthwise_initializer=self._kernel_initializer,\n        padding='same',\n        data_format=self._data_format,\n        use_bias=False,\n        name='depthwise_conv',\n    )\n    self._depthwise_norm = BatchNormLayerMaybeSynched(self._norm_type)\n\n    se_filters = int(self._hidden_size * self._se_ratio)\n    self._se = SqueezeAndExcitation(\n        se_filters=se_filters,\n        output_filters=inner_size,\n        data_format=self._data_format,\n        kernel_initializer=self._kernel_initializer,\n        bias_initializer=self._bias_initializer,\n        name='se',\n    )\n\n    self._shrink_conv = tf_keras.layers.Conv2D(\n        filters=self._hidden_size,\n        kernel_size=1,\n        strides=1,\n        padding='same',\n        data_format=self._data_format,\n        kernel_initializer=self._kernel_initializer,\n        bias_initializer=self._bias_initializer,\n        use_bias=True,\n        name='shrink_conv',\n    )\n  def call(\n      self,\n      x: tf.Tensor,\n      training: Optional[bool] = None,\n      survival_prob: Optional[Union[float, tf.Tensor]] = None,\n  ) -> tf.Tensor:\n    shortcut = x\n\n    x = self._pre_norm(x, training=training)\n    x = self._expand_conv(x)\n    x = self._expand_norm(x, training=training)\n    x = self._activation_fn(x)\n    x = self._depthwise_conv(x)\n    x = self._depthwise_norm(x, training=training)\n    x = self._activation_fn(x)\n    x = self._se(x)\n    x = self._shrink_conv(x)\n\n    x = x+shortcut\n    return x","metadata":{"executionInfo":{"elapsed":60,"status":"ok","timestamp":1750539454330,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"UvrTSJWTOFas","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:44:06.047349Z","iopub.execute_input":"2025-06-27T17:44:06.047592Z","iopub.status.idle":"2025-06-27T17:44:06.118789Z","shell.execute_reply.started":"2025-06-27T17:44:06.047570Z","shell.execute_reply":"2025-06-27T17:44:06.117886Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Expand_layer(tf.keras.layers.Layer):\n    def __init__(self, **kwargs):\n        super().__init__(**kwargs)\n    def call(self, x):\n        return tf.expand_dims(x, (1))\n        \nclass Concat_layer(tf.keras.layers.Layer):\n    def __init__(self, **kwargs):\n        super().__init__(**kwargs)\n    def call(self, x_pred, x_confidence):\n        return tf.concat([tf.expand_dims(x_pred, -1), tf.expand_dims(x_confidence, -1)], axis = -1)\n\n\nclass Concat_layer_2(tf.keras.layers.Layer):\n    def __init__(self, **kwargs):\n        super().__init__(**kwargs)\n    def call(self, x_pred, x_label, x_label_2):\n        x_label_expand = tf.tile(tf.expand_dims(tf.expand_dims(x_label, (1)), (3)), (1,70,1,2))\n        x_label_2_expand = tf.tile(tf.expand_dims(tf.expand_dims(x_label_2, (1)), (3)), (1,70,1,2))\n        x_label_concat = tf.concat([x_label_expand, x_label_2_expand], axis = 2)\n        x_pred_concat = tf.concat([x_pred, x_label_concat], axis = 2)\n        return x_pred_concat","metadata":{"executionInfo":{"elapsed":3,"status":"ok","timestamp":1750539454338,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"VY-8R6SB1RlQ","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:44:06.119720Z","iopub.execute_input":"2025-06-27T17:44:06.119960Z","iopub.status.idle":"2025-06-27T17:44:06.135633Z","shell.execute_reply.started":"2025-06-27T17:44:06.119937Z","shell.execute_reply":"2025-06-27T17:44:06.134780Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Debugee","metadata":{"id":"nVXEWVMuKgS9"}},{"cell_type":"code","source":"dim = 32\n\nx = batch[0]\nx = Expand_layer()(x)\nx = tf.keras.layers.ZeroPadding2D(padding=(1+4//2, 1+4//2), data_format=\"channels_first\")(x)\nx = tf.keras.layers.Permute((2,3,1))(x)\nx = tf.keras.layers.Conv2D(dim, (4*2, 4*2), strides = (4,4))(x)\n\nx = MBConvBlock(dim, name = f'mbconv_{0}')(x)\nx = Rearrange('b h w c -> b (h w) c')(x)\nx = Pos_embedding_layer(num_patches = 324, dim = dim)(x)\nx = Transformer_X(dim=dim, heads=8, mlp_dim=dim*2*8//3)(x)\n\n\nfor i in range(10):\n    x = Rearrange('b (h w) c -> b h w c', h = 18)(x)\n    x = MBConvBlock(dim, name = f'mbconv_{i+1}')(x)\n    x = Rearrange('b h w c -> b (h w) c')(x)\n    x = Transformer_X(dim=dim, heads=8, mlp_dim=dim*2*8//3)(x)\nx = Rearrange('b (h w) c -> b h w c', h = 18)(x)\nx = tf.keras.layers.Dense(dim*4, activation='gelu', dtype=tf.float32)(x)\nx_label = tf.keras.layers.GlobalAveragePooling2D()(x)\nx_label = tf.keras.layers.Dense(10, activation='linear', dtype=tf.float32)(x_label)\nx_label_2 = tf.keras.layers.GlobalAveragePooling2D()(x)\nx_label_2 = tf.keras.layers.Dense(1, activation='linear', dtype=tf.float32)(x_label_2)\nx_confidence = tf.keras.layers.Dense(16, dtype=tf.float32)(x)\nx = tf.keras.layers.Dense(16, dtype=tf.float32)(x)\n\n\nx_confidence = Rearrange('b h w (p1 p2 c) -> b (h p1) (w p2) c', p1=4, p2=4)(x_confidence)\nx = Rearrange('b h w (p1 p2 c) -> b (h p1) (w p2) c', p1=4, p2=4)(x)\n\nx_confidence = tf.keras.layers.Cropping2D(cropping=((1, 1), (1, 1)))(x_confidence)\nx_confidence = tf.keras.layers.Reshape([70,70])(x_confidence)\n\nx = tf.keras.layers.Cropping2D(cropping=((1, 1), (1, 1)))(x)\nx = tf.keras.layers.Reshape([70,70])(x)\n\nx_pred = Concat_layer()(x, x_confidence)\nx_pred_concat = Concat_layer_2()(x_pred, x_label, x_label_2)\n\nprint(x_pred_concat.shape)","metadata":{"executionInfo":{"elapsed":69,"status":"ok","timestamp":1750539454447,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"bAMkArpOyKeQ","outputId":"e672d198-234d-4775-d0b1-f8d32009811b","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:44:06.136523Z","iopub.execute_input":"2025-06-27T17:44:06.136754Z","iopub.status.idle":"2025-06-27T17:44:43.779946Z","shell.execute_reply.started":"2025-06-27T17:44:06.136731Z","shell.execute_reply":"2025-06-27T17:44:43.778738Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x_pred_restored = x_pred_concat[:,:,:70,:]\nx_label_restored = x_pred_concat[:,0,70:80,0]\nx_label_2_restored = x_pred_concat[:,0,80:81,0]\n\nprint(x_pred_restored.shape)\nprint(np.mean(x_pred_restored == x_pred))\n\nprint(x_label_restored.shape)\nprint(np.mean(x_label_restored == x_label))\n\nprint(x_label_2_restored.shape)\nprint(np.mean(x_label_2_restored == x_label_2))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:44:43.781178Z","iopub.execute_input":"2025-06-27T17:44:43.781459Z","iopub.status.idle":"2025-06-27T17:44:43.797078Z","shell.execute_reply.started":"2025-06-27T17:44:43.781432Z","shell.execute_reply":"2025-06-27T17:44:43.796190Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"targets = batch[1]\ntargets.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:44:43.798245Z","iopub.execute_input":"2025-06-27T17:44:43.798508Z","iopub.status.idle":"2025-06-27T17:44:43.806841Z","shell.execute_reply.started":"2025-06-27T17:44:43.798484Z","shell.execute_reply":"2025-06-27T17:44:43.806055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_label = tf.cast(targets[:,1,0,0], tf.int64)\nvel_target = targets[:,0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:44:43.807804Z","iopub.execute_input":"2025-06-27T17:44:43.808062Z","iopub.status.idle":"2025-06-27T17:44:43.819719Z","shell.execute_reply.started":"2025-06-27T17:44:43.808037Z","shell.execute_reply":"2025-06-27T17:44:43.818936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"loss = tf.nn.sparse_softmax_cross_entropy_with_logits(labels=class_label, logits=x_label_restored)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:44:43.826577Z","iopub.execute_input":"2025-06-27T17:44:43.826819Z","iopub.status.idle":"2025-06-27T17:44:43.832191Z","shell.execute_reply.started":"2025-06-27T17:44:43.826796Z","shell.execute_reply":"2025-06-27T17:44:43.831351Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loss function","metadata":{"id":"mJdLF5oIHWzN"}},{"cell_type":"code","source":"\ndef loss_fn_valid(x, preds):\n    loss = tf.nn.sparse_softmax_cross_entropy_with_logits(labels=tf.cast(x, tf.int64), logits=preds)\n    return loss\n'''\ndef loss_fn(x, preds):\n    confidence = preds[:,:,:,1]\n    preds = preds[:,:,:,0]\n    loss = tf.math.abs(x-preds)\n    loss_2 = tf.math.abs(loss-confidence)\n    loss = tf.math.reduce_mean((loss+loss_2), axis = (1,2))\n    return loss\n'''\n\ndef loss_fn(targets, x_pred_concat):\n    x_pred_restored = x_pred_concat[:,:,:70,:]\n    x_label_restored = x_pred_concat[:,0,70:80,0]\n    x_label_2_restored = x_pred_concat[:,0,80:81,0]\n\n    class_label = tf.cast(targets[:,1,0,0], tf.int64)\n    vel_target = targets[:,0]\n    loss = tf.nn.sparse_softmax_cross_entropy_with_logits(labels=class_label, logits=x_label_restored)\n    return loss","metadata":{"executionInfo":{"elapsed":4,"status":"ok","timestamp":1750539454524,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"hluW13B8Gu6n","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:50:10.649931Z","iopub.execute_input":"2025-06-27T17:50:10.650391Z","iopub.status.idle":"2025-06-27T17:50:10.657810Z","shell.execute_reply.started":"2025-06-27T17:50:10.650357Z","shell.execute_reply":"2025-06-27T17:50:10.656869Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{"id":"cTzwW5bpHaDG"}},{"cell_type":"code","source":"def get_model(dim = 128, head_dim = 256):\n    with strategy.scope():\n        inp1 = tf.keras.Input([70,70])\n        x = inp1\n\n        \n        x = Expand_layer()(x)\n        x = tf.keras.layers.ZeroPadding2D(padding=(1+4//2, 1+4//2), data_format=\"channels_first\")(x)\n        x = tf.keras.layers.Permute((2,3,1))(x)\n        x = tf.keras.layers.Conv2D(dim, (4*2, 4*2), strides = (4,4))(x)\n        \n        x = MBConvBlock(dim, name = f'mbconv_{0}')(x)\n        x = Rearrange('b h w c -> b (h w) c')(x)\n        x = Pos_embedding_layer(num_patches = 324, dim = dim)(x)\n        x = Transformer_X(dim=dim, heads=8, mlp_dim=dim*2*8//3)(x)\n        \n        \n        for i in range(10):\n            x = Rearrange('b (h w) c -> b h w c', h = 18)(x)\n            x = MBConvBlock(dim, name = f'mbconv_{i+1}')(x)\n            x = Rearrange('b h w c -> b (h w) c')(x)\n            x = Transformer_X(dim=dim, heads=8, mlp_dim=dim*2*8//3)(x)\n        x = Rearrange('b (h w) c -> b h w c', h = 18)(x)\n        x = tf.keras.layers.Dense(dim*4, activation='gelu', dtype=tf.float32)(x)\n        x_label = tf.keras.layers.GlobalAveragePooling2D()(x)\n        x_label = tf.keras.layers.Dense(10, activation='linear', dtype=tf.float32)(x_label)\n        x_label_2 = tf.keras.layers.GlobalAveragePooling2D()(x)\n        x_label_2 = tf.keras.layers.Dense(1, activation='linear', dtype=tf.float32)(x_label_2)\n        x_confidence = tf.keras.layers.Dense(16, dtype=tf.float32)(x)\n        x = tf.keras.layers.Dense(16, dtype=tf.float32)(x)\n        \n        \n        x_confidence = Rearrange('b h w (p1 p2 c) -> b (h p1) (w p2) c', p1=4, p2=4)(x_confidence)\n        x = Rearrange('b h w (p1 p2 c) -> b (h p1) (w p2) c', p1=4, p2=4)(x)\n        \n        x_confidence = tf.keras.layers.Cropping2D(cropping=((1, 1), (1, 1)))(x_confidence)\n        x_confidence = tf.keras.layers.Reshape([70,70])(x_confidence)\n        \n        x = tf.keras.layers.Cropping2D(cropping=((1, 1), (1, 1)))(x)\n        x = tf.keras.layers.Reshape([70,70])(x)\n        \n        x_pred = Concat_layer()(x, x_confidence)\n        x_pred_concat = Concat_layer_2()(x_pred, x_label, x_label_2)\n\n        \n        model = tf.keras.Model(inp1, x_pred_concat)\n        return model","metadata":{"executionInfo":{"elapsed":3,"status":"ok","timestamp":1750539454530,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"IX9l_SAoG1LC","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:50:10.676068Z","iopub.execute_input":"2025-06-27T17:50:10.676392Z","iopub.status.idle":"2025-06-27T17:50:10.688912Z","shell.execute_reply.started":"2025-06-27T17:50:10.676364Z","shell.execute_reply":"2025-06-27T17:50:10.687937Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DEBUG:\n    model = get_model(24, head_dim = 4)\nelse:\n    model = get_model()\n\nwith strategy.scope():\n    loss = loss_fn\n    optimizer = tf.keras.optimizers.AdamW(learning_rate=0.0005, weight_decay=0.001)\n    if DEBUG:\n        model.compile(loss=loss, optimizer=optimizer)\n    else:\n        model.compile(loss=loss, optimizer=optimizer, steps_per_execution = 100)\n\nmodel(batch[0])\nmodel.summary()","metadata":{"executionInfo":{"elapsed":118910,"status":"ok","timestamp":1750539573457,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"T7ObT1a6Hvaz","outputId":"ec6d70d9-3b14-44c2-a096-007c1ce1a49d","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:50:10.734597Z","iopub.execute_input":"2025-06-27T17:50:10.734871Z","iopub.status.idle":"2025-06-27T17:51:19.184970Z","shell.execute_reply.started":"2025-06-27T17:50:10.734846Z","shell.execute_reply":"2025-06-27T17:51:19.184146Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x = batch[0]\nprint(x.shape)\nx = model(x)\nprint(x.shape)\ntargets = batch[1]\nprint(loss_fn(targets,x))","metadata":{"executionInfo":{"elapsed":83381,"status":"ok","timestamp":1750539656844,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"Y2XPmqoOs0b0","outputId":"0a1e50e0-51ad-449c-e87f-365cda41ace8","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:51:19.186680Z","iopub.execute_input":"2025-06-27T17:51:19.187045Z","iopub.status.idle":"2025-06-27T17:52:14.159923Z","shell.execute_reply.started":"2025-06-27T17:51:19.186982Z","shell.execute_reply":"2025-06-27T17:52:14.159145Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Learning rate scheduler","metadata":{"id":"6bACEYTEcXoz"}},{"cell_type":"markdown","source":"Schedulaer from https://www.kaggle.com/code/irohith/aslfr-ctc-based-on-prev-comp-1st-place\n\n","metadata":{"id":"Lpx9pERGca20"}},{"cell_type":"code","source":"N_EPOCHS","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:52:14.160845Z","iopub.execute_input":"2025-06-27T17:52:14.161122Z","iopub.status.idle":"2025-06-27T17:52:14.165934Z","shell.execute_reply.started":"2025-06-27T17:52:14.161096Z","shell.execute_reply":"2025-06-27T17:52:14.165322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def lrfn(current_step, num_warmup_steps, lr_max, num_cycles=0.50, num_training_steps=N_EPOCHS):\n    if current_step < num_warmup_steps:\n        if WARMUP_METHOD == 'log':\n            return lr_max * 0.10 ** (num_warmup_steps - current_step)\n        else:\n            return lr_max * 2 ** -(num_warmup_steps - current_step)\n    else:\n        progress = float(current_step - num_warmup_steps) / float(max(1, num_training_steps - num_warmup_steps))\n\n        return max(0.0, 0.5 * (1.0 + math.cos(math.pi * float(num_cycles) * 2.0 * progress))) * lr_max\n\ndef plot_lr_schedule(lr_schedule, epochs):\n    fig = plt.figure(figsize=(20, 10))\n    plt.plot([None] + lr_schedule + [None])\n    # X Labels\n    x = np.arange(1, epochs + 1)\n    x_axis_labels = [i if epochs <= 40 or i % 5 == 0 or i == 1 else None for i in range(1, epochs + 1)]\n    plt.xlim([1, epochs])\n    plt.xticks(x, x_axis_labels) # set tick step to 1 and let x axis start at 1\n\n    # Increase y-limit for better readability\n    plt.ylim([0, max(lr_schedule) * 1.1])\n\n    # Title\n    schedule_info = f'start: {lr_schedule[0]:.1E}, max: {max(lr_schedule):.1E}, final: {lr_schedule[-1]:.1E}'\n    plt.title(f'Step Learning Rate Schedule, {schedule_info}', size=18, pad=12)\n\n    # Plot Learning Rates\n    for x, val in enumerate(lr_schedule):\n        if epochs <= 40 or x % 5 == 0 or x is epochs - 1:\n            if x < len(lr_schedule) - 1:\n                if lr_schedule[x - 1] < val:\n                    ha = 'right'\n                else:\n                    ha = 'left'\n            elif x == 0:\n                ha = 'right'\n            else:\n                ha = 'left'\n            plt.plot(x + 1, val, 'o', color='black');\n            offset_y = (max(lr_schedule) - min(lr_schedule)) * 0.02\n            plt.annotate(f'{val:.1E}', xy=(x + 1, val + offset_y), size=12, ha=ha)\n\n    plt.xlabel('Epoch', size=16, labelpad=5)\n    plt.ylabel('Learning Rate', size=16, labelpad=5)\n    plt.grid()\n    plt.show()\n\n# Learning rate for encoder\nLR_SCHEDULE = [lrfn(step, num_warmup_steps=N_WARMUP_EPOCHS, lr_max=LR_MAX, num_cycles=0.50, num_training_steps = N_EPOCHS)\n for step in range(N_EPOCHS)][:]\n#LR_SCHEDULE = list(np.linspace(LR_MAX, LR_MIN, num=N_EPOCHS))\n# Plot Learning Rate Schedule\nplot_lr_schedule(LR_SCHEDULE, epochs=N_EPOCHS)\n# Learning Rate Callback\nlr_callback = tf.keras.callbacks.LearningRateScheduler(lambda step: LR_SCHEDULE[step], verbose=0)\n\n# Custom callback to update weight decay with learning rate\nclass WeightDecayCallback(tf.keras.callbacks.Callback):\n    def __init__(self, wd_ratio=WD_RATIO):\n        self.step_counter = 0\n        self.wd_ratio = wd_ratio\n\n    def on_epoch_begin(self, epoch, logs=None):\n        model.optimizer.weight_decay = model.optimizer.learning_rate * self.wd_ratio\n        print(f'learning rate: {model.optimizer.learning_rate.numpy():.2e}, weight decay: {model.optimizer.weight_decay.numpy():.2e}')","metadata":{"executionInfo":{"elapsed":425,"status":"ok","timestamp":1750539657272,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"h1A4DmzEcX2P","outputId":"e6fa77cd-83a2-40ae-bde2-d56e31522aff","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:52:14.167597Z","iopub.execute_input":"2025-06-27T17:52:14.167828Z","iopub.status.idle":"2025-06-27T17:52:15.061399Z","shell.execute_reply.started":"2025-06-27T17:52:14.167805Z","shell.execute_reply":"2025-06-27T17:52:15.060628Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\nval_samples_num = val_steps_per_epoch*val_batch_size\nprint(val_samples_num)\n\nepoch_samples_num = steps_per_epoch*batch_size\nif DEBUG:\n    epoch_samples_num = 10\nprint(epoch_samples_num)\n'''","metadata":{"executionInfo":{"elapsed":9,"status":"ok","timestamp":1750539657283,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"v0iQOGxwOaWk","outputId":"7af97424-fc6f-422d-b049-47566537db7c","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:52:15.062368Z","iopub.execute_input":"2025-06-27T17:52:15.062622Z","iopub.status.idle":"2025-06-27T17:52:15.067297Z","shell.execute_reply.started":"2025-06-27T17:52:15.062597Z","shell.execute_reply":"2025-06-27T17:52:15.066686Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_labels(ds_for_extracting):\n    extracted_data = [x[1][:,1,0,0] for x in ds_for_extracting]\n    extracted_data = tf.concat(extracted_data, axis = 0)\n    return extracted_data\n\ndef get_features(ds_for_extracting):\n    extracted_data = [x[0] for x in ds_for_extracting]\n    extracted_data = tf.concat(extracted_data, axis = 0)\n    return extracted_data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:52:15.068109Z","iopub.execute_input":"2025-06-27T17:52:15.068337Z","iopub.status.idle":"2025-06-27T17:52:15.083470Z","shell.execute_reply.started":"2025-06-27T17:52:15.068315Z","shell.execute_reply":"2025-06-27T17:52:15.082792Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_ds_features = get_features(val_ds)\nprint(val_ds_features.shape)\nval_ds_features = val_ds_features[:10000]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:52:15.084532Z","iopub.execute_input":"2025-06-27T17:52:15.084790Z","iopub.status.idle":"2025-06-27T17:52:15.585028Z","shell.execute_reply.started":"2025-06-27T17:52:15.084764Z","shell.execute_reply":"2025-06-27T17:52:15.584119Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_ds_labels = get_labels(val_ds)\nval_ds_2_labels = get_labels(val_ds_2)\nprint(val_ds_labels.shape)\nprint(val_ds_2_labels.shape)\nval_ds_labels = val_ds_labels[:10000]\nval_ds_2_labels = val_ds_2_labels[:5000]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:52:15.585973Z","iopub.execute_input":"2025-06-27T17:52:15.586239Z","iopub.status.idle":"2025-06-27T17:52:16.213050Z","shell.execute_reply.started":"2025-06-27T17:52:15.586214Z","shell.execute_reply":"2025-06-27T17:52:16.212168Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\nval_labels = val_ds_labels\nmetrics_1 = val_ds_features\n\nscore = tf.reduce_mean(loss_fn_valid(val_labels, metrics_1))\nprint(score)\nfor i in range(10):\n    score = tf.reduce_mean(loss_fn_valid(val_labels[i*1000:(i+1)*1000], metrics_1[i*1000:(i+1)*1000]))\n    print(score)\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:52:16.214136Z","iopub.execute_input":"2025-06-27T17:52:16.214428Z","iopub.status.idle":"2025-06-27T17:52:16.219279Z","shell.execute_reply.started":"2025-06-27T17:52:16.214401Z","shell.execute_reply":"2025-06-27T17:52:16.218551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_steps_per_epoch_1 = 11000//val_batch_size\nval_steps_per_epoch_2 = 6000//val_batch_size","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:52:16.222106Z","iopub.execute_input":"2025-06-27T17:52:16.222379Z","iopub.status.idle":"2025-06-27T17:52:16.237785Z","shell.execute_reply.started":"2025-06-27T17:52:16.222352Z","shell.execute_reply":"2025-06-27T17:52:16.236940Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"jump = 20\nval_scores_list = []\nval_scores_list_2 = []\nscore_val_min_subset = [1000]\nscore_val_min = [1000]\nclass validation_callback(tf.keras.callbacks.Callback):\n    def __init__(self):\n        super().__init__()\n    def on_epoch_end(self, epoch: int, logs=None):\n        if (epoch+1)%jump == 0 or epoch == 0:\n            print('Metric: 1------------------')\n            scores = []\n            metrics_1 = model.predict(val_ds, verbose = 0, batch_size = val_batch_size, steps = val_steps_per_epoch_1)\n            metrics_1 = metrics_1[:,0,70:80,0]\n            metrics_1 = metrics_1[:10000]\n            \n            val_labels = val_ds_labels\n            \n            score = tf.reduce_mean(loss_fn_valid(val_labels, metrics_1))\n            print(score)\n            scores.append(score)\n            for i in range(10):\n                score = tf.reduce_mean(loss_fn_valid(val_labels[i*1000:(i+1)*1000], metrics_1[i*1000:(i+1)*1000]))\n                print(score)\n                scores.append(score)\n            val_scores_list.append(scores)\n\n            print('Metric: 2------------------')\n            scores = []\n            metrics_1 = model.predict(val_ds_2, verbose = 0, batch_size = val_batch_size, steps = val_steps_per_epoch_2)\n            metrics_1 = metrics_1[:,0,70:80,0]\n            metrics_1 = metrics_1[:5000]\n            \n            val_labels = val_ds_2_labels\n            \n            score = tf.reduce_mean(loss_fn_valid(val_labels, metrics_1))\n            print(score)\n            scores.append(score)\n            for i in range(10):\n                score = tf.reduce_mean(loss_fn_valid(val_labels[i*500:(i+1)*500], metrics_1[i*500:(i+1)*500]))\n                print(score)\n                scores.append(score)\n            val_scores_list_2.append(scores)","metadata":{"executionInfo":{"elapsed":23,"status":"ok","timestamp":1750539657484,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"fD59FIpyaVXQ","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:52:16.238696Z","iopub.execute_input":"2025-06-27T17:52:16.238911Z","iopub.status.idle":"2025-06-27T17:52:16.254691Z","shell.execute_reply.started":"2025-06-27T17:52:16.238889Z","shell.execute_reply":"2025-06-27T17:52:16.253964Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\nx = batch[0]\nx = model(x)\ntargets = batch[1]\n\nx_label_restored = x[:,0,70:80,0]\nclass_label = tf.cast(targets[:,1,0,0], tf.int64)\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:52:16.255528Z","iopub.execute_input":"2025-06-27T17:52:16.255756Z","iopub.status.idle":"2025-06-27T17:52:16.269820Z","shell.execute_reply.started":"2025-06-27T17:52:16.255733Z","shell.execute_reply":"2025-06-27T17:52:16.269058Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":" if DEBUG:\n    model = get_model(24, head_dim = 4)\nelse:\n    model = get_model()\n\nwith strategy.scope():\n    loss = loss_fn\n    optimizer = tf.keras.optimizers.AdamW(learning_rate=0.0005)\n\n    #checkpoint_old = tf.train.Checkpoint(model=model,optimizer=optimizer)\n    #checkpoint_old.restore(f'{checkpoint_prev_path}/ckpt_folder/ckpt-6')\n    if DEBUG:\n        model.compile(loss=loss, optimizer=optimizer)\n    else:\n        model.compile(loss=loss, optimizer=optimizer, steps_per_execution = 10)\n\ncheckpoint = tf.train.Checkpoint(model=model,optimizer=optimizer)\nos.makedirs(f'{save_folder}/ckpt_folder_best', exist_ok=True)\ncheckpoint_manager = tf.train.CheckpointManager(checkpoint, directory=f'{save_folder}/ckpt_folder_best', max_to_keep=1)\n\n\nhistory = model.fit(train_ds, verbose=2,\n                    steps_per_epoch = steps_per_epoch,\n                    epochs=training_epochs, batch_size=batch_size,\n                    callbacks=[validation_callback(), lr_callback, WeightDecayCallback()])","metadata":{"id":"MUH4B6sUlurp","trusted":true,"execution":{"iopub.status.busy":"2025-06-27T17:52:16.270726Z","iopub.execute_input":"2025-06-27T17:52:16.270961Z","iopub.status.idle":"2025-06-27T18:03:55.496228Z","shell.execute_reply.started":"2025-06-27T17:52:16.270938Z","shell.execute_reply":"2025-06-27T18:03:55.494917Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"checkpoint.save(f'{save_folder}/ckpt_folder/ckpt')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"metrics_1 = model.predict(val_ds, verbose = 0, batch_size = val_batch_size, steps = val_steps_per_epoch_1)\nmetrics_1 = metrics_1[:,0,70:80,0]\nmetrics_1 = metrics_1[:10000]\n\nval_labels = val_ds_labels","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pickle.dump(metrics_1, open('metrics_1.p', 'bw'))\npickle.dump(val_labels, open('val_labels.p', 'bw'))","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}