加载含Base64存储Numpy数组的JSON到PySpark遇编码错误,如何解决?
如何将含Base64编码Numpy数组的JSON文件加载到PySpark?
我有一个JSON输入文件,每行包含一个ID和对应的以Base64格式存储的Numpy数组,请问如何将该文件加载到PySpark中?
我尝试创建以下UDF来实现:
from pyspark.sql.functions import udf from pyspark.sql.types import ArrayType, DoubleType import base64 def decode_base64_and_convert_to_numpy(base64_string): decoded_bytes = base64.b64decode(base64_string) decoded_str = decoded_bytes.decode('utf-8') decoded_list = json.loads(decoded_str) return np.array(decoded_list) decode_udf = udf(decode_base64_and_convert_to_numpy, ArrayType(DoubleType()))
但调用该UDF时出现了编码错误:
numpy_loaded_embeddings = raw_input.withColumn('numpy_embedding', decode_udf('model_output'))
报错信息如下:
An error was encountered: An exception was thrown from the Python worker. Please see the stack trace below. Traceback (most recent call last): File "<stdin>", line 7, in decode_base64_and_convert_to_numpy UnicodeDecodeError: 'utf-8' codec can't decode byte 0x93 in position 0: invalid start byte Traceback (most recent call last): File "/mnt/yarn/usercache/livy/appcache/application_1705979808797_0003/container_1705979808797_0003_01_000001/pyspark.zip/pyspark/sql/dataframe.py", line 607, in show print(self._jdf.showString(n, 20, vertical)) File "/mnt/yarn/usercache/livy/appcache/application_1705979808797_0003/container_1705979808797_0003_01_000001/py4j-0.10.9.5-src.zip/py4j/java_gateway.py", line 1322, in __call__ answer, self.gateway_client, self.target_id, self.name) File "/mnt/yarn/usercache/livy/appcache/application_1705979808797_0003/container_1705979808797_0003_01_000001/pyspark.zip/pyspark/sql/utils.py", line 196, in deco raise converted from None pyspark.sql.utils.PythonException: An exception was thrown from the Python worker. Please see the stack trace below. Traceback (most recent call last): File "<stdin>", line 7, in decode_base64_and_convert_to_numpy UnicodeDecodeError: 'utf-8' codec can't decode byte 0x93 in position 0: invalid start byte
问题原因
错误源于对Numpy数组Base64编码逻辑的误解:Base64编码的是Numpy数组的二进制序列化数据(比如用np.save生成的字节流),而非JSON字符串的Base64编码。直接将解码后的二进制字节尝试按UTF-8解码成字符串,必然会出现编码错误,因为二进制数据并非合法的UTF-8字符。
解决方案
修正UDF逻辑,直接用numpy.frombuffer解析解码后的二进制字节,转换为Numpy数组后转为Python列表(适配Spark的ArrayType类型):
from pyspark.sql.functions import udf from pyspark.sql.types import ArrayType, DoubleType import base64 import numpy as np def decode_base64_to_numpy_array(base64_string): # 解码Base64字符串为二进制字节 decoded_bytes = base64.b64decode(base64_string) # 从二进制字节加载Numpy数组(根据实际数组类型调整dtype) np_array = np.frombuffer(decoded_bytes, dtype=np.float64) # 转为Python列表返回,适配Spark ArrayType类型 return np_array.tolist() # 注册UDF,指定返回类型为Double类型的数组 decode_udf = udf(decode_base64_to_numpy_array, ArrayType(DoubleType()))
完整执行步骤
- 加载JSON文件
raw_input = spark.read.json("path/to/your/json/file")
- 应用UDF处理列
numpy_loaded_embeddings = raw_input.withColumn('numpy_embedding', decode_udf('model_output'))
- 验证结果
numpy_loaded_embeddings.select("id", "numpy_embedding").show(truncate=False)
注意事项
- 确保
dtype=np.float64与你原始Numpy数组的数据类型一致,如果是其他类型(如float32)需对应调整,否则会出现数据解析错误。 - 如果原始Numpy数组是多维的,
np.frombuffer会返回一维数组,此时需要额外调用reshape恢复原形状,再转为列表返回。
内容的提问来源于stack exchange,提问作者219CID
相关产品推荐
相关产品推荐

