{"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},{"sourceId":247749734,"sourceType":"kernelVersion"},{"sourceId":248030507,"sourceType":"kernelVersion"}],"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-30T08:30:41.897570Z","iopub.execute_input":"2025-06-30T08:30:41.897814Z","iopub.status.idle":"2025-06-30T08:30:41.912158Z","shell.execute_reply.started":"2025-06-30T08:30:41.897790Z","shell.execute_reply":"2025-06-30T08:30:41.911560Z"}},"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-30T08:30:41.913743Z","iopub.execute_input":"2025-06-30T08:30:41.913975Z","iopub.status.idle":"2025-06-30T08:30:49.137957Z","shell.execute_reply.started":"2025-06-30T08:30:41.913953Z","shell.execute_reply":"2025-06-30T08:30:49.137029Z"}},"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-30T08:30:49.139190Z","iopub.execute_input":"2025-06-30T08:30:49.139447Z","iopub.status.idle":"2025-06-30T08:31:05.390712Z","shell.execute_reply.started":"2025-06-30T08:30:49.139420Z","shell.execute_reply":"2025-06-30T08:31:05.389958Z"}},"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-30T08:31:05.391769Z","iopub.execute_input":"2025-06-30T08:31:05.392256Z","iopub.status.idle":"2025-06-30T08:31:07.341126Z","shell.execute_reply.started":"2025-06-30T08:31:05.392228Z","shell.execute_reply":"2025-06-30T08:31:07.340410Z"}},"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-30T08:31:07.342023Z","iopub.execute_input":"2025-06-30T08:31:07.342460Z","iopub.status.idle":"2025-06-30T08:31:07.354762Z","shell.execute_reply.started":"2025-06-30T08:31:07.342433Z","shell.execute_reply":"2025-06-30T08:31:07.354139Z"}},"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-30T08:31:07.356787Z","iopub.execute_input":"2025-06-30T08:31:07.357028Z","iopub.status.idle":"2025-06-30T08:31:15.693860Z","shell.execute_reply.started":"2025-06-30T08:31:07.357005Z","shell.execute_reply":"2025-06-30T08:31:15.693069Z"}},"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-30T08:31:15.694741Z","iopub.execute_input":"2025-06-30T08:31:15.694966Z","iopub.status.idle":"2025-06-30T08:31:19.776558Z","shell.execute_reply.started":"2025-06-30T08:31:15.694942Z","shell.execute_reply":"2025-06-30T08:31:19.775388Z"}},"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-30T08:31:19.777940Z","iopub.execute_input":"2025-06-30T08:31:19.778248Z","iopub.status.idle":"2025-06-30T08:33:15.242897Z","shell.execute_reply.started":"2025-06-30T08:31:19.778220Z","shell.execute_reply":"2025-06-30T08:33:15.241621Z"}},"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_2_preds = pickle.load(open('//kaggle/input/gwi-ensemble-best-nobest/val_ensemble.p', 'br'))\n    val_ensemble_labels = pickle.load(open('/kaggle/working/val_ensemble/val_labels.p', 'br'))\n    test_ensemble = pickle.load(open('/kaggle/input/gwi-ensemble-best-nobest/test_ensemble.p', 'br'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:33:15.244598Z","iopub.execute_input":"2025-06-30T08:33:15.244929Z","iopub.status.idle":"2025-06-30T08:33:23.677697Z","shell.execute_reply.started":"2025-06-30T08:33:15.244899Z","shell.execute_reply":"2025-06-30T08:33:23.676558Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\nshuffle_rng = np.random.default_rng(seed=4224)\nfor x in models_list_all:\n    shuffle_rng.shuffle(x)\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:33:23.678999Z","iopub.execute_input":"2025-06-30T08:33:23.679333Z","iopub.status.idle":"2025-06-30T08:33:23.688370Z","shell.execute_reply.started":"2025-06-30T08:33:23.679303Z","shell.execute_reply":"2025-06-30T08:33:23.687514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#val_models_list = [x[:500] for x in models_list_all]\n#train_models_list = [x[500:] for x in models_list_all]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:33:23.689537Z","iopub.execute_input":"2025-06-30T08:33:23.689864Z","iopub.status.idle":"2025-06-30T08:33:23.703509Z","shell.execute_reply.started":"2025-06-30T08:33:23.689837Z","shell.execute_reply":"2025-06-30T08:33:23.702691Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\nif 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]\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:33:23.704518Z","iopub.execute_input":"2025-06-30T08:33:23.704806Z","iopub.status.idle":"2025-06-30T08:33:23.716091Z","shell.execute_reply.started":"2025-06-30T08:33:23.704778Z","shell.execute_reply":"2025-06-30T08:33:23.715368Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\nprint(len(train_models_list))\nprint(len(train_models_list[-1]))\nplt.imshow(train_models_list[-1][43][0])\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:33:23.717214Z","iopub.execute_input":"2025-06-30T08:33:23.717479Z","iopub.status.idle":"2025-06-30T08:33:23.732166Z","shell.execute_reply.started":"2025-06-30T08:33:23.717454Z","shell.execute_reply":"2025-06-30T08:33:23.731337Z"}},"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-30T08:33:23.733140Z","iopub.execute_input":"2025-06-30T08:33:23.733378Z","iopub.status.idle":"2025-06-30T08:33:24.176004Z","shell.execute_reply.started":"2025-06-30T08:33:23.733355Z","shell.execute_reply":"2025-06-30T08:33:24.175017Z"}},"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-30T08:33:24.177293Z","iopub.execute_input":"2025-06-30T08:33:24.177602Z","iopub.status.idle":"2025-06-30T08:33:24.187813Z","shell.execute_reply.started":"2025-06-30T08:33:24.177573Z","shell.execute_reply":"2025-06-30T08:33:24.186890Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"training_epochs = N_EPOCHS","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:33:24.188962Z","iopub.execute_input":"2025-06-30T08:33:24.189282Z","iopub.status.idle":"2025-06-30T08:33:27.389492Z","shell.execute_reply.started":"2025-06-30T08:33:24.189251Z","shell.execute_reply":"2025-06-30T08:33:27.388312Z"}},"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-30T08:33:27.390812Z","iopub.execute_input":"2025-06-30T08:33:27.391148Z","iopub.status.idle":"2025-06-30T08:33:29.723054Z","shell.execute_reply.started":"2025-06-30T08:33:27.391099Z","shell.execute_reply":"2025-06-30T08:33:29.721763Z"}},"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-30T08:33:29.727543Z","iopub.execute_input":"2025-06-30T08:33:29.727829Z","iopub.status.idle":"2025-06-30T08:33:31.547710Z","shell.execute_reply.started":"2025-06-30T08:33:29.727802Z","shell.execute_reply":"2025-06-30T08:33:31.546529Z"}},"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)\n\nval_ensemble_2_preds = np.concatenate([val_ensemble_2_preds, val_ensemble_2_preds[-1000:]], axis = 0)\ntest_ensemble = np.concatenate([test_ensemble, test_ensemble[-1000:]], axis = 0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:33:31.548887Z","iopub.execute_input":"2025-06-30T08:33:31.549228Z","iopub.status.idle":"2025-06-30T08:33:40.636559Z","shell.execute_reply.started":"2025-06-30T08:33:31.549197Z","shell.execute_reply":"2025-06-30T08:33:40.635272Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.plot(class_labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:33:40.637768Z","iopub.execute_input":"2025-06-30T08:33:40.638043Z","iopub.status.idle":"2025-06-30T08:33:49.814278Z","shell.execute_reply.started":"2025-06-30T08:33:40.638017Z","shell.execute_reply":"2025-06-30T08:33:49.813234Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_val_dataset(val_ensemble_preds, val_ensemble_labels, class_labels):\n    PAD = 0.0\n    val_ds = tf.data.Dataset.from_tensor_slices((val_ensemble_preds, val_ensemble_labels, class_labels))\n    val_ds = val_ds.map(val_ds_to_dict, tf.data.AUTOTUNE)\n    val_ds = val_ds.map(get_output, tf.data.AUTOTUNE)\n    if DEBUG:\n        val_ds = val_ds.take(64)\n    samples_num = val_ds.reduce(0, lambda x,_: x+1).numpy()\n    val_ds = val_ds.padded_batch(\n                batch_size, padding_values=(PAD,PAD),\n        padded_shapes=([70,70],[2,70,70]), drop_remainder=True)\n    val_ds = val_ds.prefetch(tf.data.AUTOTUNE)\n    return val_ds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:33:49.815518Z","iopub.execute_input":"2025-06-30T08:33:49.815823Z","iopub.status.idle":"2025-06-30T08:33:54.288223Z","shell.execute_reply.started":"2025-06-30T08:33:49.815793Z","shell.execute_reply":"2025-06-30T08:33:54.286980Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_ds = get_val_dataset(val_ensemble_preds, val_ensemble_labels, class_labels)\nval_ds_2 = get_val_dataset(val_ensemble_2_preds, val_ensemble_labels, class_labels)\ntest_ds = get_val_dataset(test_ensemble, test_ensemble*0, np.zeros((len(test_ensemble))))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:33:54.289353Z","iopub.execute_input":"2025-06-30T08:33:54.289633Z","iopub.status.idle":"2025-06-30T08:34:05.742058Z","shell.execute_reply.started":"2025-06-30T08:33:54.289605Z","shell.execute_reply":"2025-06-30T08:34:05.740857Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"laa = [x for x in val_ds.take(2)]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:34:05.743453Z","iopub.execute_input":"2025-06-30T08:34:05.743853Z","iopub.status.idle":"2025-06-30T08:34:06.282902Z","shell.execute_reply.started":"2025-06-30T08:34:05.743818Z","shell.execute_reply":"2025-06-30T08:34:06.281949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(laa[0][1][0,0])\nplt.show()\nplt.imshow(laa[0][0][0])\nplt.show()\nplt.imshow(laa[0][1][0,0]-laa[0][0][0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:34:06.284288Z","iopub.execute_input":"2025-06-30T08:34:06.284682Z","iopub.status.idle":"2025-06-30T08:34:30.397226Z","shell.execute_reply.started":"2025-06-30T08:34:06.284651Z","shell.execute_reply":"2025-06-30T08:34:30.396452Z"}},"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-30T08:34:30.398086Z","iopub.execute_input":"2025-06-30T08:34:30.398342Z","iopub.status.idle":"2025-06-30T08:34:35.341478Z","shell.execute_reply.started":"2025-06-30T08:34:30.398317Z","shell.execute_reply":"2025-06-30T08:34:35.340574Z"}},"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-30T08:34:35.342520Z","iopub.execute_input":"2025-06-30T08:34:35.342902Z","iopub.status.idle":"2025-06-30T08:34:35.746885Z","shell.execute_reply.started":"2025-06-30T08:34:35.342871Z","shell.execute_reply":"2025-06-30T08:34:35.746160Z"}},"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-30T08:34:35.748035Z","iopub.execute_input":"2025-06-30T08:34:35.748326Z","iopub.status.idle":"2025-06-30T08:34:35.797227Z","shell.execute_reply.started":"2025-06-30T08:34:35.748299Z","shell.execute_reply":"2025-06-30T08:34:35.796453Z"}},"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-30T08:34:35.798127Z","iopub.execute_input":"2025-06-30T08:34:35.798439Z","iopub.status.idle":"2025-06-30T08:34:40.199557Z","shell.execute_reply.started":"2025-06-30T08:34:35.798410Z","shell.execute_reply":"2025-06-30T08:34:40.198529Z"}},"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-30T08:34:40.200992Z","iopub.execute_input":"2025-06-30T08:34:40.201504Z","iopub.status.idle":"2025-06-30T08:34:40.215855Z","shell.execute_reply.started":"2025-06-30T08:34:40.201473Z","shell.execute_reply":"2025-06-30T08:34:40.215157Z"}},"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-30T08:34:40.216912Z","iopub.execute_input":"2025-06-30T08:34:40.217182Z","iopub.status.idle":"2025-06-30T08:34:40.287888Z","shell.execute_reply.started":"2025-06-30T08:34:40.217156Z","shell.execute_reply":"2025-06-30T08:34:40.287293Z"}},"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-30T08:34:40.288730Z","iopub.execute_input":"2025-06-30T08:34:40.288949Z","iopub.status.idle":"2025-06-30T08:34:40.302891Z","shell.execute_reply.started":"2025-06-30T08:34:40.288926Z","shell.execute_reply":"2025-06-30T08:34:40.302309Z"}},"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-30T08:34:40.303822Z","iopub.execute_input":"2025-06-30T08:34:40.304061Z","iopub.status.idle":"2025-06-30T08:35:17.782476Z","shell.execute_reply.started":"2025-06-30T08:34:40.304035Z","shell.execute_reply":"2025-06-30T08:35:17.781462Z"}},"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-30T08:35:17.783611Z","iopub.execute_input":"2025-06-30T08:35:17.783883Z","iopub.status.idle":"2025-06-30T08:35:17.799704Z","shell.execute_reply.started":"2025-06-30T08:35:17.783856Z","shell.execute_reply":"2025-06-30T08:35:17.798749Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"targets = batch[1]\ntargets.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:35:17.800816Z","iopub.execute_input":"2025-06-30T08:35:17.801135Z","iopub.status.idle":"2025-06-30T08:35:17.808032Z","shell.execute_reply.started":"2025-06-30T08:35:17.801090Z","shell.execute_reply":"2025-06-30T08:35:17.807073Z"}},"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-30T08:35:17.809056Z","iopub.execute_input":"2025-06-30T08:35:17.809321Z","iopub.status.idle":"2025-06-30T08:35:17.821810Z","shell.execute_reply.started":"2025-06-30T08:35:17.809295Z","shell.execute_reply":"2025-06-30T08:35:17.821043Z"}},"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-30T08:35:17.822796Z","iopub.execute_input":"2025-06-30T08:35:17.823067Z","iopub.status.idle":"2025-06-30T08:35:17.830028Z","shell.execute_reply.started":"2025-06-30T08:35:17.823041Z","shell.execute_reply":"2025-06-30T08:35:17.829281Z"}},"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-30T08:35:17.830951Z","iopub.execute_input":"2025-06-30T08:35:17.831198Z","iopub.status.idle":"2025-06-30T08:35:17.837766Z","shell.execute_reply.started":"2025-06-30T08:35:17.831173Z","shell.execute_reply":"2025-06-30T08:35:17.836797Z"}},"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-30T08:35:17.838903Z","iopub.execute_input":"2025-06-30T08:35:17.839166Z","iopub.status.idle":"2025-06-30T08:35:17.851212Z","shell.execute_reply.started":"2025-06-30T08:35:17.839141Z","shell.execute_reply":"2025-06-30T08:35:17.850384Z"}},"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-30T08:35:17.852130Z","iopub.execute_input":"2025-06-30T08:35:17.852363Z","iopub.status.idle":"2025-06-30T08:36:25.761844Z","shell.execute_reply.started":"2025-06-30T08:35:17.852338Z","shell.execute_reply":"2025-06-30T08:36:25.760903Z"}},"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-30T08:36:25.763012Z","iopub.execute_input":"2025-06-30T08:36:25.763314Z","iopub.status.idle":"2025-06-30T08:37:19.536844Z","shell.execute_reply.started":"2025-06-30T08:36:25.763287Z","shell.execute_reply":"2025-06-30T08:37:19.535897Z"}},"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-30T08:37:19.537859Z","iopub.execute_input":"2025-06-30T08:37:19.538130Z","iopub.status.idle":"2025-06-30T08:37:19.542780Z","shell.execute_reply.started":"2025-06-30T08:37:19.538091Z","shell.execute_reply":"2025-06-30T08:37:19.542153Z"}},"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-30T08:37:19.543634Z","iopub.execute_input":"2025-06-30T08:37:19.543860Z","iopub.status.idle":"2025-06-30T08:37:21.757014Z","shell.execute_reply.started":"2025-06-30T08:37:19.543836Z","shell.execute_reply":"2025-06-30T08:37:21.756147Z"}},"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-30T08:37:21.757994Z","iopub.execute_input":"2025-06-30T08:37:21.758264Z","iopub.status.idle":"2025-06-30T08:37:21.762493Z","shell.execute_reply.started":"2025-06-30T08:37:21.758238Z","shell.execute_reply":"2025-06-30T08:37:21.761842Z"}},"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-30T08:37:21.763314Z","iopub.execute_input":"2025-06-30T08:37:21.763545Z","iopub.status.idle":"2025-06-30T08:37:21.784789Z","shell.execute_reply.started":"2025-06-30T08:37:21.763521Z","shell.execute_reply":"2025-06-30T08:37:21.784086Z"}},"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-30T08:37:21.785626Z","iopub.execute_input":"2025-06-30T08:37:21.785872Z","iopub.status.idle":"2025-06-30T08:37:23.009368Z","shell.execute_reply.started":"2025-06-30T08:37:21.785847Z","shell.execute_reply":"2025-06-30T08:37:23.008500Z"}},"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-30T08:37:23.010335Z","iopub.execute_input":"2025-06-30T08:37:23.010584Z","iopub.status.idle":"2025-06-30T08:37:24.581145Z","shell.execute_reply.started":"2025-06-30T08:37:23.010560Z","shell.execute_reply":"2025-06-30T08:37:24.580344Z"}},"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-30T08:37:24.582024Z","iopub.execute_input":"2025-06-30T08:37:24.582270Z","iopub.status.idle":"2025-06-30T08:37:24.586862Z","shell.execute_reply.started":"2025-06-30T08:37:24.582245Z","shell.execute_reply":"2025-06-30T08:37:24.586124Z"}},"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-30T08:37:24.587782Z","iopub.execute_input":"2025-06-30T08:37:24.588008Z","iopub.status.idle":"2025-06-30T08:37:24.603827Z","shell.execute_reply.started":"2025-06-30T08:37:24.587985Z","shell.execute_reply":"2025-06-30T08:37:24.602974Z"}},"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-30T08:37:24.604903Z","iopub.execute_input":"2025-06-30T08:37:24.605171Z","iopub.status.idle":"2025-06-30T08:37:24.630139Z","shell.execute_reply.started":"2025-06-30T08:37:24.605145Z","shell.execute_reply":"2025-06-30T08:37:24.629507Z"}},"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-30T08:37:24.630988Z","iopub.execute_input":"2025-06-30T08:37:24.631226Z","iopub.status.idle":"2025-06-30T08:37:24.645145Z","shell.execute_reply.started":"2025-06-30T08:37:24.631203Z","shell.execute_reply":"2025-06-30T08:37:24.644481Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.mkdir('ckpt_folder')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:45:48.227348Z","iopub.execute_input":"2025-06-30T08:45:48.227738Z","iopub.status.idle":"2025-06-30T08:45:48.231835Z","shell.execute_reply.started":"2025-06-30T08:45:48.227711Z","shell.execute_reply":"2025-06-30T08:45:48.230980Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"source_path = '/kaggle/input/gwi-classifier-1/ckpt_folder/checkpoint'\ndestination_path = '/kaggle/working/ckpt_folder/checkpoint'\nshutil.copy2(source_path, destination_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:47:35.544918Z","iopub.execute_input":"2025-06-30T08:47:35.545760Z","iopub.status.idle":"2025-06-30T08:47:35.563167Z","shell.execute_reply.started":"2025-06-30T08:47:35.545727Z","shell.execute_reply":"2025-06-30T08:47:35.562426Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"source_path = '/kaggle/input/gwi-classifier-1/ckpt_folder/ckpt-1.data-00000-of-00001'\ndestination_path = '/kaggle/working/ckpt_folder/ckpt-1.data-00000-of-00001'\nshutil.copy2(source_path, destination_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:47:58.297166Z","iopub.execute_input":"2025-06-30T08:47:58.297969Z","iopub.status.idle":"2025-06-30T08:47:59.686797Z","shell.execute_reply.started":"2025-06-30T08:47:58.297934Z","shell.execute_reply":"2025-06-30T08:47:59.686043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"source_path = '/kaggle/input/gwi-classifier-1/ckpt_folder/ckpt-1.index'\ndestination_path = '/kaggle/working/ckpt_folder/ckpt-1.index'\nshutil.copy2(source_path, destination_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:48:22.170711Z","iopub.execute_input":"2025-06-30T08:48:22.171046Z","iopub.status.idle":"2025-06-30T08:48:22.196362Z","shell.execute_reply.started":"2025-06-30T08:48:22.171018Z","shell.execute_reply":"2025-06-30T08:48:22.195603Z"}},"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'/kaggle/working/ckpt_folder/ckpt-1')\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()])\n'''","metadata":{"id":"MUH4B6sUlurp","trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:48:46.775319Z","iopub.execute_input":"2025-06-30T08:48:46.775720Z","iopub.status.idle":"2025-06-30T08:48:55.924494Z","shell.execute_reply.started":"2025-06-30T08:48:46.775689Z","shell.execute_reply":"2025-06-30T08:48:55.923476Z"}},"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,"execution":{"iopub.status.busy":"2025-06-30T08:53:25.354526Z","iopub.execute_input":"2025-06-30T08:53:25.354931Z","iopub.status.idle":"2025-06-30T08:54:10.759104Z","shell.execute_reply.started":"2025-06-30T08:53:25.354898Z","shell.execute_reply":"2025-06-30T08:54:10.757715Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"score = 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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:56:54.732855Z","iopub.execute_input":"2025-06-30T08:56:54.734026Z","iopub.status.idle":"2025-06-30T08:56:54.754006Z","shell.execute_reply.started":"2025-06-30T08:56:54.733986Z","shell.execute_reply":"2025-06-30T08:56:54.753082Z"}},"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,"execution":{"iopub.status.busy":"2025-06-30T08:57:21.558767Z","iopub.execute_input":"2025-06-30T08:57:21.559725Z","iopub.status.idle":"2025-06-30T08:57:21.566473Z","shell.execute_reply.started":"2025-06-30T08:57:21.559685Z","shell.execute_reply":"2025-06-30T08:57:21.565521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"metrics_2 = model.predict(val_ds_2, verbose = 0, batch_size = val_batch_size, steps = val_steps_per_epoch_1)\nmetrics_2 = metrics_2[:,0,70:80,0]\nmetrics_2 = metrics_2[:10000]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:57:45.214671Z","iopub.execute_input":"2025-06-30T08:57:45.215071Z","iopub.status.idle":"2025-06-30T08:57:48.928615Z","shell.execute_reply.started":"2025-06-30T08:57:45.215041Z","shell.execute_reply":"2025-06-30T08:57:48.927477Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"score = tf.reduce_mean(loss_fn_valid(val_labels, metrics_2))\nprint(score)\nfor i in range(10):\n    score = tf.reduce_mean(loss_fn_valid(val_labels[i*1000:(i+1)*1000], metrics_2[i*1000:(i+1)*1000]))\n    print(score)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:57:55.537993Z","iopub.execute_input":"2025-06-30T08:57:55.538464Z","iopub.status.idle":"2025-06-30T08:57:55.559070Z","shell.execute_reply.started":"2025-06-30T08:57:55.538422Z","shell.execute_reply":"2025-06-30T08:57:55.558151Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pickle.dump(metrics_2, open('metrics_2.p', 'bw'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:58:20.923562Z","iopub.execute_input":"2025-06-30T08:58:20.924539Z","iopub.status.idle":"2025-06-30T08:58:20.930849Z","shell.execute_reply.started":"2025-06-30T08:58:20.924495Z","shell.execute_reply":"2025-06-30T08:58:20.929828Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_ensemble.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T08:59:58.528449Z","iopub.execute_input":"2025-06-30T08:59:58.528910Z","iopub.status.idle":"2025-06-30T08:59:58.535494Z","shell.execute_reply.started":"2025-06-30T08:59:58.528873Z","shell.execute_reply":"2025-06-30T08:59:58.534374Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_preds = model.predict(test_ds, verbose = 0, batch_size = val_batch_size)\ntest_preds = test_preds[:,0,70:80,0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:00:04.135790Z","iopub.execute_input":"2025-06-30T09:00:04.136280Z","iopub.status.idle":"2025-06-30T09:00:21.293910Z","shell.execute_reply.started":"2025-06-30T09:00:04.136238Z","shell.execute_reply":"2025-06-30T09:00:21.292696Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_ensemble_orig = pickle.load(open('/kaggle/input/gwi-ensemble-best-nobest/test_ensemble.p', 'br'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:01:36.650280Z","iopub.execute_input":"2025-06-30T09:01:36.650687Z","iopub.status.idle":"2025-06-30T09:01:37.661833Z","shell.execute_reply.started":"2025-06-30T09:01:36.650655Z","shell.execute_reply":"2025-06-30T09:01:37.660627Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_ensemble_orig.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:01:42.464829Z","iopub.execute_input":"2025-06-30T09:01:42.465253Z","iopub.status.idle":"2025-06-30T09:01:42.471207Z","shell.execute_reply.started":"2025-06-30T09:01:42.465222Z","shell.execute_reply":"2025-06-30T09:01:42.470266Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_preds.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:00:26.135195Z","iopub.execute_input":"2025-06-30T09:00:26.136146Z","iopub.status.idle":"2025-06-30T09:00:26.141565Z","shell.execute_reply.started":"2025-06-30T09:00:26.136087Z","shell.execute_reply":"2025-06-30T09:00:26.140642Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_preds_to_save = test_preds[:len(test_ensemble_orig)]\ntest_preds_to_save.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:02:13.316031Z","iopub.execute_input":"2025-06-30T09:02:13.316515Z","iopub.status.idle":"2025-06-30T09:02:13.322774Z","shell.execute_reply.started":"2025-06-30T09:02:13.316479Z","shell.execute_reply":"2025-06-30T09:02:13.321991Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pickle.dump(test_preds_to_save, open('test_preds.p', 'bw'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:02:47.869373Z","iopub.execute_input":"2025-06-30T09:02:47.869852Z","iopub.status.idle":"2025-06-30T09:02:47.888251Z","shell.execute_reply.started":"2025-06-30T09:02:47.869808Z","shell.execute_reply":"2025-06-30T09:02:47.887069Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_preds_to_save.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:06:09.865757Z","iopub.execute_input":"2025-06-30T09:06:09.866163Z","iopub.status.idle":"2025-06-30T09:06:09.871630Z","shell.execute_reply.started":"2025-06-30T09:06:09.866131Z","shell.execute_reply":"2025-06-30T09:06:09.870869Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_labels = np.argmax(test_preds_to_save, axis = 1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:07:12.512086Z","iopub.execute_input":"2025-06-30T09:07:12.512492Z","iopub.status.idle":"2025-06-30T09:07:12.522097Z","shell.execute_reply.started":"2025-06-30T09:07:12.512461Z","shell.execute_reply":"2025-06-30T09:07:12.521170Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i in range(10):\n    print(np.sum(test_labels == i)/len(test_labels))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:07:54.363361Z","iopub.execute_input":"2025-06-30T09:07:54.363832Z","iopub.status.idle":"2025-06-30T09:07:54.370228Z","shell.execute_reply.started":"2025-06-30T09:07:54.363797Z","shell.execute_reply":"2025-06-30T09:07:54.369281Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_labels = np.argmax(metrics_2, axis = 1)\nfor i in range(10):\n    print(np.sum(val_labels == i)/len(val_labels))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:10:01.633595Z","iopub.execute_input":"2025-06-30T09:10:01.634751Z","iopub.status.idle":"2025-06-30T09:10:01.641072Z","shell.execute_reply.started":"2025-06-30T09:10:01.634707Z","shell.execute_reply":"2025-06-30T09:10:01.640054Z"}},"outputs":[],"execution_count":null}]}