UDAF
UDAF是输入多行返回一组聚合值的函数,适用于典型的多模态数据处理场景,包括视频/音频/图像的统计特征聚合,多模态标签总结等。UDAF具体的定义和使用约束请参考UDAF。
UDAF注册
UDAF注册传入的必须是Python类,对于__init__、aggregate_state、accumulate、merge、finish、 __del__ 6个实例方法有特殊的含义认定,详情请参见下表;其他的Python类和实例方法不做限制,用户可以任意添加。
| 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()