在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
相关产品推荐
相关产品推荐

