如何在Rust Polars中创建自定义双参数函数表达式
问题:将自定义双参数数组处理函数封装为Polars表达式
我已经实现了一个函数,能接收两个Float64Type的ChunkedArray,生成ListType的ChunkedArray。现在需要将其封装为Polars的表达式函数,实现类似pearson相关性的双参数调用方式(即fn(mean_col, var_col)而非mean_col.fn(var_col)),输出为List类型的新列。以下是我的尝试代码,但在表达式求值时遇到瓶颈:
use polars::prelude::*; use rand::thread_rng; use rand_distr::{Normal, Distribution}; fn synthetic_data( mean_series:&ChunkedArray<Float64Type>, variance_series:&ChunkedArray<Float64Type>, ) -> PolarsResult<ChunkedArray<ListType>> { let mut rng = thread_rng(); let random_values: Vec<Vec<f64>> = mean_series.iter() .zip(variance_series.iter()) .map(|(mean, variance)| { let std_dev = variance.unwrap().sqrt(); let normal_dist = Normal::new(mean.unwrap(), std_dev).unwrap(); (0..39).map(|_| normal_dist.sample(&mut rng)).collect() }) .collect(); let mut list_chunk = ListPrimitiveChunkedBuilder::<Float64Type>::new( "intraday".into(), 5, //rows of data 39, DataType::Float64 ); for row in random_values { list_chunk.append_slice(&row); } Ok(list_chunk.finish()) } fn synthetic_data_column(s:&[Column]) -> PolarsResult<Column> { let _mean = &s[0]; let _varaince = &s[1]; let calc = synthetic_data(_mean.f64().unwrap(), _varaince.f64().unwrap()); Ok(calc?.into_column()) } fn synthetic_data_expr(mean_column: Expr, variance_column: Expr) -> Expr { mean_column.apply_many( synthetic_intraday_data_column(), &[variance_column], GetOutput::same_type(), ) }
我期望实现的效果参考Polars中pearson相关性函数的写法:
/// Compute the pearson correlation between two columns. pub fn pearson_corr(a: Expr, b: Expr) -> Expr { let input = vec![a, b]; let function = FunctionExpr::Correlation { method: CorrelationMethod::Pearson, }; Expr::Function { input, function, options: FunctionOptions { collect_groups: ApplyOptions::GroupWise, cast_options: Some(CastingRules::cast_to_supertypes()), flags: FunctionFlags::default() | FunctionFlags::RETURNS_SCALAR, ..Default::default() }, } }
修正方案
以下是修正后的完整代码,解决了表达式求值问题并优化了空值处理:
use polars::prelude::*; use rand::thread_rng; use rand_distr::{Normal, Distribution}; fn synthetic_data( mean_series: &ChunkedArray<Float64Type>, variance_series: &ChunkedArray<Float64Type>, ) -> PolarsResult<ChunkedArray<ListType>> { let mut rng = thread_rng(); let row_count = mean_series.len(); let list_len = 39; let mut list_chunk = ListPrimitiveChunkedBuilder::<Float64Type>::new( "intraday".into(), row_count, row_count * list_len, DataType::Float64, ); for (mean, variance) in mean_series.iter().zip(variance_series.iter()) { match (mean, variance) { (Some(m), Some(v)) if *v >= 0.0 => { let std_dev = v.sqrt(); let normal_dist = Normal::new(*m, std_dev).map_err(|e| PolarsError::ComputeError(e.into()))?; let values: Vec<f64> = (0..list_len).map(|_| normal_dist.sample(&mut rng)).collect(); list_chunk.append_slice(&values); } // 处理空值或负方差的情况,插入空List _ => list_chunk.append_null(), } } Ok(list_chunk.finish()) } fn synthetic_data_column(s: &[Column]) -> PolarsResult<Column> { let mean_col = s[0].f64().ok_or_else(|| { PolarsError::SchemaMismatch("expected first column to be Float64".into()) })?; let variance_col = s[1].f64().ok_or_else(|| { PolarsError::SchemaMismatch("expected second column to be Float64".into()) })?; synthetic_data(mean_col, variance_col).map(|ca| ca.into_column()) } fn synthetic_data_expr(mean_column: Expr, variance_column: Expr) -> Expr { let input = vec![mean_column, variance_column]; // 包装自定义列处理函数 let function = Function::new( "synthetic_data", synthetic_data_column, // 指定输出类型为List(Float64) GetOutput::from_type(DataType::List(Box::new(DataType::Float64))), ); Expr::Function { input, function: FunctionExpr::Udf(function), options: FunctionOptions { collect_groups: ApplyOptions::ApplyFlat, cast_options: Some(CastingRules::cast_to_supertypes()), flags: FunctionFlags::default(), ..Default::default() }, } }
关键修改说明:
- 修正表达式构建逻辑:放弃
apply_many,改用Expr::Function直接构建双参数表达式,与pearson函数的实现逻辑对齐 - 包装自定义UDF:使用
Function::new将synthetic_data_column封装为Polars可识别的用户自定义函数(UDF) - 修复硬编码问题:从输入数组长度动态获取行数,避免固定值导致的不匹配
- 添加错误处理:替换
unwrap()为match和map_err,处理空值、负方差等异常情况,避免panic - 明确输出类型:在
Function::new中指定输出为List(Float64),确保Polars能正确推断类型
内容的提问来源于stack exchange,提问作者Trevor Seibert
相关产品推荐
相关产品推荐

