{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpuV5e8","dataSources":[{"sourceType":"competition","sourceId":21154,"databundleVersionId":1243559}],"dockerImageVersionId":31331,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-24T20:08:29.535857Z","iopub.execute_input":"2026-03-24T20:08:29.536186Z","iopub.status.idle":"2026-03-24T20:08:29.585642Z","shell.execute_reply.started":"2026-03-24T20:08:29.536163Z","shell.execute_reply":"2026-03-24T20:08:29.584475Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\n\nimport jax\nimport keras\nimport tensorflow as tf\n\nprint(\"JAX devices:\", jax.devices())\nprint(\"Keras backend:\", keras.backend.backend())\n\ndef setup_environment(competition_name : str = \"tpu-getting-started\", base_batch_per_core : int = 16):\n\n    env_config = {\n        \"is_kaggle\" : False,\n        \"device_type\" : \"UNKNOWN\",\n        \"num_devices\" : 0,\n        \"global_batch_size\" : base_batch_per_core,\n        \"data_path\": \"./data\",\n        \"output_path\": \"./output\"\n    }\n\n    try:\n        from kaggle_datasets import KaggleDatasets\n        env_config['is_kaggle'] = True\n        env_config['data_path'] = \"/kaggle/input/competitions/\" + competition_name\n        env_config['output_path'] = \"/kaggle/working\"\n        print(\"Running on Kaggle\")\n    except ImportError:\n        env_config['is_kaggle'] = False\n        print(\"Running on Local\")\n    \n    devices = jax.devices()\n    env_config['num_devices'] = len(devices)\n    env_config['device_type'] = devices[0].device_kind\n    print(\"Number of devices:\", env_config['num_devices'])\n    print(\"Device type:\", env_config['device_type'])\n    \n    dp_strategy = keras.distribution.DataParallel(devices=devices)\n    keras.distribution.set_distribution(dp_strategy)\n\n    env_config[\"global_batch_size\"] = base_batch_per_core * env_config[\"num_devices\"]\n\n    return env_config\n    \nif __name__ == \"__main__\":\n    config = setup_environment(\"tpu-getting-started\")\n    \n    print(config)\n\n        \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-24T20:08:29.586362Z","iopub.execute_input":"2026-03-24T20:08:29.586522Z","iopub.status.idle":"2026-03-24T20:08:29.592402Z","shell.execute_reply.started":"2026-03-24T20:08:29.586507Z","shell.execute_reply":"2026-03-24T20:08:29.591496Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from kaggle_datasets import KaggleDatasets\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-24T20:08:29.592800Z","iopub.execute_input":"2026-03-24T20:08:29.592958Z","iopub.status.idle":"2026-03-24T20:08:29.603415Z","shell.execute_reply.started":"2026-03-24T20:08:29.592936Z","shell.execute_reply":"2026-03-24T20:08:29.602511Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_PATH = config[\"data_path\"]\nBATCH_SIZE = config[\"global_batch_size\"]\nIMAGE_SIZE = [512, 512] \nAUTO = tf.data.AUTOTUNE\n\n\nTRAIN_FILENAMES = tf.io.gfile.glob(f'{DATA_PATH}/tfrecords-jpeg-{IMAGE_SIZE[0]}x{IMAGE_SIZE[0]}/train/*.tfrec')\nVAL_FILENAMES = tf.io.gfile.glob(f'{DATA_PATH}/tfrecords-jpeg-{IMAGE_SIZE[0]}x{IMAGE_SIZE[0]}/val/*.tfrec')\nTEST_FILENAMES = tf.io.gfile.glob(f'{DATA_PATH}/tfrecords-jpeg-{IMAGE_SIZE[0]}x{IMAGE_SIZE[0]}/test/*.tfrec')\n\n\ndef decode_image(image_data):\n    image = tf.image.decode_jpeg(image_data, channels=3)\n    image = tf.cast(image, tf.float32) / 255.0  \n    image = tf.reshape(image, [*IMAGE_SIZE, 3]) \n    return image\n\ndef read_labeled_tfrecord(example):\n    LABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"class\": tf.io.FixedLenFeature([], tf.int64),\n    }\n    example = tf.io.parse_single_example(example, LABELED_TFREC_FORMAT)\n    image = decode_image(example['image'])\n    label = tf.cast(example['class'], tf.int32)\n    return image, label\n\ndef read_unlabeled_tfrecord(example):\n    UNLABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"id\": tf.io.FixedLenFeature([], tf.string),\n    }\n    example = tf.io.parse_single_example(example, UNLABELED_TFREC_FORMAT)\n    image = decode_image(example['image'])\n    idnum = example['id']\n    return image, idnum\n\n\ndef load_dataset(filenames, labeled=True, ordered=False):\n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.deterministic = False \n\n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=AUTO)\n    dataset = dataset.with_options(ignore_order)\n    \n    parse_fn = read_labeled_tfrecord if labeled else read_unlabeled_tfrecord\n    dataset = dataset.map(parse_fn, num_parallel_calls=AUTO)\n    return dataset\n\ndef get_training_dataset():\n    dataset = load_dataset(TRAIN_FILENAMES, labeled=True)\n    dataset = dataset.repeat() \n    dataset = dataset.shuffle(2048)\n    dataset = dataset.batch(BATCH_SIZE, drop_remainder=True) \n    dataset = dataset.prefetch(AUTO) \n    return dataset\n\ndef get_validation_dataset():\n    dataset = load_dataset(VAL_FILENAMES, labeled=True, ordered=False)\n    dataset = dataset.batch(BATCH_SIZE, drop_remainder=False) \n    dataset = dataset.cache() \n    dataset = dataset.prefetch(AUTO)\n    return dataset\n\ndef get_test_dataset(ordered=True):\n    dataset = load_dataset(TEST_FILENAMES, labeled=False, ordered=ordered)\n    dataset = dataset.batch(BATCH_SIZE, drop_remainder=False)\n    dataset = dataset.prefetch(AUTO)\n    return dataset\n\n\nds_train = get_training_dataset()\nds_valid = get_validation_dataset()\nds_test = get_test_dataset()\n\nprint(\"Data pipeline construction completed\")\nprint(f\"Data source: {DATA_PATH}\")\nprint(f\"Train dataset Batch Size: {BATCH_SIZE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-24T20:08:29.604004Z","iopub.execute_input":"2026-03-24T20:08:29.604157Z","iopub.status.idle":"2026-03-24T20:08:30.136734Z","shell.execute_reply.started":"2026-03-24T20:08:29.604143Z","shell.execute_reply":"2026-03-24T20:08:30.135345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUM_TRAINING_IMAGES = 12753\nNUM_VALIDATION_IMAGES = 3712\nNUM_TEST_IMAGES = 7382\n\n\nSTEPS_PER_EPOCH = NUM_TRAINING_IMAGES // BATCH_SIZE\n\n\ndef build_model():\n    base_model = keras.applications.EfficientNetV2B0(\n        input_shape=[*IMAGE_SIZE, 3],\n        include_top=False, \n        weights='imagenet'\n    )\n    \n    base_model.trainable = True \n\n    inputs = keras.Input(shape=[*IMAGE_SIZE, 3])\n    x = base_model(inputs)\n    \n    x = keras.layers.GlobalAveragePooling2D()(x)\n    \n    outputs = keras.layers.Dense(104, activation='softmax')(x)\n    \n    return keras.Model(inputs=inputs, outputs=outputs)\n\n\n\n\nmodel = build_model()\n\n# Compile and Train\n\nmodel.compile(\n    optimizer=keras.optimizers.Adam(learning_rate=1e-4),\n    loss='sparse_categorical_crossentropy',\n    metrics=['sparse_categorical_accuracy']\n)\n\nprint(\"Compile complete. Start training on TPU\")\n\nearly_stopping = keras.callbacks.EarlyStopping(\n    monitor='val_sparse_categorical_accuracy', \n    patience=10, \n    restore_best_weights=True\n)\n\nhistory = model.fit(\n    ds_train,\n    validation_data=ds_valid,\n    epochs=40,\n    steps_per_epoch=STEPS_PER_EPOCH,\n    callbacks=[early_stopping]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-24T20:08:30.137672Z","iopub.execute_input":"2026-03-24T20:08:30.137854Z","iopub.status.idle":"2026-03-24T20:18:24.559781Z","shell.execute_reply.started":"2026-03-24T20:08:30.137837Z","shell.execute_reply":"2026-03-24T20:18:24.557563Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import math\nimport numpy as np\nimport pandas as pd\n\n\nNUM_TEST_IMAGES = 7382\npad_elements = math.ceil(NUM_TEST_IMAGES / BATCH_SIZE) * BATCH_SIZE - NUM_TEST_IMAGES\n\nunbatched_ds = ds_test.unbatch()\n\n\npadded_test_ds = unbatched_ds.concatenate(unbatched_ds.take(pad_elements)).batch(BATCH_SIZE, drop_remainder=True)\n\n\ntest_images_ds = padded_test_ds.map(lambda image, idnum: image)\n\ntest_ids_ds = unbatched_ds.map(lambda image, idnum: idnum)\n\nprint(\"Start predicting on TPU\")\nprobabilities = model.predict(test_images_ds)\n\n# 拔刀斩：砍掉尾部凑数的假预测，严格切回 7382\nprobabilities = probabilities[:NUM_TEST_IMAGES]\npredictions = np.argmax(probabilities, axis=-1)\n\nprint(\"Decoding underlying image ID byte stream\")\ntest_ids_bytes = list(test_ids_ds.as_numpy_iterator())\ntest_ids = [id_byte.decode('utf-8') for id_byte in test_ids_bytes]\n\nassert len(test_ids) == len(predictions), f\"Fatal error: The number of IDs ({len(test_ids)}) does not match the number of predictions ({len(predictions)})!\"\n\nprint(\"Generating submission.csv...\")\nsubmission = pd.DataFrame({\n    'id': test_ids,\n    'label': predictions\n})\nsubmission.to_csv('submission.csv', index=False)\n\nprint(\"Success!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-24T20:18:24.561353Z","iopub.execute_input":"2026-03-24T20:18:24.561575Z","iopub.status.idle":"2026-03-24T20:18:49.886610Z","shell.execute_reply.started":"2026-03-24T20:18:24.561557Z","shell.execute_reply":"2026-03-24T20:18:49.885229Z"}},"outputs":[],"execution_count":null}]}