You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Polars性能优化问题:逐行生成DataFrame遇性能瓶颈

优化大型Polars DataFrame逐行处理的性能方案

针对你在大型Polars DataFrame中逐行运行第三方算法的性能瓶颈,核心优化方向是用Polars矢量化操作替代逐行map_elements,并优化预处理和模型预测的批量处理逻辑,以下是具体方案:

一、核心优化思路

  1. 替换逐行生成小DataFrame的逻辑,用Polars原生的explode和广播操作批量生成候选数据集
  2. 移除不必要的Pandas转换,用Polars原生列表运算处理嵌入向量
  3. 对第三方ML模型采用批量预测,避免单条/小批量调用的开销
  4. 批量写入SQL,减少数据库IO次数

二、具体代码实现

1. 矢量化生成候选DataFrame(替代原build_candidate_dataframe)

原逐行生成小DataFrame的方式在大型数据集上效率极低,改用explode展开列表列,一次性生成所有候选行:

import polars as pl
import numpy as np

def build_candidate_dataframe_vectorized(df_items, df_products):
    # 展开相似度分数和位置列表,生成所有候选行
    exploded_df = df_items.explode(
        ["SEARCH_SIMILARITY_SCORE", "SEARCH_POSITION"]
    ).rename({
        "SEARCH_SIMILARITY_SCORE": "SIMILARITY_SCORE",
        "SEARCH_POSITION": "POSITION"
    })
    
    # 自动广播查询相关字段到所有展开后的行
    candidate_df = exploded_df.with_columns([
        pl.col("PRODUCT_INFO").alias("QUERY"),
        pl.col("SKU").alias("QUERY_SKU"),
        pl.col("YELLOW_CAT").alias("QUERY_LEAF"),
        pl.col("CATL3").alias("QUERY_CAT"),
        pl.col("EMBEDDINGS").alias("QUERY_EMBEDDINGS")
    ])
    
    # 与产品表关联,获取相似产品信息
    candidate_df = candidate_df.join(
        df_products.select([
            "SKU", "EMBEDDINGS", "INDEX", "DESCRIPTION", "CATL3", "YELLOW_CAT"
        ]),
        left_on="POSITION",
        right_on="INDEX",
        how="left"
    ).rename({
        "DESCRIPTION": "SIMILAR_PRODUCT_INFO",
        "CATL3": "SIMILAR_PRODUCT_CAT",
        "YELLOW_CAT": "SIMILAR_PRODUCT_LEAF"
    })
    
    # 按相似度降序排序(按需保留)
    return candidate_df.sort("SIMILARITY_SCORE", descending=True)

2. 优化预处理与批量模型预测(替代原process_candidate_output)

移除Pandas转换,用Polars原生列表运算处理嵌入向量,同时将模型预测改为批量操作:

def process_candidate_output_vectorized(candidate_df, ml_model, ml_comp_model):
    # Polars原生计算合并嵌入向量,无需转Pandas
    candidate_df = candidate_df.with_columns(
        ((pl.col("QUERY_EMBEDDINGS") + pl.col("EMBEDDINGS")) / 2).alias("COMBINED_EMBEDDINGS")
    )
    
    # 筛选需要的列
    target_cols = [
        "QUERY", "QUERY_SKU", "QUERY_CAT", "QUERY_LEAF",
        "SIMILAR_PRODUCT_INFO", "SIMILAR_PRODUCT_CAT", "SIMILAR_PRODUCT_LEAF",
        "SIMILARITY_SCORE", "COMBINED_EMBEDDINGS", "SKU", "POSITION"
    ]
    candidate_df = candidate_df.select(target_cols)
    
    # 过滤掉查询产品自身(矢量化操作,无需取单行值)
    candidate_df = candidate_df.filter(pl.col("SKU") != pl.col("QUERY_SKU"))
    
    # 批量执行补全预测
    if not candidate_df.is_empty():
        # 提取嵌入向量为Numpy矩阵,供模型批量预测
        embeddings_mat = np.array(candidate_df["COMBINED_EMBEDDINGS"].to_list())
        comp_preds = ml_model.predict(embeddings_mat)
        comp_probs = ml_model.predict_proba(embeddings_mat)[:, 1]
        
        candidate_df = candidate_df.with_columns([
            pl.Series("COMPLEMENTARY_PREDICTIONS", comp_preds),
            pl.Series("COMPLEMENT_PROBABILITY", comp_probs)
        ])
        
        # 过滤预测为1的结果
        candidate_df = candidate_df.filter(pl.col("COMPLEMENTARY_PREDICTIONS") == 1)
        
        # 批量执行配件预测
        if not candidate_df.is_empty():
            acc_embeddings_mat = np.array(candidate_df["COMBINED_EMBEDDINGS"].to_list())
            acc_preds = ml_comp_model.predict(acc_embeddings_mat)
            acc_probs = ml_comp_model.predict_proba(acc_embeddings_mat)[:, 1]
            
            candidate_df = candidate_df.with_columns([
                pl.Series("ACCESSORY_PREDICTIONS", acc_preds),
                pl.Series("LABEL_PROBABILITY", acc_probs)
            ])
    
    # 按预测概率降序排序
    return candidate_df.sort("LABEL_PROBABILITY", descending=True)

3. 批量处理主流程(替代原逐行map_elements)

将所有步骤整合为批量处理,避免逐行循环:

def batch_process(df_items, df_products, ml_model, ml_comp_model, db_engine, current_datetime):
    try:
        # 1. 批量生成候选数据集
        candidate_df = build_candidate_dataframe_vectorized(df_items, df_products)
        
        # 2. 批量预处理与模型预测
        processed_df = process_candidate_output_vectorized(candidate_df, ml_model, ml_comp_model)
        
        # 3. 批量写入SQL(需修改原write_validate_complements为批量写入逻辑)
        write_validate_complements_batch(processed_df, df_items, current_datetime, db_engine)
        
    except Exception as e:
        print(f"批量处理异常: {repr(e)}")

# 调用示例
# batch_process(df_items_sm_ex, df_products, rfc, rfc_comp, engine, current_datetime)

三、关键优化点说明

  • 矢量化操作替代逐行循环:Polars的explode和广播操作是底层优化的矢量化逻辑,比map_elements逐行处理快10-100倍(取决于数据集大小)
  • 避免数据格式转换:移除原代码中to_pandas()的转换,直接用Polars原生列表运算,减少内存开销和转换时间
  • 批量模型预测:第三方ML模型(如Scikit-learn)的批量预测效率远高于单条调用,避免了多次模型初始化和数据传递的开销
  • 批量写入SQL:将所有处理结果一次性写入数据库,减少连接建立和IO次数,提升写入速度

如果数据集过大导致内存不足,可以进一步拆分数据块(用Polars的partition_by或分块读取),分批次处理后再合并写入。

内容的提问来源于stack exchange,提问作者tastycakezs

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.22 00:14:51