You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在Google Colab中调整WiFi/蓝牙频谱图分类模型的标签与输出文本

解决方案

一、修改预测输出文本

你可以通过创建类别映射字典,将数字索引转换为对应的中文类别名称,替代原有的数字输出。

方法1:手动映射(适合已知类别顺序)

直接替换原代码中的print部分:

# 替换原有的print语句
class_mapping = {0: "蓝牙", 1: "WiFi"}
predicted_class_name = class_mapping[predicted_class]
print('预测类别:', predicted_class_name)

方法2:自动获取类别映射(更可靠)

利用flow_from_directory生成的class_indices属性,自动获取文件夹到索引的映射,避免手动写错顺序:

# 在创建training_set后添加这行代码保存类别映射
class_indices = training_set.class_indices
# 反转字典得到索引到类别名称的映射
class_mapping = {v: k for k, v in class_indices.items()}

# 预测时替换print语句
predicted_class_name = class_mapping[predicted_class]
print('预测类别:', predicted_class_name)

二、利用LabelImg的TXT标注重新训练模型

LabelImg导出的TXT文件(默认YOLO格式)包含类别索引信息,你需要自定义数据生成器来读取图像和对应的标注文件,替代原来的flow_from_directory。

前提假设

  • 图像和标注文件同名(如img1.jpg对应img1.txt)
  • TXT文件第一行为类别索引(0对应蓝牙,1对应WiFi)
  • 图像和标注分别存放在images和labels文件夹中

自定义数据生成器代码

import os
import cv2
import numpy as np

# 配置路径
IMAGE_DIR = '/content/drive/MyDrive/My training dataset/images'
LABEL_DIR = '/content/drive/MyDrive/My training dataset/labels'
TARGET_SIZE = (224, 224)
BATCH_SIZE = 16
NUM_CLASSES = 2  # 类别总数

def custom_data_generator(img_dir, label_dir):
    # 获取所有图像路径
    img_paths = [os.path.join(img_dir, f) for f in os.listdir(img_dir) if f.lower().endswith(('.png', '.jpg', '.jpeg'))]
    while True:
        # 打乱数据顺序
        np.random.shuffle(img_paths)
        for batch_start in range(0, len(img_paths), BATCH_SIZE):
            batch_imgs = []
            batch_labels = []
            batch_paths = img_paths[batch_start:batch_start+BATCH_SIZE]
            for path in batch_paths:
                # 读取并预处理图像
                img = cv2.imread(path)
                img = cv2.resize(img, TARGET_SIZE)
                img = img / 255.0  # 归一化到0-1
                batch_imgs.append(img)
                
                # 读取对应标注文件
                label_filename = os.path.basename(path).replace('.jpg', '.txt').replace('.png', '.txt')
                label_path = os.path.join(label_dir, label_filename)
                with open(label_path, 'r') as f:
                    # 读取第一行的类别索引
                    class_idx = int(f.readline().split()[0])
                    # 转换为one-hot编码(适配categorical_crossentropy损失)
                    one_hot_label = np.zeros(NUM_CLASSES)
                    one_hot_label[class_idx] = 1
                    batch_labels.append(one_hot_label)
            
            yield np.array(batch_imgs), np.array(batch_labels)

# 初始化训练生成器(有验证集的话按同样逻辑创建)
train_generator = custom_data_generator(IMAGE_DIR, LABEL_DIR)
# 计算训练步数
train_steps = len(os.listdir(IMAGE_DIR)) // BATCH_SIZE

# 训练模型
r = model.fit(
    train_generator,
    epochs=10,
    steps_per_epoch=train_steps
)

注意事项

  • 如果TXT文件存储的是类别名称(如"蓝牙")而非索引,需先创建映射:class_name_to_idx = {"蓝牙":0, "WiFi":1},再将读取代码改为class_idx = class_name_to_idx[class_name]。
  • 若想简化标签处理,可将损失函数改为sparse_categorical_crossentropy,此时无需转换为one-hot编码,直接传入类别索引即可。

内容的提问来源于stack exchange,提问作者PSY

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 12:54:52