{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.15","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":21154,"databundleVersionId":1243559,"sourceType":"competition"}],"dockerImageVersionId":30788,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction\n\n任务：搭建机器学习模型构建对104种花朵正确识别\n\n使用平台提供的TPU资源，每周20小时。","metadata":{}},{"cell_type":"markdown","source":"#  Step 1: \n\n下载efficientnet、导入需要的库、打印输出Tensorflow的版本，还可以打印输出全部文件路径","metadata":{}},{"cell_type":"code","source":"!pip install -q efficientnet\n\nimport math, re, os\nimport numpy as np\nimport tensorflow as tf\nimport pandas as pd\nfrom matplotlib import pyplot as plt\nimport efficientnet.tfkeras as efn\nfrom sklearn.metrics import f1_score, precision_score, recall_score, confusion_matrix\nprint(\"Tensorflow version \" + tf.__version__)","metadata":{"execution":{"iopub.status.busy":"2024-10-29T11:56:17.324424Z","iopub.execute_input":"2024-10-29T11:56:17.324774Z","iopub.status.idle":"2024-10-29T11:56:38.637972Z","shell.execute_reply.started":"2024-10-29T11:56:17.324733Z","shell.execute_reply":"2024-10-29T11:56:38.636856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 可对输入中的全部文件路径打印输出\n# 先注释掉，否则结果太长了\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#    for filename in filenames:\n#        print(os.path.join(dirname, filename))","metadata":{"execution":{"iopub.status.busy":"2024-10-29T11:56:57.151619Z","iopub.execute_input":"2024-10-29T11:56:57.152309Z","iopub.status.idle":"2024-10-29T11:56:57.250523Z","shell.execute_reply.started":"2024-10-29T11:56:57.152268Z","shell.execute_reply":"2024-10-29T11:56:57.249443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step 2: 检测并配置TPU环境\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":"try:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()  # Detect TPU\n    print('Running on TPU:', tpu.master())\nexcept ValueError:\n    tpu = None\n\nif tpu:\n    tf.config.experimental_connect_to_cluster(tpu) # 连接到tpu集群\n    tf.tpu.experimental.initialize_tpu_system(tpu) # 连接到tpu系统\n    strategy = tf.distribute.TPUStrategy(tpu) # 创建tpu分布式策略\nelse:\n    strategy = tf.distribute.get_strategy() \n    \n\n    \nAUTO = tf.data.experimental.AUTOTUNE # 让程序自动选择最优的线程并行个数\n\nprint(\"REPLICAS: \", strategy.num_replicas_in_sync) #输出副本数","metadata":{"execution":{"iopub.status.busy":"2024-10-23T15:06:59.405736Z","iopub.execute_input":"2024-10-23T15:06:59.406125Z","iopub.status.idle":"2024-10-23T15:07:08.458688Z","shell.execute_reply.started":"2024-10-23T15:06:59.406098Z","shell.execute_reply":"2024-10-23T15:07:08.457886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step 3: 获取竞赛数据集路径并简单设置\nGet GCS Path\nWhen used with TPUs, datasets need to be stored in a Google Cloud Storage bucket. 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":"markdown","source":"# Load Data\n\nWhen used with TPUs, datasets are often serialized into TFRecords. 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":"from kaggle_datasets import KaggleDatasets\n\n# 数据集已放入INPUT\n# /kaggle/input/tpu-getting-started/tfrecords-jpeg-512x512/train/10-512x512-798.tfrec 以512大小的训练集为例，这是每一个tfrec格式的路径\n# 其中GCS_DS_PATH即为 /kaggle/input/tpu-getting-started，此路径下有四个不同大小的数据集（各个数据集下分训练集、验证集、测试集）\n# 各个训练集、验证集、测试集下的tfrec格式文件包含了多个（图片加标签）信息，其中 10-512x512-798.tfrec，就表示第10组，大小为512*512，的798个（图片加标签信息）\nGCS_DS_PATH = KaggleDatasets().get_gcs_path('tpu-getting-started')\nprint(GCS_DS_PATH) \n\nTRAINING_FILENAMES = tf.io.gfile.glob(GCS_DS_PATH + '/tfrecords-jpeg-512x512' + '/train/*.tfrec') # 512大小的训练集路径\nVALIDATION_FILENAMES = tf.io.gfile.glob(GCS_DS_PATH + '/tfrecords-jpeg-512x512' + '/val/*.tfrec') # 512大小的验证集路径","metadata":{"execution":{"iopub.status.busy":"2024-10-23T15:07:08.459777Z","iopub.execute_input":"2024-10-23T15:07:08.460044Z","iopub.status.idle":"2024-10-23T15:07:08.468619Z","shell.execute_reply.started":"2024-10-23T15:07:08.460017Z","shell.execute_reply":"2024-10-23T15:07:08.467944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 设置参数","metadata":{}},{"cell_type":"code","source":"IMAGE_SIZE = [512,512] # at this size, a GPU will run out of memory. Use the TPU\nEPOCHS = 20\nBATCH_SIZE = 20 * strategy.num_replicas_in_sync\n\n# 无论数据集图像大小，对应的三个集合内的数据量相等\nNUM_TRAINING_IMAGES = 12753\nNUM_TEST_IMAGES = 7382\nNUM_VALIDATION_IMAGES = 3712\n\n# 104种花的名称，训练集、验证集内标签为数字序号，能够与这个CLASSES对应\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']\n\nSTEPS_PER_EPOCH = NUM_TRAINING_IMAGES // BATCH_SIZE","metadata":{"execution":{"iopub.status.busy":"2024-10-23T15:07:08.470329Z","iopub.execute_input":"2024-10-23T15:07:08.470579Z","iopub.status.idle":"2024-10-23T15:07:08.480006Z","shell.execute_reply.started":"2024-10-23T15:07:08.470553Z","shell.execute_reply":"2024-10-23T15:07:08.479347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step 4：数据处理和获取\n\n函数包括：1、图像数据解码\n2、读取带种类标签的tfrecord数据（训练集验证集）\n3、读取不带种类标签但带有id的tfrecord数据（测试集）\n4、对图像简单的增强（翻转等），丰富训练集，增强模型鲁棒性\n5、将训练集和验证集合并作为训练集\n6、加载数据集\n7、加载训练数据集\n8、加载验证数据集\n9、加载测试数据集","metadata":{}},{"cell_type":"code","source":"# 1\ndef decode_image(image_data):\n    image = tf.image.decode_jpeg(image_data, channels=3)\n    image = tf.cast(image, tf.float32) / 255.0  # convert image to floats in [0, 1] range\n    image = tf.reshape(image, [*IMAGE_SIZE, 3]) # explicit size needed for TPU\n    return image\n\n# 2\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\n# 3\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\n\n# 4\n# 按水平 (从左向右) 随机翻转图像.返回图片的参数image和label\ndef data_augment(image, label, seed=2020):\n    # TensorFlow函数：tf.image.random_flip_left_right\n    # 按水平 (从左向右) 随机翻转图像.\n    # 以1比2的概率,输出image沿着第二维翻转的内容,即,width.否则按原样输出图像.\n    # 参数：\n    # image：形状为[height, width, channels]的三维张量.\n    # seed：一个Python整数,用于创建一个随机种子.查看tf.set_random_seed行为.\n    # 返回：一个与image具有相同类型和形状的三维张量.\n    image = tf.image.random_flip_left_right(image, seed=seed)\n#    image = tf.image.random_jpeg_quality(image, 85, 100, seed=seed)\n#     image = tf.image.random_flip_up_down(image, seed=seed)\n#     image = tf.image.random_brightness(image, 0.1, seed=seed)\n#   image = tf.image.random_jpeg_quality(image, 85, 100, seed=seed)\n#     image = tf.image.random_saturation(image, 0, 2)\n    return image, label  \n\n\n# 5\n# 将训练集和验证集合并\ndef get_train_valid_datasets():\n    dataset = load_dataset(TRAINING_FILENAMES, labeled=True)\n    # + VALIDATION_FILENAMES\n    # 将数据转换并行化\n    # 加载训练集，第一个参数为训练集路径，第二个参数表示有标签\n    dataset = dataset.map(data_augment, num_parallel_calls=AUTO)\n    # 重复此数据集count次数\n    # 函数形式：repeat(count=None)\n    # 参数count:(可选）表示数据集应重复的次数。默认行为（如果count是None或-1）是无限期重复的数据集。\n    dataset = dataset.repeat() # 数据集必须重复几个轮次\n    dataset = dataset.shuffle(2048) # 将数据打乱，括号中数值越大，混乱程度越大\n    dataset = dataset.batch(BATCH_SIZE)\n    # pipeline（管道）读取数据，在训练时预取下一批（自动调整预取缓冲区大小）\n    dataset = dataset.prefetch(AUTO)\n    return dataset\n\n# 6\ndef load_dataset(filenames, labeled=True, ordered=False):\n    ignore_order = tf.data.Options()\n    ignore_order.experimental_deterministic = False  # Allow non-deterministic order for TPU efficiency\n    \n    dataset = tf.data.TFRecordDataset(filenames)\n    dataset = dataset.with_options(ignore_order)\n    dataset = dataset.map(\n        read_labeled_tfrecord if labeled else read_unlabeled_tfrecord,\n        num_parallel_calls=tf.data.AUTOTUNE\n    )\n    return dataset\n\n# 7\ndef get_training_dataset():\n    dataset = load_dataset(\n        tf.io.gfile.glob(GCS_DS_PATH + '/tfrecords-jpeg-512x512/train/*.tfrec'), \n        labeled=True\n    )\n    dataset = dataset.shuffle(2048)  # Shuffle before repeating\n    dataset = dataset.repeat()  # Ensure enough data for all epochs\n    dataset = dataset.batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)\n    return dataset\n\n# 8\ndef get_validation_dataset():\n    dataset = load_dataset(\n        tf.io.gfile.glob(GCS_DS_PATH + '/tfrecords-jpeg-512x512/val/*.tfrec'), \n        labeled=True\n    )\n    dataset = dataset.batch(BATCH_SIZE).cache().prefetch(tf.data.AUTOTUNE)\n    return dataset\n# 9\ndef get_test_dataset(ordered=False):\n    dataset = load_dataset(tf.io.gfile.glob(GCS_DS_PATH + '/tfrecords-jpeg-512x512/test/*.tfrec'), labeled=False, ordered=ordered)\n    dataset = dataset.batch(BATCH_SIZE)\n    return dataset","metadata":{"execution":{"iopub.status.busy":"2024-10-23T15:07:08.480889Z","iopub.execute_input":"2024-10-23T15:07:08.481106Z","iopub.status.idle":"2024-10-23T15:07:08.658481Z","shell.execute_reply.started":"2024-10-23T15:07:08.481084Z","shell.execute_reply":"2024-10-23T15:07:08.657784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step 5: 图像可视化\n\nLet's take a moment to look at some of the images in the dataset.\n\n这些函数可用于：\n1、展示训练集及其标签\n2、展示测试集合\n3、展示验证集、判断结果以及是否正确","metadata":{}},{"cell_type":"code","source":"# 设置numpy数组基本属性，设置显示15个数字，用于插入换行符的每行字符数（默认为75）。\n# 当数组数目过大时，设置显示几个数字，其余用省略号\n# 用于插入换行符的每行字符数（默认为75）。\nnp.set_printoptions(threshold=15, linewidth=80)\n \n# 将小批量图片和标签处理为numpy向量格式\ndef batch_to_numpy_images_and_labels(data):\n    images, labels = data \n    numpy_images = images.numpy() # 将图像转换为numpy向量格式\n    numpy_labels = labels.numpy() # 将label标签转换为numpy向量格式\n    if numpy_labels.dtype == object: # 在这种情况下为二进制字符串，它们是图像ID字符串\n        numpy_labels = [None for _ in enumerate(numpy_images)]\n    # 如果没有标签，只有图像ID，则对标签返回None（测试数据就是这种情况）\n    return numpy_images, numpy_labels\n \n# 把实际类型和模型预测出来的模型一起显示在图片上方，这是用给验证集的，当对验证集预测完标签后和验证集的实际标签进行比较\n# label,图片中花朵的实际类别\n# current_label，当前我们预测的类别\ndef title_from_label_and_target(label, current_label):\n    # 如果没有预测的类别，则返回实际类别，比如训练集\n    if current_label is None:\n        return CLASSES[label], True\n    current = (label == current_label) # 判断一下实际类别和我们预测的类别是否一致\n    # 如果一致，则返回OK，不一致则返回NO加实际类别\n    return \"{} [{}{}{}]\".format(CLASSES[label], 'OK' if current else 'NO', u\"\\u2192\" if not current else '',\n                                CLASSES[current_label] if not current else ''), current\n \n# 绘制一朵花\ndef display_one_flower(image, title, subplot, red=False, titlesize=16):\n    plt.subplot(*subplot)\n    plt.axis('off') # 不显示坐标尺寸\n    plt.imshow(image) # 函数负责对图像进行处理，并显示其格式；而plt.show()则是将plt.imshow()处理后的函数显示出来。\n    if len(title) > 0:\n        #绘制图片的标题\n        plt.title(title, fontsize=int(titlesize) if not red else int(titlesize/1.2), color='red' if red else 'black', \n                  fontdict={'verticalalignment':'center'}, pad=int(titlesize/1.5))\n    return (subplot[0], subplot[1], subplot[2]+1)\n    \n\n# 展示小批量图片，我们在下面的代码中经常展示20张照片\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    # 读取图片和实际标签数据，而且这些数据被转换成numpy向量的格式\n    images, labels = batch_to_numpy_images_and_labels(databatch)\n    # 如果没有实际标签，比如测试集，那么我们需要将labels变量设为每个元素都为none\n    if labels is None:\n        labels = [None for _ in enumerate(images)]\n        \n    # 删除不适合矩形的数据，即一次只显示正好满足矩形数量的图片\n    rows = int(math.sqrt(len(images)))\n    cols = len(images) // rows\n        \n    # 画布大小和间距\n    FIGSIZE = 13.0\n    SPACING = 0.1\n    subplot=(rows,cols,1)\n    if rows < cols:\n        # 如果行大于列\n        plt.figure(figsize=(FIGSIZE, FIGSIZE / cols * rows))\n    else:\n        plt.figure(figsize=(FIGSIZE / rows * cols, FIGSIZE))\n    \n    # 显示\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 # 经过测试可以在1x1到10x10图像上工作的魔术公式\n        subplot = display_one_flower(image, title, subplot, not correct, titlesize=dynamic_titlesize)\n    \n    # 布局\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()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"接下来加载训练集、测试集、验证集，但仅用于可视化。打印输出训练集及其标签示例、验证集及其标签示例、测试集及其id示例","metadata":{}},{"cell_type":"code","source":"# 注意此时的的加载仅用于可视化（后续训练、验证、预测的加载操作与之分开）\ndss_train = get_training_dataset()\ntests_ds = get_test_dataset(ordered=True)\ndss_valid = get_validation_dataset()\n\n\n# 训练集数据展示\nprint(\"Training data shapes:\")\n# 输出训练集前3个小批量的图像数据形状、标签形状\nfor image, label in dss_train.take(3):\n    print(image.numpy().shape, label.numpy().shape)\n\n    # 训练数据标签示例\nprint(\"Training data label examples:\", label.numpy())\n\n\n\n# 验证集数据展示\nprint(\"Validation data shapes:\")\n# 输出验证集前3个小批量的图像数据形状、标签形状\nfor image, label in dss_valid.take(3):\n    print(image.numpy().shape, label.numpy().shape)\n    \n    # 验证数据标签示例\nprint(\"Validation data label examples:\", label.numpy())\n \n    \n # 测试集数据展示   \nprint(\"Test data shapes:\")\n# 输出测试集前3个小批量的图像数据形状、标签形状\nfor image, idnum in tests_ds.take(3):\n    print(image.numpy().shape, idnum.numpy().shape)\n    \n    # 测试集的id示例\nprint(\"Test data IDs:\", idnum.numpy().astype('U')) # U=unicode string","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"接下来对训练集和测试集可视化，但不对验证集可视化","metadata":{}},{"cell_type":"code","source":"# 训练集可视化\ntrset = dss_train.unbatch().batch(20) # 将训练集分成大小为20的小批量\ntrbatch = iter(trset) # 首先获得Iterator对象","metadata":{"execution":{"iopub.status.busy":"2024-10-27T16:02:55.110589Z","iopub.execute_input":"2024-10-27T16:02:55.110821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 运行该单元格以获取训练集合下一组图像，并绘图展示，同时还能够显示花朵的种类\ndisplay_batch_of_images(next(trbatch))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 测试集可视化\ntesset = tests_ds.unbatch().batch(20) # 将训练集分成大小为20的小批量\ntesbatch = iter(tesset) # 首先获得Iterator对象","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 再次运行该单元格以获取下一组图像，并绘图展示，不显示标签和id\ndisplay_batch_of_images(next(tesbatch))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step 5: Define Model\nFor this problem, we'll use a model called efficientnetB7.","metadata":{}},{"cell_type":"code","source":"# optimizer='adam'\n# Model creation and training within strategy scope\nwith strategy.scope():\n    enet = efn.EfficientNetB7(\n    input_shape=(512, 512, 3),\n        weights='imagenet',\n        include_top=False\n    )\n\n    model = tf.keras.Sequential([\n        enet,\n        tf.keras.layers.GlobalAveragePooling2D(),\n        tf.keras.layers.Dense(104, activation='softmax')\n    ])\n\n    model.compile(\n        optimizer=tf.keras.optimizers.Adam(),\n        loss='sparse_categorical_crossentropy',\n        metrics=['sparse_categorical_accuracy']\n    )\n\n    model.summary()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step 6: Training\nThe '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":"code","source":"# 加载数据集用于训练、验证和测试\n# 注意get_training_dataset()只加载训练集且并不翻转增强，而get_train_valid_datasets()会将训练集和验证集混合作为训练集（并翻转增强）\nds_train = get_train_valid_datasets()\nds_valid = get_validation_dataset()\ntest_ds = get_test_dataset(ordered=True)\n\n\n# Train the model\nhistorical = model.fit(\n        ds_train,\n        steps_per_epoch=STEPS_PER_EPOCH,\n        epochs=EPOCHS,\n        validation_data=ds_valid\n)\n\nhistorical.history.keys()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step7：基于训练结果，评估模型","metadata":{}},{"cell_type":"code","source":"# 绘制混淆矩阵的函数\ndef display_confusion_matrix(cmat, score, precision, recall):\n    plt.figure(figsize=(15,15))  # 设置画布大小\n    ax = plt.gca() # 返回当前axes(matplotlib.axes.Axes) 获取当前子图\n    ax.matshow(cmat, cmap='Reds') # 绘制矩阵\n    ax.set_xticks(range(104))  # 根据花朵类别数（其实就是104）设置x轴范围\n    ax.set_xticklabels(CLASSES, fontdict={'fontsize': 7}) # 设置x轴下标字体的大小\n    plt.setp(ax.get_xticklabels(), rotation=45, ha=\"left\", rotation_mode=\"anchor\") # 更换x轴下标角度\n    ax.set_yticks(range(104))  # 根据花朵类别数（其实就是104）设置y轴范围\n    ax.set_yticklabels(CLASSES, fontdict={'fontsize': 7}) # 设置y轴下标字体的大小\n    plt.setp(ax.get_yticklabels(), rotation=45, ha=\"right\", rotation_mode=\"anchor\") # 更换y轴下标角度\n    titlestring = \"\"\n    if score is not None:\n        titlestring += 'f1 = {:.3f} '.format(score) # 更改格式为有3位小数的浮点数\n    if precision is not None:\n        titlestring += '\\nprecision = {:.3f} '.format(precision) # 更改格式为有3位小数的浮点数\n    if recall is not None:\n        titlestring += '\\nrecall = {:.3f} '.format(recall) # 更改格式为有3位小数的浮点数\n    if len(titlestring) > 0:\n        ax.text(101, 1, titlestring, fontdict={'fontsize': 18, 'horizontalalignment':'right', 'verticalalignment':'top', 'color':'#804040'}) #添加文本注释\n    plt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 绘制模型各训练轮次的loss、categorical_accuracy曲线\nimport pandas as pd\nhistory_frame = pd.DataFrame( historical.history)\nhistory_frame.loc[:, ['loss', 'val_loss']].plot()\nhistory_frame.loc[:, ['sparse_categorical_accuracy', 'val_sparse_categorical_accuracy']].plot();","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 进行验证集的预测类别和真实类别的对比，分别输出，并对验证集可视化\n# 因为我们要分割数据集并分别对图像和标签进行迭代，所以顺序很重要。\ncmdataset = get_validation_dataset()  # 验证集\nimages_ds = cmdataset.map(lambda image, label: image)  # 图像集\nlabels_ds = cmdataset.map(lambda image, label: label).unbatch() # 标签集 \ncm_correct_labels = next(iter(labels_ds.batch(NUM_VALIDATION_IMAGES))).numpy() # get everything as one batch\ncm_probabilities = model.predict(images_ds) # 图片在104个类别上的概率\ncm_predictions = np.argmax(cm_probabilities, axis=-1) # 其中最大的概率表示这个图片的预测类别\n\nprint(\"Correct   labels: \", cm_correct_labels.shape, cm_correct_labels) # 输出正确（实际）标签的形状、输出正确标签 \nprint(\"Predicted labels: \", cm_predictions.shape, cm_predictions) # 输出预测标签的形状、输出预测标签\n\n# 验证集可视化\ndataset = ds_valid.unbatch().batch(20)  #将验证集分成大小为20的小批量\nbatch = iter(dataset) # 将数据集转化为Iterator对象","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 再次运行该单元格以获取下一组图像\nimages, labels = next(batch) # 获取验证集的下一个批量\nprobabilities = model.predict(images) # 图片在104个类别上的概率\npredictions = np.argmax(probabilities, axis=-1) # 其中最大的概率表示这个图片的预测类别\ndisplay_batch_of_images((images, labels), predictions) # 展示一个批量的图片，图片标题为预测标签+预测标签是否正确（OK或NO）\n# 举个例子：标题为wild rose（NO->watercress），这个图片实际是豆瓣花，但是预测为野玫瑰，所以它是错的。所以它的标签为 野玫瑰（NO->豆瓣花）","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 计算混淆矩阵、f1_score分数、精确率、召回率，并输出展示\n# 注意此评估是我们基于验证集的评估，提交测试结果，平台给出的f1_score分数是以及测试集的，因此与我们计算的不相等，往往略低。但f1_score的计算规则固定。\n\n# 参数为实际标签和预测的标签\ncmat = confusion_matrix(cm_correct_labels, cm_predictions, labels=range(104))\n# 计算f1分数\nscore = f1_score(cm_correct_labels, cm_predictions, labels=range(104), average='macro')\n# 计算精确率\nprecision = precision_score(cm_correct_labels, cm_predictions, labels=range(104), average='macro')\n# 计算召回率\nrecall = recall_score(cm_correct_labels, cm_predictions, labels=range(104), average='macro')\n# 归一化\ncmat = (cmat.T / cmat.sum(axis=1)).T # normalized\n# 绘制混淆矩阵\ndisplay_confusion_matrix(cmat, score, precision, recall)\n# 输出f1分数、精确率、召回率\nprint('f1 score: {:.3f}, precision: {:.3f}, recall: {:.3f}'.format(score, precision, recall))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setp8：对测试集预测，并生成提交结果","metadata":{}},{"cell_type":"code","source":"# 预测并输出结果\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) # 输出预测结果\n\n# 预测结果填入提交文件\nprint('Generating submission.csv file...')\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') # all in one batch\nnp.savetxt('submission.csv', np.rec.fromarrays([test_ids, predictions]), fmt=['%s', '%d'], delimiter=',', header='id,label', comments='')\n","metadata":{"execution":{"iopub.status.busy":"2024-10-23T15:07:57.016490Z","iopub.execute_input":"2024-10-23T15:07:57.016809Z","iopub.status.idle":"2024-10-23T15:08:12.783946Z","shell.execute_reply.started":"2024-10-23T15:07:57.016779Z","shell.execute_reply":"2024-10-23T15:08:12.782528Z"},"trusted":true},"execution_count":null,"outputs":[]}]}