PySpark多标签文本分类报错:ValueError整数转换无效
问题分析与快速修复方案
嘿,这个错误我之前踩过坑!核心问题其实很直白:你的label列看起来是数组格式,但实际存储的是字符串类型,而不是PySpark能识别的ArrayType(IntegerType())。
为啥会报错?
你提到已经用多标签二值化把label转成了数组形式,但大概率是后续的存储/读取环节出了问题——比如把数据存成CSV时,数组会被自动序列化成类似"[1, 0, 1, 0, 1]"的字符串;或者二值化后的输出列没被正确保留为数组类型,导致模型训练时,OneVsRest+LinearSVC需要数值型数组输入,却收到了字符串,尝试把整个字符串转成int就直接炸了,也就是你看到的ValueError: invalid literal for int() with base 10: '[1, 0, 1, 0, 1, 1, 1, 0, 0]'。
怎么解决?
分两步走,先确认问题,再修复:
先验证label列的真实类型
运行下面的代码看看列类型:df.printSchema()如果输出里
label的类型是string,那实锤就是这个问题了。把字符串格式的"数组"转成真正的数值数组
这里给你两种靠谱的方法:- 方法一:用
from_json(推荐,更稳妥)
利用PySpark的JSON解析函数,把字符串转成指定类型的数组:from pyspark.sql.types import ArrayType, IntegerType from pyspark.sql.functions import from_json # 定义目标数组类型的schema label_schema = ArrayType(IntegerType()) # 转换label列 df = df.withColumn("label", from_json(df["label"], label_schema)) - 方法二:正则分割+类型转换(适合格式特别规整的情况)
如果你的字符串数组格式很统一(比如首尾是[],元素用逗号分隔),可以用正则去掉括号再分割转类型:from pyspark.sql.functions import split, regexp_replace, col df = df.withColumn("label", split(regexp_replace(col("label"), r"^\[|\]$", ""), ",") .cast(ArrayType(IntegerType())) )
- 方法一:用
再次验证类型
转换后再跑一遍df.printSchema(),确认label的类型是array<int>,之后再执行你的Pipeline训练就没问题了。
额外提醒
以后处理多标签数据时,尽量用Parquet格式存储,它能保留PySpark的复杂类型(比如数组、结构体),不会像CSV那样把序列化成字符串,能避免很多这类坑。
内容的提问来源于stack exchange,提问作者Kertis van Kertis
相关产品推荐
相关产品推荐

