如何从PySpark DataFrame指定列提取数据生成一维字符串数组
问题原因
df.select('name').collect() 返回的是pyspark.sql.Row对象的列表,每个Row对应一条记录,包裹了你要的name字段值,所以直接转numpy数组会得到二维嵌套结构。
可行实现方案
方案1:直接提取字段生成字符串列表(推荐,无需依赖numpy)
直接在collect后通过列表推导式提取name字段值,得到的就是纯字符串组成的一维动态容器(Python列表本身就是动态扩容的容器):
# 直接得到字符串列表 name_str_list = [row["name"] for row in df.select("name").collect()] # 如果需要转numpy一维字符串数组,再加下面这行 name_np_array = np.array(name_str_list)
方案2:对已有的numpy二维数组做降维处理
如果要保留numpy处理逻辑,直接调用flatten()方法拉平二维数组即可:
data_array = np.asarray(df.select('name').collect()).flatten() # 此时data_array就是一维字符串数组
方案3:循环手动追加到动态容器
你提到的动态容器可以直接用Python原生列表实现,示例如下:
data_array = np.asarray(df.select('name').collect()) res_container = [] for x in data_array: s = x[0] res_container.append(s) # res_container 就是你需要的字符串列表
注意事项
如果你的DataFrame数据量极大,不要直接调用collect(),该操作会把全量数据拉取到Spark Driver节点的内存中,容易触发内存溢出,仅建议小数据量场景下使用上述方案。
内容的提问来源于stack exchange,提问作者DataSteve
相关产品推荐
相关产品推荐

