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

如何在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()
        },
    }
}

关键修改说明:

  1. 修正表达式构建逻辑:放弃apply_many,改用Expr::Function直接构建双参数表达式,与pearson函数的实现逻辑对齐
  2. 包装自定义UDF:使用Function::new将synthetic_data_column封装为Polars可识别的用户自定义函数(UDF)
  3. 修复硬编码问题:从输入数组长度动态获取行数,避免固定值导致的不匹配
  4. 添加错误处理:替换unwrap()为match和map_err,处理空值、负方差等异常情况,避免panic
  5. 明确输出类型:在Function::new中指定输出为List(Float64),确保Polars能正确推断类型

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 16:34:53