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

TensorFlow:如何从CSV创建自定义wine_quality数据集替代tensorflow_datasets内置版本

基于CSV创建兼容TFDS内置wine_quality格式的自定义数据集方法

以下两种实现方案均可完全对齐内置数据集的格式,可直接复用原有配套代码:

方法一:轻量适配(无需注册TFDS数据集,修改量最小)

仅需新增一个自定义加载函数,替换原有tfds.load调用即可,操作步骤如下:

  1. 确认CSV文件字段与内置wine_quality对齐,列顺序参考:fixed acidity、volatile acidity、citric acid、residual sugar、chlorides、free sulfur dioxide、total sulfur dioxide、density、pH、sulphates、alcohol、quality(最后一列为标签列)
  2. 自定义加载函数代码:
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)
  1. 替换原有加载代码:
# 原有代码
# 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目录中,操作步骤如下:

  1. 新建数据集文件夹my_wine_quality,在文件夹内创建两个文件:__init__.py、my_wine_quality.py
  2. 在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"],
            }
  1. 在my_wine_quality文件夹下执行命令构建并注册数据集:
tfds build
  1. 注册完成后即可直接使用原有调用方式加载自定义数据集,仅需修改name参数即可:
dataset = tfds.load(name="my_wine_quality", as_supervised=True, split="train")

两种方法返回的数据集类型均为PrefetchDataset,特征结构、数据类型与官方内置wine_quality完全一致,原有配套代码无需调整即可直接运行

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 14:54:01