PySpark中并行化循环调用API 优化35万行数据处理性能
在PySpark中并行调用API优化35万行数据处理性能
原代码的核心问题是通过collect()将全量数据拉取到Driver节点串行循环调用API,完全没有利用Spark的分布式计算能力,导致35万行数据处理效率极低。下面是两种基于PySpark的并行化实现方案:
方案一:普通自定义UDF(逐行处理)
通过注册UDF,让每个Executor节点的Task并行处理部分数据,避免全量数据拉取到Driver:
from pyspark.sql import functions as F from pyspark.sql.types import StructType, StructField, StringType, IntegerType import requests # 定义UDF返回的Schema:包含id和错误响应码 error_schema = StructType([ StructField("id", StringType(), nullable=False), StructField("responseCode", IntegerType(), nullable=True) ]) def check_api_response(url, row_id): try: # 设置超时时间,避免请求卡住 response = requests.get(url, timeout=10) # 主动抛出HTTP错误状态码对应的异常 response.raise_for_status() resp_json = response.json() # 只记录响应码不等于1的情况 if resp_json.get("responseCode") != 1: return (row_id, resp_json.get("responseCode")) else: return (row_id, None) except Exception: # 捕获所有异常,用特殊码-999标记异常情况 return (row_id, -999) # 注册UDF api_check_udf = F.udf(check_api_response, error_schema) # 分布式处理数据,过滤出有错误的记录 error_df = df_url.select( api_check_udf(F.col("url"), F.col("id")).alias("error_info") ).select("error_info.*").filter(F.col("responseCode").isNotNull()) # 将错误结果收集到Driver并转为字典 dict_error = {row.id: row.responseCode for row in error_df.collect()} print(dict_error)
方案二:Pandas批量UDF(更高效)
逐行UDF存在较多序列化开销,使用Pandas批量UDF可以一次性处理多行数据,大幅提升效率:
from pyspark.sql import functions as F from pyspark.sql.types import StructType, StructField, StringType, IntegerType import requests import pandas as pd from requests.adapters import HTTPAdapter from urllib3.util.retry import Retry # 创建带重试机制的Session,提升API调用稳定性 def create_retry_session(): session = requests.Session() # 设置重试策略:针对常见的服务端错误和限流码重试 retry_strategy = Retry( total=3, backoff_factor=1, status_forcelist=[429, 500, 502, 503, 504] ) adapter = HTTPAdapter(max_retries=retry_strategy) session.mount("http://", adapter) session.mount("https://", adapter) return session # 批量处理函数:接收Pandas Series,返回DataFrame def batch_process_api(urls: pd.Series, row_ids: pd.Series) -> pd.DataFrame: session = create_retry_session() results = [] for url, row_id in zip(urls, row_ids): resp_code = None try: response = session.get(url, timeout=10) response.raise_for_status() resp_json = response.json() if resp_json.get("responseCode") != 1: resp_code = resp_json.get("responseCode") except Exception: resp_code = -999 results.append({"id": row_id, "responseCode": resp_code}) return pd.DataFrame(results) # 定义返回Schema error_schema = StructType([ StructField("id", StringType(), nullable=False), StructField("responseCode", IntegerType(), nullable=True) ]) # 注册Pandas UDF batch_api_udf = F.pandas_udf(batch_process_api, error_schema) # 分布式处理并过滤错误记录 error_df = df_url.select( batch_api_udf(F.col("url"), F.col("id")).alias("error_info") ).select("error_info.*").filter(F.col("responseCode").isNotNull()) # 转为错误字典 dict_error = {row.id: row.responseCode for row in error_df.collect()} print(dict_error)
关键优化注意事项
- 避免全量拉取数据:绝对不要用
collect()把所有数据拉到Driver,始终让Executor分布式处理。 - 控制API并发:如果目标API有QPS限制,可在UDF中加入
time.sleep(0.1)之类的延迟,或使用线程池(注意每个Task的线程数不要超过Executor资源限制)。 - 调整Spark并行度:通过
spark.sql.shuffle.partitions设置合适的并行度(比如35万行可设为200-500,匹配集群CPU核心数),让数据均匀分布到更多Task中并行处理。 - 异常处理:必须捕获API调用的所有异常(超时、连接失败、JSON解析错误等),避免单个请求失败导致整个Task崩溃。
- 复用Session:使用带重试的Session复用TCP连接,减少建立连接的开销,提升调用效率。
内容的提问来源于stack exchange,提问作者learner
相关产品推荐
相关产品推荐

