{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":21154,"databundleVersionId":1243559,"sourceType":"competition"},{"sourceId":37130068,"sourceType":"kernelVersion"}],"dockerImageVersionId":30646,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n# Introduction #\n\nWelcome to the [**Petals to the Metal**](https://www.kaggle.com/c/tpu-getting-started) competition! In this competition, you’re challenged to build a machine learning model to classify 104 types of flowers based on their images.\n\nIn this tutorial notebook, you'll learn how to build an image classifier in Keras and train it on a [Tensor Processing Unit (TPU)](https://www.kaggle.com/docs/tpu). At the end, you'll have a complete project you can build off of with ideas of your own.\n\n<blockquote style=\"margin-right:auto; margin-left:auto; background-color: #ebf9ff; padding: 1em; margin:24px;\">\n    <strong>Fork This Notebook!</strong><br>\nCreate your own editable copy of this notebook by clicking on the <strong>Copy and Edit</strong> button in the top right corner.\n</blockquote>","metadata":{}},{"cell_type":"markdown","source":"# Step 1: Imports #\n\nWe begin by importing several Python packages.","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras import layers\nfrom tensorflow.keras.applications import DenseNet201 # 引入 DenseNet\nimport pandas as pd\nimport numpy as np\nimport re       # 处理正则\nimport os       # 处理系统路径\nimport math     # <--- 刚才缺失的数学库！\nimport matplotlib.pyplot as plt # <--- 画图库，你也肯定需要！\nfrom tensorflow.keras import mixed_precision\n\n# 开启混合精度，速度翻倍！\npolicy = mixed_precision.Policy('mixed_float16')\nmixed_precision.set_global_policy(policy)\n\nprint(\"Mixed Precision enabled\")\n# 建议使用 224 (标准) 或 331 (更高清，但更慢)\nIMAGE_SIZE = [331,331] \n\n# 检测硬件\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.experimental.TPUStrategy(tpu)\n    print(\"Running on TPU\")\nexcept ValueError:\n    strategy = tf.distribute.get_strategy()\n    print(\"Running on GPU/CPU\")\n\n# 如果是 GPU，BATCH_SIZE 不要太大，32 (16*2) 是个安全值\nBATCH_SIZE = 16 * strategy.num_replicas_in_sync\nprint(f\"Batch Size: {BATCH_SIZE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T11:59:38.656695Z","iopub.execute_input":"2025-12-24T11:59:38.657043Z","iopub.status.idle":"2025-12-24T11:59:50.076297Z","shell.execute_reply.started":"2025-12-24T11:59:38.657009Z","shell.execute_reply":"2025-12-24T11:59:50.075415Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 2: Distribution Strategy #\n\nA TPU has eight different *cores* and each of these cores acts as its own accelerator. (A TPU is sort of like having eight GPUs in one machine.) We tell TensorFlow how to make use of all these cores at once through a **distribution strategy**. Run the following cell to create the distribution strategy that we'll later apply to our model.","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\n\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver() \n    print('Running on TPU ', tpu.master())\nexcept ValueError:\n    tpu = None\n\nif tpu:\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.TPUStrategy(tpu)\nelse:\n    # --- 修改部分开始 ---\n    # 检查是否有 GPU\n    gpus = tf.config.list_physical_devices('GPU')\n    if len(gpus) > 1:\n        # 如果有多个 GPU，使用 MirroredStrategy 来启用“双核/多核”并行\n        strategy = tf.distribute.MirroredStrategy()\n        print(f\"检测到 {len(gpus)} 个 GPU，已启用 MirroredStrategy (双核/多卡模式)！\")\n    else:\n        # 只有一个 GPU 或没有 GPU，使用默认策略\n        strategy = tf.distribute.get_strategy()\n        print(\"未检测到多 GPU，使用默认策略。\")\n    # --- 修改部分结束 ---\n\nprint(\"REPLICAS (并行数量): \", strategy.num_replicas_in_sync)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T11:59:50.077878Z","iopub.execute_input":"2025-12-24T11:59:50.078348Z","iopub.status.idle":"2025-12-24T11:59:50.873486Z","shell.execute_reply.started":"2025-12-24T11:59:50.078324Z","shell.execute_reply":"2025-12-24T11:59:50.872558Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We'll use the distribution strategy when we create our neural network model. Then, TensorFlow will distribute the training among the eight TPU cores by creating eight different *replicas* of the model, one for each core.\n\n# Step 3: Loading the Competition Data #\n\n## Get GCS Path ##\n\nWhen used with TPUs, datasets need to be stored in a [Google Cloud Storage bucket](https://cloud.google.com/storage/). You can use data from any public GCS bucket by giving its path just like you would data from `'/kaggle/input'`. The following will retrieve the GCS path for this competition's dataset.","metadata":{}},{"cell_type":"code","source":"from kaggle_datasets import KaggleDatasets\n\nGCS_DS_PATH = KaggleDatasets().get_gcs_path('tpu-getting-started')\nprint(GCS_DS_PATH) # what do gcs paths look like?","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T11:59:50.874797Z","iopub.execute_input":"2025-12-24T11:59:50.875156Z","iopub.status.idle":"2025-12-24T11:59:51.140417Z","shell.execute_reply.started":"2025-12-24T11:59:50.875123Z","shell.execute_reply":"2025-12-24T11:59:51.139529Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"You can use data from any public dataset here on Kaggle in just the same way. If you'd like to use data from one of your private datasets, see [here](https://www.kaggle.com/docs/tpu#tpu3pt5).\n\n## Load Data ##\n\nWhen used with TPUs, datasets are often serialized into [TFRecords](https://www.kaggle.com/ryanholbrook/tfrecords-basics). This is a format convenient for distributing data to each of the TPUs cores. We've hidden the cell that reads the TFRecords for our dataset since the process is a bit long. You could come back to it later for some guidance on using your own datasets with TPUs.","metadata":{}},{"cell_type":"code","source":"\nGCS_PATH = GCS_DS_PATH + '/tfrecords-jpeg-512x512'\nAUTO = tf.data.experimental.AUTOTUNE\n\nTRAINING_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/train/*.tfrec')\nVALIDATION_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/val/*.tfrec')\nTEST_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/test/*.tfrec') \n\nCLASSES = ['pink primrose',    'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea',     'wild geranium',     'tiger lily',           'moon orchid',              'bird of paradise', 'monkshood',        'globe thistle',         # 00 - 09\n           'snapdragon',       \"colt's foot\",               'king protea',      'spear thistle', 'yellow iris',       'globe-flower',         'purple coneflower',        'peruvian lily',    'balloon flower',   'giant white arum lily', # 10 - 19\n           'fire lily',        'pincushion flower',         'fritillary',       'red ginger',    'grape hyacinth',    'corn poppy',           'prince of wales feathers', 'stemless gentian', 'artichoke',        'sweet william',         # 20 - 29\n           'carnation',        'garden phlox',              'love in the mist', 'cosmos',        'alpine sea holly',  'ruby-lipped cattleya', 'cape flower',              'great masterwort', 'siam tulip',       'lenten rose',           # 30 - 39\n           'barberton daisy',  'daffodil',                  'sword lily',       'poinsettia',    'bolero deep blue',  'wallflower',           'marigold',                 'buttercup',        'daisy',            'common dandelion',      # 40 - 49\n           'petunia',          'wild pansy',                'primula',          'sunflower',     'lilac hibiscus',    'bishop of llandaff',   'gaura',                    'geranium',         'orange dahlia',    'pink-yellow dahlia',    # 50 - 59\n           'cautleya spicata', 'japanese anemone',          'black-eyed susan', 'silverbush',    'californian poppy', 'osteospermum',         'spring crocus',            'iris',             'windflower',       'tree poppy',            # 60 - 69\n           'gazania',          'azalea',                    'water lily',       'rose',          'thorn apple',       'morning glory',        'passion flower',           'lotus',            'toad lily',        'anthurium',             # 70 - 79\n           'frangipani',       'clematis',                  'hibiscus',         'columbine',     'desert-rose',       'tree mallow',          'magnolia',                 'cyclamen ',        'watercress',       'canna lily',            # 80 - 89\n           'hippeastrum ',     'bee balm',                  'pink quill',       'foxglove',      'bougainvillea',     'camellia',             'mallow',                   'mexican petunia',  'bromelia',         'blanket flower',        # 90 - 99\n           'trumpet creeper',  'blackberry lily',           'common tulip',     'wild rose']                                                                                                                                               # 100 - 102\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    \n    # 【关键修复】强制 Resize，不保留长宽比，确保所有图片尺寸一模一样\n    # 只要这里生效，就不会报 shape 错误\n    image = tf.image.resize(image, IMAGE_SIZE) \n    \n    return image\n\ndef read_labeled_tfrecord(example):\n    LABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring\n        \"class\": tf.io.FixedLenFeature([], tf.int64),  # shape [] means single element\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 # returns a dataset of (image, label) pairs\n\ndef read_unlabeled_tfrecord(example):\n    UNLABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring\n        \"id\": tf.io.FixedLenFeature([], tf.string),  # shape [] means single element\n        # class is missing, this competitions's challenge is to predict flower classes for the test dataset\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 # returns a dataset of image(s)\n\ndef load_dataset(filenames, labeled=True, ordered=False):\n    # Read from TFRecords. For optimal performance, reading from multiple files at once and\n    # disregarding data order. Order does not matter since we will be shuffling the data anyway.\n\n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.experimental_deterministic = False # disable order, increase speed\n\n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=AUTO) # automatically interleaves reads from multiple files\n    dataset = dataset.with_options(ignore_order) # uses data as soon as it streams in, rather than in its original order\n    dataset = dataset.map(read_labeled_tfrecord if labeled else read_unlabeled_tfrecord, num_parallel_calls=AUTO)\n    # returns a dataset of (image, label) pairs if labeled=True or (image, id) pairs if labeled=False\n    return dataset","metadata":{"_kg_hide-input":true,"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T11:59:51.142743Z","iopub.execute_input":"2025-12-24T11:59:51.143004Z","iopub.status.idle":"2025-12-24T11:59:51.291784Z","shell.execute_reply.started":"2025-12-24T11:59:51.142982Z","shell.execute_reply":"2025-12-24T11:59:51.291126Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create Data Pipelines ##\n\nIn this final step we'll use the `tf.data` API to define an efficient data pipeline for each of the training, validation, and test splits.","metadata":{}},{"cell_type":"code","source":"def data_augment(image, label):\n    \"\"\"\n    增强版数据增强，不使用外部库\n    \"\"\"\n    # ====== 基础增强 ======\n    # 随机左右翻转\n    image = tf.image.random_flip_left_right(image)\n    \n    # 随机上下翻转\n    if tf.random.uniform([]) > 0.7:\n        image = tf.image.random_flip_up_down(image)\n    \n    # ====== 随机旋转（使用rot90实现） ======\n    if tf.random.uniform([]) > 0.3:\n        k = tf.random.uniform([], minval=0, maxval=4, dtype=tf.int32)\n        image = tf.image.rot90(image, k=k)\n    \n    # ====== 随机缩放裁剪 ======\n    if tf.random.uniform([]) > 0.4:\n        # 随机缩放 (0.65-1.0)\n        scale = tf.random.uniform([], 0.65, 1.0)\n        h = tf.cast(tf.cast(IMAGE_SIZE[0], tf.float32) * scale, tf.int32)\n        w = tf.cast(tf.cast(IMAGE_SIZE[1], tf.float32) * scale, tf.int32)\n        \n        # 确保尺寸有效\n        h = tf.maximum(h, 1)\n        w = tf.maximum(w, 1)\n        \n        # 随机偏移\n        offset_h = tf.random.uniform([], 0, tf.maximum(1, IMAGE_SIZE[0] - h), dtype=tf.int32)\n        offset_w = tf.random.uniform([], 0, tf.maximum(1, IMAGE_SIZE[1] - w), dtype=tf.int32)\n        \n        # 裁剪并调整大小\n        image = tf.image.crop_to_bounding_box(image, offset_h, offset_w, h, w)\n        image = tf.image.resize(image, IMAGE_SIZE)\n    \n    # ====== 颜色增强 ======\n    # 随机亮度\n    image = tf.image.random_brightness(image, max_delta=0.25)\n    \n    # 随机对比度\n    image = tf.image.random_contrast(image, lower=0.7, upper=1.3)\n    \n    # 随机饱和度\n    image = tf.image.random_saturation(image, lower=0.7, upper=1.3)\n    \n    # 随机色调\n    image = tf.image.random_hue(image, max_delta=0.1)\n    \n    # ====== 随机遮挡（CutOut） ======\n    if tf.random.uniform([]) > 0.5:\n        # 随机遮挡大小（10%-30%）\n        mask_h = tf.cast(tf.random.uniform([], 0.1, 0.3) * IMAGE_SIZE[0], tf.int32)\n        mask_w = tf.cast(tf.random.uniform([], 0.1, 0.3) * IMAGE_SIZE[1], tf.int32)\n        \n        # 确保尺寸有效\n        mask_h = tf.maximum(mask_h, 1)\n        mask_w = tf.maximum(mask_w, 1)\n        \n        # 随机位置\n        y = tf.random.uniform([], 0, IMAGE_SIZE[0] - mask_h, dtype=tf.int32)\n        x = tf.random.uniform([], 0, IMAGE_SIZE[1] - mask_w, dtype=tf.int32)\n        \n        # 创建遮挡\n        # 方法1：设置为平均颜色（更自然）\n        avg_color = tf.reduce_mean(image, axis=[0, 1], keepdims=True)\n        mask = tf.ones([mask_h, mask_w, 3]) * avg_color\n        \n        # 创建更新后的图像\n        top = image[:y, :, :]\n        middle_top = image[y:y+mask_h, :x, :]\n        middle_bottom = image[y:y+mask_h, x+mask_w:, :]\n        bottom = image[y+mask_h:, :, :]\n        \n        # 重建图像\n        image_top = top\n        image_middle = tf.concat([middle_top, mask, middle_bottom], axis=1)\n        image_bottom = bottom\n        \n        image = tf.concat([image_top, image_middle, image_bottom], axis=0)\n    \n    # ====== 噪声增强 ======\n    if tf.random.uniform([]) > 0.3:\n        noise = tf.random.normal(tf.shape(image), mean=0.0, stddev=0.02)\n        image = image + noise\n    \n    # ====== 随机模糊（通过下采样实现） ======\n    if tf.random.uniform([]) > 0.7:\n        # 临时缩小再放大实现模糊效果\n        small_size = [IMAGE_SIZE[0] // 2, IMAGE_SIZE[1] // 2]\n        image = tf.image.resize(image, small_size)\n        image = tf.image.resize(image, IMAGE_SIZE)\n    \n    # ====== 确保数值范围 ======\n    image = tf.clip_by_value(image, 0.0, 1.0)\n    \n    return image, label\n\n\n# --- 新增代码开始：CutMix 和 One-Hot 辅助函数 ---\n\n# 确保定义了分类数量\nCLASSES_NUM = 104\n\ndef one_hot(image, label):\n    \"\"\"\n    将整数标签转换为 One-hot 编码。\n    CutMix/MixUp 需要标签是概率分布（如 [0, 0.5, 0, 0.5...]），\n    而不是单一整数。\n    \"\"\"\n    label = tf.cast(label, tf.int32)\n    label = tf.one_hot(label, CLASSES_NUM)\n    label = tf.cast(label, tf.float32)\n    return image, label\n\ndef cutmix(image, label, PROBABILITY=1.0):\n    \"\"\"\n    CutMix 数据增强函数：\n    在 Batch 层面操作，随机将一张图的一部分剪切粘贴到另一张图上，\n    并按面积比例混合标签。\n    \"\"\"\n    # input image shape: [batch, size, size, 3]\n    # input label shape: [batch, classes]\n    \n    DIM = IMAGE_SIZE[0] # 获取图片尺寸\n    \n    imgs = []\n    labs = []\n    \n    for j in range(BATCH_SIZE):\n        # 决定是否进行 CutMix\n        P = tf.cast(tf.random.uniform([], 0, 1) <= PROBABILITY, tf.int32)\n        \n        # 随机选择 Batch 中的另一张图片 k\n        k = tf.cast(tf.random.uniform([], 0, BATCH_SIZE), tf.int32)\n        \n        # 生成剪切框 (Bounding Box)\n        x = tf.cast(tf.random.uniform([], 0, DIM), tf.int32)\n        y = tf.cast(tf.random.uniform([], 0, DIM), tf.int32)\n        \n        # Beta 分布生成剪切比例\n        b = tf.random.uniform([], 0, 1) \n        w = tf.cast(DIM * tf.math.sqrt(1-b), tf.int32) * P\n        \n        xa = tf.maximum(0, x - w // 2)\n        yb = tf.maximum(0, y - w // 2)\n        xb = tf.minimum(DIM, x + w // 2)\n        ya = tf.minimum(DIM, y + w // 2)\n        \n        # 替换图像区域\n        image2 = image[k]\n        one = image[j, ya:yb, 0:xa, :]\n        two = image2[ya:yb, xa:xb, :] # 来自图片 k 的补丁\n        three = image[j, ya:yb, xb:DIM, :]\n        \n        # 拼接\n        middle = tf.concat([one, two, three], axis=1)\n        img = tf.concat([image[j, 0:ya, :, :], middle, image[j, yb:DIM, :, :]], axis=0)\n        imgs.append(img)\n        \n        # 计算新标签权重\n        a = tf.cast(w * w / DIM / DIM, tf.float32)\n        # 标签混合：(1-a) * 原图标签 + a * 混合图标签\n        labs.append((1-a) * label[j] + a * label[k])\n        \n    image2 = tf.reshape(tf.stack(imgs), (BATCH_SIZE, DIM, DIM, 3))\n    label2 = tf.reshape(tf.stack(labs), (BATCH_SIZE, CLASSES_NUM))\n    \n    return image2, label2\n# --- 新增代码结束 ---\n# -----------------------------------------------------------\n# 注意：你需要确保在 get_training_dataset 函数里调用了它！\n# -----------------------------------------------------------\ndef get_training_dataset():\n    dataset = load_dataset(TRAINING_FILENAMES, labeled=True)\n    dataset = dataset.map(data_augment, num_parallel_calls=AUTO)\n    # 注意：这里不需要 one_hot，也不需要 cutmix\n    dataset = dataset.repeat()\n    dataset = dataset.shuffle(2048)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTO)\n    return dataset\n\ndef get_validation_dataset(ordered=False):\n    dataset = load_dataset(VALIDATION_FILENAMES, labeled=True, ordered=ordered)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTO)\n    return dataset\n\ndef get_test_dataset(ordered=False):\n    dataset = load_dataset(TEST_FILENAMES, labeled=False, ordered=ordered)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTO)\n    return dataset\n\ndef count_data_items(filenames):\n    # the number of data items is written in the name of the .tfrec\n    # files, i.e. flowers00-230.tfrec = 230 data items\n    n = [int(re.compile(r\"-([0-9]*)\\.\").search(filename).group(1)) for filename in filenames]\n    return np.sum(n)\n\nNUM_TRAINING_IMAGES = count_data_items(TRAINING_FILENAMES)\nNUM_VALIDATION_IMAGES = count_data_items(VALIDATION_FILENAMES)\nNUM_TEST_IMAGES = count_data_items(TEST_FILENAMES)\nprint('Dataset: {} training images, {} validation images, {} unlabeled test images'.format(NUM_TRAINING_IMAGES, NUM_VALIDATION_IMAGES, NUM_TEST_IMAGES))\n","metadata":{"_kg_hide-input":true,"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T11:59:51.293160Z","iopub.execute_input":"2025-12-24T11:59:51.293451Z","iopub.status.idle":"2025-12-24T11:59:51.319065Z","shell.execute_reply.started":"2025-12-24T11:59:51.293430Z","shell.execute_reply":"2025-12-24T11:59:51.318235Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"This next cell will create the datasets that we'll use with Keras during training and inference. Notice how we scale the size of the batches to the number of TPU cores.","metadata":{}},{"cell_type":"code","source":"ds_train = get_training_dataset()\nds_valid = get_validation_dataset()\nds_test = get_test_dataset()\n\nprint(\"Training:\", ds_train)\nprint (\"Validation:\", ds_valid)\nprint(\"Test:\", ds_test)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T11:59:51.320118Z","iopub.execute_input":"2025-12-24T11:59:51.320392Z","iopub.status.idle":"2025-12-24T11:59:52.268908Z","shell.execute_reply.started":"2025-12-24T11:59:51.320373Z","shell.execute_reply":"2025-12-24T11:59:52.268016Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"These datasets are `tf.data.Dataset` objects. You can think about a dataset in TensorFlow as a *stream* of data records. The training and validation sets are streams of `(image, label)` pairs.","metadata":{}},{"cell_type":"code","source":"np.set_printoptions(threshold=15, linewidth=80)\n\nprint(\"Training data shapes:\")\nfor image, label in ds_train.take(3):\n    print(image.numpy().shape, label.numpy().shape)\nprint(\"Training data label examples:\", label.numpy())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T11:59:52.270146Z","iopub.execute_input":"2025-12-24T11:59:52.270485Z","iopub.status.idle":"2025-12-24T12:00:03.960258Z","shell.execute_reply.started":"2025-12-24T11:59:52.270452Z","shell.execute_reply":"2025-12-24T12:00:03.959283Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The test set is a stream of `(image, idnum)` pairs; `idnum` here is the unique identifier given to the image that we'll use later when we make our submission as a `csv` file.","metadata":{}},{"cell_type":"code","source":"print(\"Test data shapes:\")\nfor image, idnum in ds_test.take(3):\n    print(image.numpy().shape, idnum.numpy().shape)\nprint(\"Test data IDs:\", idnum.numpy().astype('U')) # U=unicode string","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T12:00:03.961258Z","iopub.execute_input":"2025-12-24T12:00:03.961531Z","iopub.status.idle":"2025-12-24T12:00:04.673444Z","shell.execute_reply.started":"2025-12-24T12:00:03.961508Z","shell.execute_reply":"2025-12-24T12:00:04.672447Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 4: Explore Data #\n\nLet's take a moment to look at some of the images in the dataset.","metadata":{}},{"cell_type":"code","source":"\nfrom matplotlib import pyplot as plt\n\ndef batch_to_numpy_images_and_labels(data):\n    images, labels = data\n    numpy_images = images.numpy()\n    numpy_labels = labels.numpy()\n    if numpy_labels.dtype == object: # binary string in this case,\n                                     # these are image ID strings\n        numpy_labels = [None for _ in enumerate(numpy_images)]\n    # If no labels, only image IDs, return None for labels (this is\n    # the case for test data)\n    return numpy_images, numpy_labels\n\ndef title_from_label_and_target(label, correct_label):\n    if correct_label is None:\n        return CLASSES[label], True\n    correct = (label == correct_label)\n    return \"{} [{}{}{}]\".format(CLASSES[label], 'OK' if correct else 'NO', u\"\\u2192\" if not correct else '',\n                                CLASSES[correct_label] if not correct else ''), correct\n\ndef display_one_flower(image, title, subplot, red=False, titlesize=16):\n    plt.subplot(*subplot)\n    plt.axis('off')\n    plt.imshow(image)\n    if len(title) > 0:\n        plt.title(title, fontsize=int(titlesize) if not red else int(titlesize/1.2), color='red' if red else 'black', fontdict={'verticalalignment':'center'}, pad=int(titlesize/1.5))\n    return (subplot[0], subplot[1], subplot[2]+1)\n    \ndef display_batch_of_images(databatch, predictions=None):\n    \"\"\"This will work with:\n    display_batch_of_images(images)\n    display_batch_of_images(images, predictions)\n    display_batch_of_images((images, labels))\n    display_batch_of_images((images, labels), predictions)\n    \"\"\"\n    # data\n    images, labels = batch_to_numpy_images_and_labels(databatch)\n    if labels is None:\n        labels = [None for _ in enumerate(images)]\n        \n    # auto-squaring: this will drop data that does not fit into square\n    # or square-ish rectangle\n    rows = int(math.sqrt(len(images)))\n    cols = len(images)//rows\n        \n    # size and spacing\n    FIGSIZE = 13.0\n    SPACING = 0.1\n    subplot=(rows,cols,1)\n    if rows < cols:\n        plt.figure(figsize=(FIGSIZE,FIGSIZE/cols*rows))\n    else:\n        plt.figure(figsize=(FIGSIZE/rows*cols,FIGSIZE))\n    \n    # display\n    for i, (image, label) in enumerate(zip(images[:rows*cols], labels[:rows*cols])):\n        title = '' if label is None else CLASSES[label]\n        correct = True\n        if predictions is not None:\n            title, correct = title_from_label_and_target(predictions[i], label)\n        dynamic_titlesize = FIGSIZE*SPACING/max(rows,cols)*40+3 # magic formula tested to work from 1x1 to 10x10 images\n        subplot = display_one_flower(image, title, subplot, not correct, titlesize=dynamic_titlesize)\n    \n    #layout\n    plt.tight_layout()\n    if label is None and predictions is None:\n        plt.subplots_adjust(wspace=0, hspace=0)\n    else:\n        plt.subplots_adjust(wspace=SPACING, hspace=SPACING)\n    plt.show()\n\n\ndef display_training_curves(training, validation, title, subplot):\n    if subplot%10==1: # set up the subplots on the first call\n        plt.subplots(figsize=(10,10), facecolor='#F0F0F0')\n        plt.tight_layout()\n    ax = plt.subplot(subplot)\n    ax.set_facecolor('#F8F8F8')\n    ax.plot(training)\n    ax.plot(validation)\n    ax.set_title('model '+ title)\n    ax.set_ylabel(title)\n    #ax.set_ylim(0.28,1.05)\n    ax.set_xlabel('epoch')\n    ax.legend(['train', 'valid.'])","metadata":{"_kg_hide-input":true,"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T12:00:04.674703Z","iopub.execute_input":"2025-12-24T12:00:04.675043Z","iopub.status.idle":"2025-12-24T12:00:04.691887Z","shell.execute_reply.started":"2025-12-24T12:00:04.675010Z","shell.execute_reply":"2025-12-24T12:00:04.691018Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"You can display a single batch of images from a dataset with another of our helper functions. The next cell will turn the dataset into an iterator of batches of 20 images.","metadata":{}},{"cell_type":"code","source":"ds_iter = iter(ds_train.unbatch().batch(20))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T12:00:04.695693Z","iopub.execute_input":"2025-12-24T12:00:04.696045Z","iopub.status.idle":"2025-12-24T12:00:04.770501Z","shell.execute_reply.started":"2025-12-24T12:00:04.696024Z","shell.execute_reply":"2025-12-24T12:00:04.769840Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Use the Python `next` function to pop out the next batch in the stream and display it with the helper function.","metadata":{}},{"cell_type":"code","source":"one_batch = next(ds_iter)\ndisplay_batch_of_images(one_batch)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T12:00:04.771474Z","iopub.execute_input":"2025-12-24T12:00:04.771767Z","iopub.status.idle":"2025-12-24T12:00:18.054649Z","shell.execute_reply.started":"2025-12-24T12:00:04.771743Z","shell.execute_reply":"2025-12-24T12:00:18.053760Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"By defining `ds_iter` and `one_batch` in separate cells, you only need to rerun the cell above to see a new batch of images.","metadata":{}},{"cell_type":"markdown","source":"# Step 5: Define Model #\n\nNow we're ready to create a neural network for classifying images! We'll use what's known as **transfer learning**. With transfer learning, you reuse part of a pretrained model to get a head-start on a new dataset.\n\nFor this tutorial, we'll to use a model called **VGG16** pretrained on [ImageNet](http://image-net.org/)). Later, you might want to experiment with [other models](https://www.tensorflow.org/api_docs/python/tf/keras/applications) included with Keras. ([Xception](https://www.tensorflow.org/api_docs/python/tf/keras/applications/Xception) wouldn't be a bad choice.)\n\nThe distribution strategy we created earlier contains a [context manager](https://docs.python.org/3/reference/compound_stmts.html#with), `strategy.scope`. This context manager tells TensorFlow how to divide the work of training among the eight TPU cores. When using TensorFlow with a TPU, it's important to define your model in a `strategy.scope()` context.","metadata":{}},{"cell_type":"code","source":"\n# 引入 EfficientNet\nfrom tensorflow.keras.applications import EfficientNetB5\n\nwith strategy.scope():\n    # 使用 EfficientNetB5 (比 DenseNet201 更准)\n    pretrained_model = EfficientNetB5(\n        weights='imagenet', \n        include_top=False, \n        input_shape=[*IMAGE_SIZE, 3]\n    )\n    \n    # 关键点：一定要解冻！\n    pretrained_model.trainable = True \n    \n    model = tf.keras.Sequential([\n        pretrained_model,\n        tf.keras.layers.GlobalAveragePooling2D(),\n        # 加大 Dropout 防止过拟合\n        tf.keras.layers.Dropout(0.4),\n        tf.keras.layers.Dense(len(CLASSES), activation='softmax')\n    ])\n# 在第174行（model.summary()之前）添加：\ndef label_smoothing_loss(y_true, y_pred, smoothing=0.1):\n    \"\"\"\n    标签平滑损失函数，减少过拟合\n    \"\"\"\n    y_true = tf.cast(y_true, tf.int32)\n    # 将稀疏标签转换为one-hot\n    y_true_one_hot = tf.one_hot(y_true, depth=len(CLASSES))\n    \n    # 应用标签平滑\n    y_true_smoothed = y_true_one_hot * (1 - smoothing) + smoothing / len(CLASSES)\n    \n    # 计算交叉熵损失\n    loss = tf.keras.losses.categorical_crossentropy(y_true_smoothed, y_pred, from_logits=False)\n    return tf.reduce_mean(loss)\n\nprint(\"✅ 标签平滑损失函数已定义\")    \nmodel.compile(\n    optimizer='adam',\n    # 因为没有 one-hot，必须用 sparse\n    loss='sparse_categorical_crossentropy', \n    metrics=['sparse_categorical_accuracy']\n)\nmodel.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T12:00:18.055792Z","iopub.execute_input":"2025-12-24T12:00:18.056048Z","iopub.status.idle":"2025-12-24T12:00:31.900953Z","shell.execute_reply.started":"2025-12-24T12:00:18.056026Z","shell.execute_reply":"2025-12-24T12:00:31.900157Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The `'sparse_categorical'` versions of the loss and metrics are appropriate for a classification task with more than two labels, like this one.","metadata":{}},{"cell_type":"markdown","source":"# Step 6: Training #\n\n## Learning Rate Schedule ##\n\nWe'll train this network with a special learning rate schedule.","metadata":{}},{"cell_type":"code","source":"# 定义动态学习率调度器\n# 这里的关键是：max_lr 会根据你的显卡数量自动调整！\ndef build_lrfn(lr_start=0.00001, lr_max=0.00005, \n               lr_min=0.00001, lr_rampup_epochs=5, \n               lr_sustain_epochs=0, lr_exp_decay=.8):\n    \n    def lrfn(epoch):\n        if epoch < lr_rampup_epochs:\n            # 预热阶段 (Warmup)\n            lr = (lr_max - lr_start) / lr_rampup_epochs * epoch + lr_start\n        elif epoch < lr_rampup_epochs + lr_sustain_epochs:\n            # 保持阶段\n            lr = lr_max\n        else:\n            # 衰减阶段\n            lr = (lr_max - lr_min) * lr_exp_decay**(epoch - lr_rampup_epochs - lr_sustain_epochs) + lr_min\n        return lr\n    return lrfn\n\n#  关键修正：让学习率随显卡数量自动扩大\n# 如果是 GPU T4 x2，这里就会自动变成 0.0001，比原来大一倍，学得更快！\nlr_fn = build_lrfn(\n    lr_start=0.00001,\n    lr_max=0.00005 * strategy.num_replicas_in_sync, \n    lr_min=0.00001,\n    lr_rampup_epochs=5,\n    lr_sustain_epochs=0,\n    lr_exp_decay=0.8\n)\n\nlr_callback = tf.keras.callbacks.LearningRateScheduler(lr_fn, verbose=1)\n\n# 画出来确认一下\nrng = [i for i in range(25)]\ny = [lr_fn(x) for x in rng]\nplt.plot(rng, y)\nplt.title(\"Learning Rate Schedule\")\nplt.show()","metadata":{"_kg_hide-input":true,"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T12:00:31.902231Z","iopub.execute_input":"2025-12-24T12:00:31.902794Z","iopub.status.idle":"2025-12-24T12:00:32.057678Z","shell.execute_reply.started":"2025-12-24T12:00:31.902759Z","shell.execute_reply":"2025-12-24T12:00:32.056879Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Fit Model ##\n\nAnd now we're ready to train the model. After defining a few parameters, we're good to go!","metadata":{}},{"cell_type":"code","source":"# Define training epochs\n# 建议增加到 20 或 25\nEPOCHS = 20 \nSTEPS_PER_EPOCH = NUM_TRAINING_IMAGES // BATCH_SIZE\n# 在第214行（history = model.fit之前）添加：\nimport os\nos.makedirs('/kaggle/working/checkpoints', exist_ok=True)\n\n# 改进的回调函数\ncallbacks = [\n    # 学习率调度\n    lr_callback,\n    \n    # 早停\n    tf.keras.callbacks.EarlyStopping(\n        monitor='val_sparse_categorical_accuracy',\n        patience=10,\n        restore_best_weights=True,\n        verbose=1,\n        mode='max'\n    ),\n    \n    # 模型检查点\n    tf.keras.callbacks.ModelCheckpoint(\n        '/kaggle/working/checkpoints/best_model.h5',\n        monitor='val_sparse_categorical_accuracy',\n        save_best_only=True,\n        save_weights_only=True,\n        mode='max',\n        verbose=1\n    ),\n    \n    # 学习率衰减\n    tf.keras.callbacks.ReduceLROnPlateau(\n        monitor='val_loss',\n        factor=0.5,\n        patience=3,\n        min_lr=1e-7,\n        verbose=1\n    ),\n    \n    # CSV记录器\n    tf.keras.callbacks.CSVLogger('/kaggle/working/training_log.csv'),\n]\n\nprint(\"✅ 回调函数已配置\")\n\n# 修改第215行的训练代码为：\nhistory = model.fit(\n    ds_train,\n    validation_data=ds_valid,\n    epochs=EPOCHS,\n    steps_per_epoch=STEPS_PER_EPOCH,\n    callbacks=callbacks,  # 使用新的回调函数列表\n    verbose=1\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T12:00:32.058711Z","iopub.execute_input":"2025-12-24T12:00:32.058952Z","iopub.status.idle":"2025-12-24T13:51:35.693023Z","shell.execute_reply.started":"2025-12-24T12:00:32.058931Z","shell.execute_reply":"2025-12-24T13:51:35.691731Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"This next cell shows how the loss and metrics progressed during training. Thankfully, it converges!","metadata":{}},{"cell_type":"code","source":"display_training_curves(\n    history.history['loss'],\n    history.history['val_loss'],\n    'loss',\n    211,\n)\ndisplay_training_curves(\n    history.history['sparse_categorical_accuracy'],\n    history.history['val_sparse_categorical_accuracy'],\n    'accuracy',\n    212,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T13:51:35.695569Z","iopub.execute_input":"2025-12-24T13:51:35.696296Z","iopub.status.idle":"2025-12-24T13:51:36.175197Z","shell.execute_reply.started":"2025-12-24T13:51:35.696260Z","shell.execute_reply":"2025-12-24T13:51:36.174350Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 7: Evaluate Predictions #\n\nBefore making your final predictions on the test set, it's a good idea to evaluate your model's predictions on the validation set. This can help you diagnose problems in training or suggest ways your model could be improved. We'll look at two common ways of validation: plotting the **confusion matrix** and **visual validation**.","metadata":{}},{"cell_type":"code","source":"\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import f1_score, precision_score, recall_score, confusion_matrix\n\ndef display_confusion_matrix(cmat, score, precision, recall):\n    plt.figure(figsize=(15,15))\n    ax = plt.gca()\n    ax.matshow(cmat, cmap='Reds')\n    ax.set_xticks(range(len(CLASSES)))\n    ax.set_xticklabels(CLASSES, fontdict={'fontsize': 7})\n    plt.setp(ax.get_xticklabels(), rotation=45, ha=\"left\", rotation_mode=\"anchor\")\n    ax.set_yticks(range(len(CLASSES)))\n    ax.set_yticklabels(CLASSES, fontdict={'fontsize': 7})\n    plt.setp(ax.get_yticklabels(), rotation=45, ha=\"right\", rotation_mode=\"anchor\")\n    titlestring = \"\"\n    if score is not None:\n        titlestring += 'f1 = {:.3f} '.format(score)\n    if precision is not None:\n        titlestring += '\\nprecision = {:.3f} '.format(precision)\n    if recall is not None:\n        titlestring += '\\nrecall = {:.3f} '.format(recall)\n    if len(titlestring) > 0:\n        ax.text(101, 1, titlestring, fontdict={'fontsize': 18, 'horizontalalignment':'right', 'verticalalignment':'top', 'color':'#804040'})\n    plt.show()\n    \ndef display_training_curves(training, validation, title, subplot):\n    if subplot%10==1: # set up the subplots on the first call\n        plt.subplots(figsize=(10,10), facecolor='#F0F0F0')\n        plt.tight_layout()\n    ax = plt.subplot(subplot)\n    ax.set_facecolor('#F8F8F8')\n    ax.plot(training)\n    ax.plot(validation)\n    ax.set_title('model '+ title)\n    ax.set_ylabel(title)\n    #ax.set_ylim(0.28,1.05)\n    ax.set_xlabel('epoch')\n    ax.legend(['train', 'valid.'])","metadata":{"_kg_hide-input":true,"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T13:51:36.176460Z","iopub.execute_input":"2025-12-24T13:51:36.177077Z","iopub.status.idle":"2025-12-24T13:51:36.564116Z","shell.execute_reply.started":"2025-12-24T13:51:36.177042Z","shell.execute_reply":"2025-12-24T13:51:36.563155Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Confusion Matrix ##\n\nA [confusion matrix](https://en.wikipedia.org/wiki/Confusion_matrix) shows the actual class of an image tabulated against its predicted class. It is one of the best tools you have for evaluating the performance of a classifier.\n\nThe following cell does some processing on the validation data and then creates the matrix with the `confusion_matrix` function included in [`scikit-learn`](https://scikit-learn.org/stable/index.html).","metadata":{}},{"cell_type":"code","source":"cmdataset = get_validation_dataset(ordered=True)\nimages_ds = cmdataset.map(lambda image, label: image)\nlabels_ds = cmdataset.map(lambda image, label: label).unbatch()\n\ncm_correct_labels = next(iter(labels_ds.batch(NUM_VALIDATION_IMAGES))).numpy()\ncm_probabilities = model.predict(images_ds)\ncm_predictions = np.argmax(cm_probabilities, axis=-1)\n\nlabels = range(len(CLASSES))\ncmat = confusion_matrix(\n    cm_correct_labels,\n    cm_predictions,\n    labels=labels,\n)\ncmat = (cmat.T / cmat.sum(axis=1)).T # normalize","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T13:51:36.566683Z","iopub.execute_input":"2025-12-24T13:51:36.571146Z","iopub.status.idle":"2025-12-24T13:52:13.583682Z","shell.execute_reply.started":"2025-12-24T13:51:36.571092Z","shell.execute_reply":"2025-12-24T13:52:13.582930Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"You might be familiar with metrics like [F1-score](https://en.wikipedia.org/wiki/F1_score) or [precision and recall](https://en.wikipedia.org/wiki/Precision_and_recall). This cell will compute these metrics and display them with a plot of the confusion matrix. (These metrics are defined in the Scikit-learn module `sklearn.metrics`; we've imported them in the helper script for you.)","metadata":{}},{"cell_type":"code","source":"score = f1_score(\n    cm_correct_labels,\n    cm_predictions,\n    labels=labels,\n    average='macro',\n)\nprecision = precision_score(\n    cm_correct_labels,\n    cm_predictions,\n    labels=labels,\n    average='macro',\n)\nrecall = recall_score(\n    cm_correct_labels,\n    cm_predictions,\n    labels=labels,\n    average='macro',\n)\ndisplay_confusion_matrix(cmat, score, precision, recall)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T13:52:13.584789Z","iopub.execute_input":"2025-12-24T13:52:13.585047Z","iopub.status.idle":"2025-12-24T13:52:15.205225Z","shell.execute_reply.started":"2025-12-24T13:52:13.585025Z","shell.execute_reply":"2025-12-24T13:52:15.204373Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visual Validation ##\n\nIt can also be helpful to look at some examples from the validation set and see what class your model predicted. This can help reveal patterns in the kinds of images your model has trouble with.\n\nThis cell will set up the validation set to display 20 images at a time -- you can change this to display more or fewer, if you like.","metadata":{}},{"cell_type":"code","source":"dataset = get_validation_dataset()\ndataset = dataset.unbatch().batch(20)\nbatch = iter(dataset)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T13:52:15.206384Z","iopub.execute_input":"2025-12-24T13:52:15.206720Z","iopub.status.idle":"2025-12-24T13:52:15.272190Z","shell.execute_reply.started":"2025-12-24T13:52:15.206697Z","shell.execute_reply":"2025-12-24T13:52:15.271543Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"And here is a set of flowers with their predicted species. Run the cell again to see another set.","metadata":{}},{"cell_type":"code","source":"images, labels = next(batch)\nprobabilities = model.predict(images)\npredictions = np.argmax(probabilities, axis=-1)\ndisplay_batch_of_images((images, labels), predictions)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T13:52:15.273080Z","iopub.execute_input":"2025-12-24T13:52:15.273348Z","iopub.status.idle":"2025-12-24T13:52:18.832935Z","shell.execute_reply.started":"2025-12-24T13:52:15.273327Z","shell.execute_reply":"2025-12-24T13:52:18.831599Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 8: Make Test Predictions #\n\nOnce you're satisfied with everything, you're ready to make predictions on the test set.","metadata":{}},{"cell_type":"code","source":"import gc\n# 在第294行（import gc之后）添加：\ndef test_time_augmentation(model, dataset, n_aug=5):\n    \"\"\"\n    测试时增强：对每个测试图像应用多种增强，然后平均预测结果\n    \"\"\"\n    all_predictions = []\n    \n    for i in range(n_aug):\n        # 对数据集应用不同的增强\n        if i == 0:\n            # 原始图像\n            aug_dataset = dataset\n        elif i == 1:\n            # 水平翻转\n            aug_dataset = dataset.map(lambda img, idnum: (tf.image.flip_left_right(img), idnum))\n        elif i == 2:\n            # 垂直翻转\n            aug_dataset = dataset.map(lambda img, idnum: (tf.image.flip_up_down(img), idnum))\n        elif i == 3:\n            # 旋转90度\n            aug_dataset = dataset.map(lambda img, idnum: (tf.image.rot90(img), idnum))\n        elif i == 4:\n            # 旋转270度\n            aug_dataset = dataset.map(lambda img, idnum: (tf.image.rot90(img, k=3), idnum))\n        else:\n            # 随机亮度调整\n            aug_dataset = dataset.map(lambda img, idnum: \n                (tf.clip_by_value(tf.image.random_brightness(img, max_delta=0.1), 0.0, 1.0), idnum))\n        \n        # 预测\n        test_images_ds = aug_dataset.map(lambda image, idnum: image)\n        probabilities = model.predict(test_images_ds, verbose=0)\n        all_predictions.append(probabilities)\n    \n    # 平均所有增强的预测\n    avg_predictions = np.mean(all_predictions, axis=0)\n    final_predictions = np.argmax(avg_predictions, axis=-1)\n    \n    return final_predictions\n\nprint(\"✅ 测试时增强函数已定义\")\n# 1. 此时模型已经训练完毕，不再需要训练集和验证集了\n# 删除这些巨大的数据集对象，释放 RAM\ntry:\n    del ds_train\n    del ds_valid\nexcept NameError:\n    pass # 防止变量名不存在报错\n\n# 2. 强制运行 Python 的垃圾回收机制\ngc.collect()\n\nprint(\"✅ 内存清理完成！旧的训练数据已释放，准备加载测试集进行预测...\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T13:52:18.834255Z","iopub.execute_input":"2025-12-24T13:52:18.834529Z","iopub.status.idle":"2025-12-24T13:52:20.234235Z","shell.execute_reply.started":"2025-12-24T13:52:18.834507Z","shell.execute_reply":"2025-12-24T13:52:20.233344Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_ds = get_test_dataset(ordered=True)\n\nprint('Computing predictions...')\ntest_images_ds = test_ds.map(lambda image, idnum: image)\nprobabilities = model.predict(test_images_ds)\npredictions = np.argmax(probabilities, axis=-1)\nprint(predictions)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T13:52:20.235689Z","iopub.execute_input":"2025-12-24T13:52:20.236309Z","iopub.status.idle":"2025-12-24T13:53:06.453765Z","shell.execute_reply.started":"2025-12-24T13:52:20.236267Z","shell.execute_reply":"2025-12-24T13:53:06.452838Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We'll generate a file `submission.csv`. This file is what you'll submit to get your score on the leaderboard.","metadata":{}},{"cell_type":"code","source":"print('Generating submission.csv file...')\n\n# Get image ids from test set and convert to unicode\ntest_ids_ds = test_ds.map(lambda image, idnum: idnum).unbatch()\ntest_ids = next(iter(test_ids_ds.batch(NUM_TEST_IMAGES))).numpy().astype('U')\n\n# Write the submission file\nnp.savetxt(\n    'submission.csv',\n    np.rec.fromarrays([test_ids, predictions]),\n    fmt=['%s', '%d'],\n    delimiter=',',\n    header='id,label',\n    comments='',\n)\n\n# Look at the first few predictions\n!head submission.csv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T13:53:06.455094Z","iopub.execute_input":"2025-12-24T13:53:06.455463Z","iopub.status.idle":"2025-12-24T13:53:18.240014Z","shell.execute_reply.started":"2025-12-24T13:53:06.455418Z","shell.execute_reply":"2025-12-24T13:53:18.238881Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 9: Make a submission #\n\nIf you haven't already, create your own editable copy of this notebook by clicking on the **Copy and Edit** button in the top right corner. Then, submit to the competition by following these steps:\n\n1. Begin by clicking on the blue **Save Version** button in the top right corner of the window.  This will generate a pop-up window.  \n2. Ensure that the **Save and Run All** option is selected, and then click on the blue **Save** button.\n3. This generates a window in the bottom left corner of the notebook.  After it has finished running, click on the number to the right of the **Save Version** button.  This pulls up a list of versions on the right of the screen.  Click on the ellipsis **(...)** to the right of the most recent version, and select **Open in Viewer**.  This brings you into view mode of the same page. You will need to scroll down to get back to these instructions.\n4. Click on the **Output** tab on the right of the screen.  Then, click on the file you would like to submit, and click on the blue **Submit** button to submit your results to the leaderboard.\n\nYou have now successfully submitted to the competition!\n\nIf you want to keep working to improve your performance, select the blue **Edit** button in the top right of the screen. Then you can change your code and repeat the process. There's a lot of room to improve, and you will climb up the leaderboard as you work.\n","metadata":{}},{"cell_type":"markdown","source":"---\n\n\n\n\n*Have questions or comments? Visit the [Learn Discussion forum](https://www.kaggle.com/learn-forum/161321) to chat with other Learners.*","metadata":{}}]}