更新时间:2026-08-24 GMT+08:00
分享

UDAF

UDAF是输入多行返回一组聚合值的函数,适用于典型的多模态数据处理场景,包括视频/音频/图像的统计特征聚合,多模态标签总结等。UDAF具体的定义和使用约束请参考UDAF

UDAF注册

UDAF注册传入的必须是Python类,对于__init__、aggregate_state、accumulate、merge、finish、 __del__ 6个实例方法有特殊的含义认定,详情请参见下表;其他的Python类和实例方法不做限制,用户可以任意添加。

表1 注册UDAF时的特殊实例方法

Python Class方法

是否必须

参数

含义

适用场景

__init__(self, *args, **kwargs)

通过with_arguments方法传入参数,只允许传入标量Scalar。

UDAF的构造方法。

初始化UDAF属性(如保存参数、打开文件、建立连接等)。

aggregate_state(self)

无。

可序列化的聚合中间态,用于跨分区传递。

accumulate与merge间共享状态。

accumulate(self, *args, **kwargs)

通过UDAF算子传入参数,可以传入标量Scalar和列名Column。

增量更新聚合中间态。

正常数据扫过阶段。

merge(self, other_state)

来自其它分区的中间态。

合并其它分区状态到当前中间态。

分布式reduce阶段,多分区聚合结果归并。

finish(self)

无。

输出最终聚合结果。

聚合收尾。

__del__(self)

不支持传入参数。

UDAF的析构方法。

UDAF资源清理(如关闭文件、断开网络连接等)。

示例

import os
import json
import datetime as pydt
import ibis.expr.datatypes as dt
import pandas as pd
from typing import Any, Dict, Optional, Tuple
import aura_frame as aura
from aura_frame.multimodal import ai_lake
from aura_frame.multimodal.function import AggregateFnBuilder

target_database = "test"

class _AggState:
    def __init__(self):
        self.sum_px_qty = 0.0
        self.sum_qty = 0.0
        self.rows_seen = 0
        self.rows_used = 0

class PythonVWAP:
    # ... (same as before, but use pydt.date instead of dt.date for type hints)

con = ai_lake.connect(
    aura_endpoint=os.getenv("aura_endpoint"),
    aura_endpoint_name=os.getenv("aura_endpoint_name"),
    aura_workspace_id=os.getenv("aura_workspace_id"),
    lf_catalog_name=os.getenv("lf_catalog_name"),
    access_key=os.getenv("access_key"),
    secret_key=os.getenv("secret_key"),
    default_database=target_database,
    use_single_cn_mode=True,
)

try:
    con.create_agg_function(
        PythonVWAP,
        database=target_database,
        signature=aura.Signature(
            parameters=[
                aura.Parameter(name="quantity", annotation=int),
                aura.Parameter(name="price", annotation=float),
                aura.Parameter(name="date", annotation=dt.date),
                aura.Parameter(name="symbol", annotation=str),
            ],
            return_annotation=str,
        ),
        volatility=aura.udf.VolatilityType.IMMUTABLE,
        strict=False,
    )
    ds = con.load_dataset("stock_table", database=target_database)
    vwap_handler = con.get_function("PythonVWAP", database=target_database)
    udaf_builder = AggregateFnBuilder(
        fn=vwap_handler,
        on=[ds.quantity, ds.price, ds.date, ds.symbol],
        as_col="VWAP_column",
        num_cpus=0.5,
    )
    ds = ds.aggregate([udaf_builder], by=[ds.symbol])
    res = ds.execute()
    df = res.copy()
    parsed = df["VWAP_column"].apply(json.loads).apply(pd.Series)
    df = pd.concat([df.drop(columns=["VWAP_column"]), parsed], axis=1)
    print(df)
finally:
    con.close()

相关文档