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

如何使用TensorFlow/Keras在Python中加载并分割GTZAN频谱图数据集

加载GTZAN频谱图数据集并分割训练/测试集(TensorFlow/Keras版)

嘿,刚接触机器学习和Keras的话,GTZAN这种按文件夹分类的数据集其实很好处理!我之前做音频分类的时候也用过这个数据集,给你分享两个最实用的方法,都是开箱即用的:

方法一:用tf.keras.utils.image_dataset_from_directory(推荐,TensorFlow 2.3+)

这是TensorFlow官方推荐的现代方法,自动帮你处理文件夹结构、标签分配,还能直接分割训练/测试集,省超多事:

步骤1:导入依赖

import tensorflow as tf
import numpy as np

步骤2:定义基础参数

# 你的数据集根路径(根据实际路径调整)
data_dir = "./data"
# 输入图片的尺寸(GTZAN频谱图常见尺寸是288x432,按你实际图调整)
img_height, img_width = 288, 432
# 批量大小,根据你的显存情况改
batch_size = 32
# 测试集占比,比如设为20%
test_split = 0.2

步骤3:加载并分割数据集

# 加载训练集
train_ds = tf.keras.utils.image_dataset_from_directory(
  data_dir,
  validation_split=test_split,
  subset="training",
  seed=123,  # 固定随机种子,保证每次分割结果一致
  image_size=(img_height, img_width),
  batch_size=batch_size)

# 加载测试集
test_ds = tf.keras.utils.image_dataset_from_directory(
  data_dir,
  validation_split=test_split,
  subset="validation",
  seed=123,
  image_size=(img_height, img_width),
  batch_size=batch_size)

步骤4:转换成你习惯的x_train, y_train格式

上面返回的是TensorFlow的Dataset对象,要是你需要numpy数组格式,可以转一下:

def dataset_to_numpy(dataset):
    x_list = []
    y_list = []
    for images, labels in dataset:
        x_list.append(images.numpy())
        y_list.append(labels.numpy())
    return np.concatenate(x_list), np.concatenate(y_list)

x_train, y_train = dataset_to_numpy(train_ds)
x_test, y_test = dataset_to_numpy(test_ds)

# 顺便看看类别名称(就是你文件夹的名字:pop、hip hop这些)
class_names = train_ds.class_names
print("所有流派类别:", class_names)

可选:优化训练速度

把数据缓存到内存,避免重复读磁盘,训练会快很多:

train_ds = train_ds.cache().prefetch(buffer_size=tf.data.AUTOTUNE)
test_ds = test_ds.cache().prefetch(buffer_size=tf.data.AUTOTUNE)

方法二:用ImageDataGenerator(兼容旧版Keras)

如果你的TensorFlow版本比较老,用这个经典方法也没问题:

步骤1:导入依赖

from tensorflow.keras.preprocessing.image import ImageDataGenerator
import numpy as np

步骤2:初始化数据生成器

# 归一化像素值到0-1,同时设置测试集分割比例
datagen = ImageDataGenerator(rescale=1./255, validation_split=test_split)

步骤3:加载训练/测试集

train_generator = datagen.flow_from_directory(
    data_dir,
    target_size=(img_height, img_width),
    batch_size=batch_size,
    class_mode='categorical',  # 多分类场景用这个,自动生成one-hot标签
    subset='training',
    seed=123)

test_generator = datagen.flow_from_directory(
    data_dir,
    target_size=(img_height, img_width),
    batch_size=batch_size,
    class_mode='categorical',
    subset='validation',
    seed=123)

转换成numpy数组格式

生成器是批量返回数据的,要拿到全部数据得循环读取:

# 获取训练集全部数据
x_train, y_train = train_generator.next()
for _ in range(train_generator.samples // batch_size):
    batch_x, batch_y = train_generator.next()
    x_train = np.concatenate([x_train, batch_x])
    y_train = np.concatenate([y_train, batch_y])

# 获取测试集全部数据
x_test, y_test = test_generator.next()
for _ in range(test_generator.samples // batch_size):
    batch_x, batch_y = test_generator.next()
    x_test = np.concatenate([x_test, batch_x])
    y_test = np.concatenate([y_test, batch_y])

几个小提醒

  • 一定要把img_height, img_width改成你实际频谱图的尺寸,别直接用我给的数值!
  • 如果需要把类别索引转成one-hot编码,方法一可以用tf.one_hot(y_train, depth=len(class_names)),方法二已经自动生成one-hot了
  • 固定seed=123是为了保证每次运行分割的训练/测试集都一样,方便复现实验结果

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 04:12:28