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

用Polars原生API替代低效map_elements实现疾病编码映射

用Polars原生API替代低效map_elements实现疾病编码逻辑

问题背景

我有一个从CSV读取的Polars DataFrame,包含age、diagnosis字段,需要新增code列,值由diagnosis和age共同决定。目前通过map_elements实现了逻辑,但Polars官方警告该方法远慢于原生表达式API,仅应作为最后手段。尝试用map_batches替代时出现报错,求高效解决方案。

当前可用但低效的map_elements实现

import polars as pl

disease_codes = {
"malaria": {"under_5": "a", "over_5": "b"},
"PUD": "c",
"Asthma": "d",
}

def return_code(row):
    diagnosis = row["diagnosis"]
    age = row["age"]
    dx_return = disease_codes.get(diagnosis, "Undefined")
    if type(dx_return) == dict:
        if age >= 5:
            return dx_return.get("over_5")
        return dx_return.get("under_5")
    return dx_return

df.with_columns(pl.struct(["diagnosis", "age"]).map_elements(return_code).alias("code"))

报错的map_batches尝试

def return_code(diagnosis, age):
    dx_return = disease_codes.get(diagnosis, "Undefined")
    if type(dx_return) == dict:
        if age >= 5:
            return dx_return.get("over_5")
        return dx_return.get("under_5")
    return dx_return

df.with_columns(
    (pl.struct(["diagnosis", "age"]).map_batches(
        lambda x: return_code(x.struct.field("diagnosis"), x.struct.field("age"))
    )).alias("code")
)

报错信息:

TypeError: cannot use `__getitem__` on Series of dtype Struct([Field('diagnosis', Utf8), Field('age', Int64)]) with argument 'diagnosis' of type 'str'

解决方案

方法一:纯Polars原生表达式(性能最优,推荐)

完全用Polars原生条件表达式实现逻辑,避免任何Python层面的循环,是效率最高的方案。

直接写条件链(适合规则较少的场景)

df = df.with_columns(
    pl.when(pl.col("diagnosis") == "PUD")
      .then("c")
      .when(pl.col("diagnosis") == "Asthma")
      .then("d")
      .when(pl.col("diagnosis") == "malaria")
      .then(pl.when(pl.col("age") >= 5).then("b").otherwise("a"))
      .otherwise("Undefined")
      .alias("code")
)

动态构建表达式(适合规则较多的场景)

如果疾病编码规则较多,可以拆分规则并动态生成表达式,便于维护:

# 拆分固定编码和需按年龄分支的编码
fixed_codes = {k: v for k, v in disease_codes.items() if not isinstance(v, dict)}
age_based_codes = {k: v for k, v in disease_codes.items() if isinstance(v, dict)}

# 初始化默认值为"Undefined"的表达式
expr = pl.lit("Undefined")

# 添加固定编码规则
for diag, code in fixed_codes.items():
    expr = pl.when(pl.col("diagnosis") == diag).then(code).otherwise(expr)

# 添加按年龄分支的编码规则
for diag, age_map in age_based_codes.items():
    expr = pl.when(pl.col("diagnosis") == diag)
             .then(pl.when(pl.col("age") >= 5).then(age_map["over_5"]).otherwise(age_map["under_5"]))
             .otherwise(expr)

df = df.with_columns(expr.alias("code"))

方法二:修正后的map_batches用法

你之前的map_batches报错,是因为函数接收的是Series对象而非单行数据,必须做批量向量运算,不能直接用字典get处理Series。修正后的实现如下:

import numpy as np

def return_code_batch(struct_series):
    # 提取diagnosis和age的Series
    diagnosis = struct_series.struct.field("diagnosis")
    age = struct_series.struct.field("age")
    
    # 初始化结果数组,默认值为"Undefined"
    result = np.full(len(diagnosis), "Undefined", dtype=str)
    
    # 处理固定编码的疾病
    for diag, code in fixed_codes.items():
        mask = diagnosis == diag
        result[mask] = code
    
    # 处理疟疾的年龄分支
    malaria_mask = diagnosis == "malaria"
    under5_mask = malaria_mask & (age < 5)
    over5_mask = malaria_mask & (age >= 5)
    result[under5_mask] = disease_codes["malaria"]["under_5"]
    result[over5_mask] = disease_codes["malaria"]["over_5"]
    
    return pl.Series(result)

df = df.with_columns(
    pl.struct(["diagnosis", "age"]).map_batches(return_code_batch).alias("code")
)

该方法比map_elements高效,但仍不如纯原生表达式。

报错原因说明

你之前的return_code函数试图直接对Series调用disease_codes.get(diagnosis, ...),但diagnosis是Series类型(不是单个字符串),字典get方法无法处理Series,因此抛出类型错误。map_batches要求函数接收Series并返回Series,必须基于向量运算做批量处理,不能逐行操作。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 00:34:57