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

如何创建适配TextClassifier的自定义CSV数据集?

适配Keras 3 TextClassifier的自定义CSV数据集构建方法

你的核心问题在于之前的代码要么列名不匹配、未处理字符串标签到数值的转换,要么数据集格式不符合TextClassifier.fit()的输入要求(模型期望输入是(text, label)结构的数据集,其中text为字符串类型,label为整数类型)。以下是两种可行的解决方案:

方案一:基于tf.data的CSV读取方案

该方案直接用TensorFlow的CSV读取工具,修正列名匹配问题并添加标签转换:

import tensorflow as tf

def create_dataset_from_csv(filename, batch_size=16):
    # 定义正负标签到整数的映射
    label_map = {"pos": 1, "neg": 0}
    
    # 读取CSV,确保列名与你的文件一致(假设CSV表头为text,label)
    dataset = tf.data.experimental.make_csv_dataset(
        filename,
        batch_size=batch_size,
        label_name="label",  # 对应CSV中的label列名
        select_columns=["text", "label"],  # 选择需要的列
        field_delim=",",
        header=True,  # 如果CSV有表头则设为True,无表头则改为False
        num_epochs=None  # 交由fit()控制训练轮数
    )
    
    # 将批量数据转换为模型期望的(text, label)元组格式
    def preprocess_batch(features, label):
        text = features["text"]
        # 将字符串标签映射为整数
        numeric_label = tf.map_fn(lambda x: label_map[x], label, dtype=tf.int32)
        return text, numeric_label
    
    dataset = dataset.map(preprocess_batch)
    return dataset

方案二:基于Pandas的灵活处理方案

如果需要更灵活的文本预处理(比如清洗、过滤),用Pandas读取CSV后转换为tf.data.Dataset会更直观:

import pandas as pd
import tensorflow as tf

def create_dataset_from_csv_pandas(filename, batch_size=16):
    # 读取CSV文件
    df = pd.read_csv(filename)
    # 将字符串标签转换为模型需要的整数
    df["label"] = df["label"].map({"pos": 1, "neg": 0})
    
    # 转换为tf.data.Dataset并添加优化
    dataset = tf.data.Dataset.from_tensor_slices((df["text"].values, df["label"].values))
    # 打乱数据、分批、预取提升训练效率
    dataset = dataset.shuffle(buffer_size=len(df)).batch(batch_size).prefetch(tf.data.AUTOTUNE)
    return dataset

使用示例

# 加载自定义数据集
train_dataset = create_dataset_from_csv("train.csv")
test_dataset = create_dataset_from_csv("test.csv")

# 或者使用Pandas版本
# train_dataset = create_dataset_from_csv_pandas("train.csv")
# test_dataset = create_dataset_from_csv_pandas("test.csv")

# 初始化并训练模型
classifier = hub.models.TextClassifier.from_preset(
    "bert_base_multi",
    num_classes=2
)
classifier.fit(train_dataset, validation_data=test_dataset, epochs=3)

原方法失效原因说明

  • 第一个方法中列名使用了"Class"和"Text",与你的CSV实际列名不匹配,导致无法正确读取数据;
  • 第二个方法返回的是数组而非tf.data.Dataset,虽然fit()支持数组输入,但大数据场景下内存占用过高,且缺乏训练优化;
  • 第三个方法中create_tf_dataset_from_csv函数未定义,同时未处理字符串标签到整数的转换,导致数据集格式不符合要求。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 21:07:21