用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
相关产品推荐
相关产品推荐

