TensorFlow:如何从CSV创建自定义wine_quality数据集替代tensorflow_datasets内置版本
基于CSV创建兼容TFDS内置wine_quality格式的自定义数据集方法
以下两种实现方案均可完全对齐内置数据集的格式,可直接复用原有配套代码:
方法一:轻量适配(无需注册TFDS数据集,修改量最小)
仅需新增一个自定义加载函数,替换原有tfds.load调用即可,操作步骤如下:
- 确认CSV文件字段与内置wine_quality对齐,列顺序参考:fixed acidity、volatile acidity、citric acid、residual sugar、chlorides、free sulfur dioxide、total sulfur dioxide、density、pH、sulphates、alcohol、quality(最后一列为标签列)
- 自定义加载函数代码:
import tensorflow as tf import tensorflow_datasets as tfds def load_custom_wine_quality(csv_path, as_supervised=True, split="train"): # 对齐内置数据集的字段定义与数据类型 CSV_COLUMNS = [ "fixed acidity", "volatile acidity", "citric acid", "residual sugar", "chlorides", "free sulfur dioxide", "total sulfur dioxide", "density", "pH", "sulphates", "alcohol", "quality" ] DEFAULTS = [ tf.constant(0, dtype=tf.float32), tf.constant(0, dtype=tf.float32), tf.constant(0, dtype=tf.float32), tf.constant(0, dtype=tf.float32), tf.constant(0, dtype=tf.float32), tf.constant(0, dtype=tf.float32), tf.constant(0, dtype=tf.float32), tf.constant(0, dtype=tf.float64), tf.constant(0, dtype=tf.float32), tf.constant(0, dtype=tf.float64), tf.constant(0, dtype=tf.float64), tf.constant(0, dtype=tf.int32) ] def parse_csv_row(row): columns = tf.io.decode_csv(row, record_defaults=DEFAULTS) features = dict(zip(CSV_COLUMNS[:-1], columns[:-1])) label = columns[-1] return (features, label) if as_supervised else features # 加载CSV并解析,skip_header_lines根据你的CSV是否有表头调整 dataset = tf.data.TextLineDataset(csv_path, skip_header_lines=1) dataset = dataset.map(parse_csv_row, num_parallel_calls=tf.data.AUTOTUNE) # 数据集切分,可根据需求调整比例 total_size = len(dataset) train_size = int(0.8 * total_size) if split == "train": dataset = dataset.take(train_size) elif split == "test": dataset = dataset.skip(train_size) # 对齐内置数据集的Prefetch格式 return dataset.prefetch(tf.data.AUTOTUNE)
- 替换原有加载代码:
# 原有代码 # dataset = tfds.load(name="wine_quality", as_supervised=True, split="train") # 替换为如下代码,填入你的CSV文件路径 dataset = load_custom_wine_quality("./your_wine_data.csv", as_supervised=True, split="train")
方法二:注册自定义TFDS数据集(完全兼容原有tfds.load调用)
如果需要完全复用原有tfds.load调用方式,可将自定义数据集注册到本地TFDS目录中,操作步骤如下:
- 新建数据集文件夹
my_wine_quality,在文件夹内创建两个文件:__init__.py、my_wine_quality.py - 在
my_wine_quality.py中写入数据集定义代码:
import tensorflow as tf import tensorflow_datasets as tfds import pandas as pd class MyWineQuality(tfds.core.GeneratorBasedBuilder): VERSION = tfds.core.Version("1.0.0") RELEASE_NOTES = { "1.0.0": "自定义wine_quality数据集", } def _info(self) -> tfds.core.DatasetInfo: return tfds.core.DatasetInfo( builder=self, features=tfds.features.FeaturesDict({ "fixed acidity": tf.float32, "volatile acidity": tf.float32, "citric acid": tf.float32, "residual sugar": tf.float32, "chlorides": tf.float32, "free sulfur dioxide": tf.float32, "total sulfur dioxide": tf.float32, "density": tf.float64, "pH": tf.float32, "sulphates": tf.float64, "alcohol": tf.float64, "quality": tf.int32, }), # 对齐as_supervised=True的返回格式 supervised_keys=(["fixed acidity", "volatile acidity", "citric acid", "residual sugar", "chlorides", "free sulfur dioxide", "total sulfur dioxide", "density", "pH", "sulphates", "alcohol"], "quality"), ) def _split_generators(self, dl_manager: tfds.download.DownloadManager): # 填入你的CSV文件绝对路径 csv_path = "/path/to/your_wine_data.csv" df = pd.read_csv(csv_path) # 切分训练测试集,可自定义比例 train_df = df.sample(frac=0.8, random_state=42) test_df = df.drop(train_df.index) return { "train": self._generate_examples(train_df), "test": self._generate_examples(test_df), } def _generate_examples(self, df): for idx, row in df.iterrows(): yield idx, { "fixed acidity": row["fixed acidity"], "volatile acidity": row["volatile acidity"], "citric acid": row["citric acid"], "residual sugar": row["residual sugar"], "chlorides": row["chlorides"], "free sulfur dioxide": row["free sulfur dioxide"], "total sulfur dioxide": row["total sulfur dioxide"], "density": row["density"], "pH": row["pH"], "sulphates": row["sulphates"], "alcohol": row["alcohol"], "quality": row["quality"], }
- 在
my_wine_quality文件夹下执行命令构建并注册数据集:
tfds build
- 注册完成后即可直接使用原有调用方式加载自定义数据集,仅需修改name参数即可:
dataset = tfds.load(name="my_wine_quality", as_supervised=True, split="train")
两种方法返回的数据集类型均为PrefetchDataset,特征结构、数据类型与官方内置wine_quality完全一致,原有配套代码无需调整即可直接运行
内容的提问来源于stack exchange,提问作者Zhen Hao
相关产品推荐
相关产品推荐

