Python Polars处理超大GIS数据集列函数应用内存不足问题求助
处理超内存GIS数据集:解析字符串坐标并计算平均经纬度
针对2500万行的超内存数据集,以下是Polars及替代库的高效解决方案:
Polars 最优方案
Polars的矢量化操作能避免Python UDF的内存开销,结合流式处理可解决OOM问题:
核心思路
- 用
pl.scan_csv()懒加载CSV,避免全量数据入内存 - 用内置的
str.json_decode()解析字符串化的坐标数组(替代Python UDF的json.loads) - 用Polars数组函数计算平均经纬度,全程矢量化
- 流式写入Parquet,避免内存溢出
代码示例
import polars as pl # 懒加载CSV文件(超内存场景必须用scan而非read_csv) lazy_df = pl.scan_csv("your_input.csv") # 解析坐标并计算平均经纬度 processed_lazy_df = lazy_df.with_columns( # 将字符串解析为List[List[float]]类型的坐标数组 coords=pl.col("geometry.coordinates").str.json_decode(), # 提取所有点的经度,计算平均值 avg_lon=pl.col("geometry.coordinates") .str.json_decode() .arr.eval(pl.element()[0]) .arr.mean(), # 提取所有点的纬度,计算平均值 avg_lat=pl.col("geometry.coordinates") .str.json_decode() .arr.eval(pl.element()[1]) .arr.mean() ).drop("geometry.coordinates") # 可选:删除无需保留的原坐标列 # 流式写入Parquet,分块处理数据避免OOM processed_lazy_df.write_parquet("output_polars.parquet", streaming=True)
为什么之前的方法失败
- 直接
cast(pl.List)无效:Polars无法自动将JSON格式的字符串转为List类型,必须用str.json_decode()做专门解析 map_elements+json.loads导致OOM:Python UDF会触发全量数据加载,且每行调用Python解释器带来巨大内存开销sink_parquet报错:默认标准引擎不支持该操作,write_parquet(streaming=True)是Polars官方推荐的超内存数据写入方式
替代方案:Dask 分块处理
如果Polars仍有兼容性问题,可使用Dask的分块机制处理超内存数据:
import dask.dataframe as dd import json import pandas as pd def compute_avg_coords(coords_str): coords = json.loads(coords_str) lons = [point[0] for point in coords] lats = [point[1] for point in coords] return pd.Series([sum(lons)/len(lons), sum(lats)/len(lats)]) # 读取CSV为Dask DataFrame(自动分块) dask_df = dd.read_csv("your_input.csv") # 计算平均经纬度,指定输出列类型 dask_df[["avg_lon", "avg_lat"]] = dask_df["geometry.coordinates"].apply( compute_avg_coords, meta={"avg_lon": float, "avg_lat": float} ) # 写入Parquet(分块写入) dask_df.drop("geometry.coordinates").to_parquet("output_dask.parquet")
大规模集群方案:PySpark
若数据集规模远超单节点内存,可使用PySpark进行分布式处理:
from pyspark.sql import SparkSession from pyspark.sql.functions import udf from pyspark.sql.types import StructType, StructField, DoubleType import json # 初始化Spark会话 spark = SparkSession.builder.appName("GIS_Avg_Coords").getOrCreate() # 定义UDF:解析坐标并计算平均经纬度 @udf(returnType=StructType([ StructField("avg_lon", DoubleType()), StructField("avg_lat", DoubleType()) ])) def calculate_avg_coords(coords_str): coords = json.loads(coords_str) lons = [p[0] for p in coords] lats = [p[1] for p in coords] return (sum(lons)/len(lons), sum(lats)/len(lats)) # 读取CSV文件 spark_df = spark.read.csv("your_input.csv", header=True, inferSchema=True) # 计算并提取平均经纬度 spark_df = spark_df.withColumn("avg_coords", calculate_avg_coords(spark_df["geometry.coordinates"])) spark_df = spark_df.select("*", "avg_coords.avg_lon", "avg_coords.avg_lat").drop("avg_coords", "geometry.coordinates") # 写入Parquet spark_df.write.parquet("output_spark.parquet")
内容的提问来源于stack exchange,提问作者Lionel Peer
相关产品推荐
相关产品推荐

