{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.16","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"},{"sourceId":11367935,"sourceType":"datasetVersion","datasetId":7116013},{"sourceId":11368499,"sourceType":"datasetVersion","datasetId":7116445},{"sourceId":11368545,"sourceType":"datasetVersion","datasetId":7116479},{"sourceId":11368547,"sourceType":"datasetVersion","datasetId":7116481},{"sourceId":11376433,"sourceType":"datasetVersion","datasetId":7122462},{"sourceId":11376448,"sourceType":"datasetVersion","datasetId":7122476},{"sourceId":11376464,"sourceType":"datasetVersion","datasetId":7122489},{"sourceId":11376742,"sourceType":"datasetVersion","datasetId":7122712},{"sourceId":11376935,"sourceType":"datasetVersion","datasetId":7122866},{"sourceId":11377083,"sourceType":"datasetVersion","datasetId":7122981},{"sourceId":332421,"sourceType":"modelInstanceVersion","modelInstanceId":278650,"modelId":299553}],"dockerImageVersionId":30920,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Training a TensorFlow model on TPUs results in the loss becoming infinite, after convergence, at some constant step.  I first observed it on the WFI data from the Yale/UNC-CH competition, but managed to gather a Minimal Working Example which fails on dummy data, whatever the model, whatever the batch size, whatever the number of epochs (as long as there are enough to reach this point, obviously), whatever the number of loaded files, whatever the tensor size.","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\n\nresolver = tf.distribute.cluster_resolver.TPUClusterResolver(tpu='local')\ntf.config.experimental_connect_to_cluster(resolver)\ntf.tpu.experimental.initialize_tpu_system(resolver)\nstrategy = tf.distribute.TPUStrategy(resolver)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T11:05:49.634144Z","iopub.execute_input":"2025-04-21T11:05:49.634360Z","iopub.status.idle":"2025-04-21T11:06:20.086568Z","shell.execute_reply.started":"2025-04-21T11:05:49.634335Z","shell.execute_reply":"2025-04-21T11:06:20.085746Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The issue seems to occur whatever the batch size : I tried 1024 on the original data, then 512, then 256, then 6.","metadata":{}},{"cell_type":"code","source":"#BATCH_SIZE = 256\nBATCH_SIZE = 6\nNUM_BATCHES = 50  # fails at #42 anyway","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T11:06:20.087893Z","iopub.execute_input":"2025-04-21T11:06:20.088156Z","iopub.status.idle":"2025-04-21T11:06:20.092592Z","shell.execute_reply.started":"2025-04-21T11:06:20.088130Z","shell.execute_reply":"2025-04-21T11:06:20.091175Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The issue occurs whatever the model, even with this straightforward one.","metadata":{}},{"cell_type":"code","source":"class DummyModel(tf.keras.Model):\n    def call(self, x, training=False):\n        batch_size = tf.shape(x)[0]\n        return tf.ones([batch_size, 7, 7, 1], dtype=tf.float32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T11:28:31.702038Z","iopub.execute_input":"2025-04-21T11:28:31.702459Z","iopub.status.idle":"2025-04-21T11:28:31.707533Z","shell.execute_reply.started":"2025-04-21T11:28:31.702410Z","shell.execute_reply":"2025-04-21T11:28:31.706516Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Observation : the loss becomes infinite à batch #41, at all epochs.","metadata":{}},{"cell_type":"code","source":"N_total = BATCH_SIZE * NUM_BATCHES\nX_all = tf.ones((N_total, 10, 7, 5), dtype=tf.float32)\ny_all = tf.ones((N_total, 7, 7, 1), dtype=tf.float32) / 2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T11:28:32.234147Z","iopub.execute_input":"2025-04-21T11:28:32.234601Z","iopub.status.idle":"2025-04-21T11:28:32.240643Z","shell.execute_reply.started":"2025-04-21T11:28:32.234562Z","shell.execute_reply":"2025-04-21T11:28:32.239557Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = tf.data.Dataset.from_tensor_slices((X_all, y_all))\nds = ds.batch(BATCH_SIZE)\ndummytrain = ds.prefetch(1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T11:28:32.682171Z","iopub.execute_input":"2025-04-21T11:28:32.682555Z","iopub.status.idle":"2025-04-21T11:28:32.698561Z","shell.execute_reply.started":"2025-04-21T11:28:32.682527Z","shell.execute_reply":"2025-04-21T11:28:32.697200Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with strategy.scope():\n    model = DummyModel()\n    model.compile(optimizer='adam', loss='mae')\n\nh = model.fit(\n    dummytrain,\n    epochs=2\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T11:28:33.121119Z","iopub.execute_input":"2025-04-21T11:28:33.121505Z","iopub.status.idle":"2025-04-21T11:28:36.723438Z","shell.execute_reply.started":"2025-04-21T11:28:33.121441Z","shell.execute_reply":"2025-04-21T11:28:36.721627Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The model state seems correct, since it works without recompiling :","metadata":{}},{"cell_type":"code","source":"h = model.fit(\n    dummytrain.take(2),\n    epochs=2\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T11:14:05.802796Z","iopub.execute_input":"2025-04-21T11:14:05.803257Z","iopub.status.idle":"2025-04-21T11:14:08.189227Z","shell.execute_reply.started":"2025-04-21T11:14:05.803139Z","shell.execute_reply":"2025-04-21T11:14:08.187606Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The issue is not a model.fit bug :","metadata":{}},{"cell_type":"code","source":"with strategy.scope():\n    model = DummyModel()\n    model.compile(optimizer='adam', loss='mae')\n    \nfor i, (xb, yb) in enumerate(dummytrain):\n    loss = model.train_on_batch(xb, yb)\n    print(f\"{i}: loss = {loss}\")\n    if tf.math.is_nan(loss) or tf.math.is_inf(loss):\n        print(f\"🚨 NaN or Inf detected at batch {i}\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T11:29:50.972962Z","iopub.execute_input":"2025-04-21T11:29:50.973356Z","execution_failed":"2025-04-21T11:30:02.515Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}