跳到主要内容

自定义函数

自定义UDF函数

案例

自定义UDAF函数

参考

https://cloud.tencent.com/developer/article/1918556 https://cloud.tencent.com/developer/article/1846460

案例一 (多个参数)

public class MergeArrAccountStatTopN extends AbstractGenericUDAFResolver {
@Override
public GenericUDAFEvaluator getEvaluator(GenericUDAFParameterInfo info) throws SemanticException {
return getEvaluator(info.getParameters());
}

@Override
public GenericUDAFEvaluator getEvaluator(TypeInfo[] parameters) throws SemanticException {
if (parameters.length != 2) {
throw new UDFArgumentTypeException(parameters.length - 2,
"Exactly one argument is expected.");
}

if (parameters[0].getCategory() != ObjectInspector.Category.PRIMITIVE) {
throw new UDFArgumentTypeException(0,
"Only primitive type arguments are accepted but "
+ parameters[0].getTypeName() + " is passed.");
}
switch (((PrimitiveTypeInfo) parameters[0]).getPrimitiveCategory()) {
case STRING:
// 在这做类型检查以及选择指定的 Evaluator
return new MergeArrAccountStatTopNEvaluator();
default:
throw new UDFArgumentTypeException(0,
"Only string type arguments are accepted but " + parameters[0].getTypeName() + " is passed.");
}
}
}

总结结果存储

public class MergeArrAccountStatTopNAggregationBuffer extends GenericUDAFEvaluator.AbstractAggregationBuffer {
private Integer topKNum;
private HashMap<String, com.alibaba.fastjson.JSONObject> map = Maps.newHashMap();

public void addJsonArr(String jsonArr) {
synchronized (jsonArr) {
if (!Strings.isNullOrEmpty(jsonArr)) {
JSONArray jsonArray = JSON.parseArray(jsonArr);
for (Object o : jsonArray) {
com.alibaba.fastjson.JSONObject jsonObject = JSON.parseObject(o.toString());
String sourceId = jsonObject.getString("source_id");
if (map.containsKey(sourceId)) {
JSONObject existsObj = map.get(sourceId);
existsObj.put("works_num", existsObj.getIntValue("works_num") + jsonObject.getIntValue("works_num"));
existsObj.put("interactive", existsObj.getIntValue("interactive") + jsonObject.getIntValue("interactive"));
map.put(sourceId, existsObj);
} else {
map.put(sourceId, jsonObject);
}
}
}
}
}

public HashMap<String, JSONObject> getMap() {
return map;
}

public void setMap(HashMap<String, JSONObject> map) {
this.map = map;
}

public Integer getTopKNum() {
return topKNum;
}

public void setTopKNum(Integer topKNum) {
this.topKNum = topKNum;
}
}

核心逻辑

public class MergeArrAccountStatTopNEvaluator extends GenericUDAFEvaluator {
static final Log LOG = LogFactory.getLog(MergeArrAccountStatTopNEvaluator.class.getName());

// Iterate 输入
PrimitiveObjectInspector iterateInputJsonArrStringOIParamOne;
PrimitiveObjectInspector iterateInputTopNumOIParamTow;

// Merge 输入
private StructObjectInspector mergeInputStructOI;
private StringObjectInspector jsonArrStringFieldOI;
private IntObjectInspector topNumFieldOI;
private StructField jsonArrString;
private StructField topNum;
// TerminatePartial 输出
private Object[] partialResult;
// Terminate 输出
Text result;

/**
* 这个方法主要是指定的时候,每一个阶段的类型声明,包括输入和输出的类型
*
* @param mode The mode of aggregation.
* @param parameters
* @return
* @throws HiveException
*/
@Override
public ObjectInspector init(Mode mode, ObjectInspector[] parameters) throws HiveException {
super.init(mode, parameters);
// map输入数据类型说明
// 初始化输入参数(map阶段)这里先声明输入map端的数据类型
if (mode == Mode.PARTIAL1 || mode == Mode.COMPLETE) {
// 原始数据 (输入的第一个参数和第二个参数)
iterateInputJsonArrStringOIParamOne = (PrimitiveObjectInspector) parameters[0];
iterateInputTopNumOIParamTow = (PrimitiveObjectInspector) parameters[1];
} else {
// Merge 输入类型说明(struct的数据也就是第一个参数,是下面terminatePartial的输出)
// 部分聚合数据(这里是map到reduce的merger阶段的数据类型说明)
mergeInputStructOI = (StructObjectInspector) parameters[0];
jsonArrString = mergeInputStructOI.getStructFieldRef("json_arr_string");
topNum = mergeInputStructOI.getStructFieldRef("top_num");
jsonArrStringFieldOI = (StringObjectInspector) jsonArrString.getFieldObjectInspector();
topNumFieldOI = (IntObjectInspector) topNum.getFieldObjectInspector();
}
// TerminatePartial 输出的数据类型声明,这里要声明可以序列化的类型(由于这里声明的是struct类型的输出,那么上面的Merge就是struct的输入)
// 初始化输出
if (mode == Mode.PARTIAL1 || mode == Mode.PARTIAL2) {
// // 最终结果
result = new Text("");
partialResult = new Object[2];
partialResult[0] = result;
partialResult[1] = new IntWritable(0);
// 部分聚合结果
// 字段类型
ArrayList<ObjectInspector> structFieldOIs = new ArrayList<>();
structFieldOIs.add(PrimitiveObjectInspectorFactory.writableStringObjectInspector);
structFieldOIs.add(PrimitiveObjectInspectorFactory.writableIntObjectInspector);
// 字段名称
ArrayList<String> structFieldNames = new ArrayList<>();
structFieldNames.add("json_arr_string");
structFieldNames.add("top_num");
return ObjectInspectorFactory.getStandardStructObjectInspector(structFieldNames, structFieldOIs);
} else {
// Terminate 输出,由于结果是字符串,所以这里用到了text类型
result = new Text("");
return PrimitiveObjectInspectorFactory.writableStringObjectInspector;
}
}

public AggregationBuffer getNewAggregationBuffer() throws HiveException {
MergeArrAccountStatTopNAggregationBuffer mergeArrAccountStatTopNAggregationBuffer = new MergeArrAccountStatTopNAggregationBuffer();
reset(mergeArrAccountStatTopNAggregationBuffer);
return mergeArrAccountStatTopNAggregationBuffer;
}

@Override
public void reset(AggregationBuffer agg) throws HiveException {
((MergeArrAccountStatTopNAggregationBuffer) agg).setMap(Maps.newHashMap());
((MergeArrAccountStatTopNAggregationBuffer) agg).setTopKNum(0);
}

/**
* map 阶段的数据输入
*
* @param agg udaf进行group by以后的一条条数据
* @param parameters The objects of parameters.
* @throws HiveException
*/
public void iterate(AggregationBuffer agg, Object[] parameters) throws HiveException {
Object jsonArrp = parameters[0];
Object tokp = parameters[1];
if (jsonArrp != null) {
try {
//得到第一个参数
String jsonStrArr = PrimitiveObjectInspectorUtils.getString(jsonArrp, iterateInputJsonArrStringOIParamOne);
//得到第二个参数
Integer tokNum = PrimitiveObjectInspectorUtils.getInt(tokp, iterateInputTopNumOIParamTow);
if (!Strings.isNullOrEmpty(jsonStrArr)) {
//写入中间结果对象
((MergeArrAccountStatTopNAggregationBuffer) agg).setTopKNum(tokNum);
((MergeArrAccountStatTopNAggregationBuffer) agg).addJsonArr(jsonStrArr);
}
} catch (RuntimeException e) {
LOG.warn(getClass().getSimpleName() + " " + StringUtils.stringifyException(e));
LOG.warn(getClass().getSimpleName() + " ignoring similar exceptions.");
}
}
}

/**
* map 阶段过去以后在这个阶段就要输出了
* @param agg 声明的中间聚合对象,map阶段执行完以后的输入
* @return partialResult terminate的序列化输出
* @throws HiveException
*/
public Object terminatePartial(AggregationBuffer agg) throws HiveException {
MergeArrAccountStatTopNAggregationBuffer mergeArrAccountStatTopNAggregationBuffer=(MergeArrAccountStatTopNAggregationBuffer) agg;
HashMap<String, JSONObject> map = mergeArrAccountStatTopNAggregationBuffer.getMap();
Integer topKNum = mergeArrAccountStatTopNAggregationBuffer.getTopKNum();
Text jsonArrStrTemp = new Text(JSON.toJSONString(map.values()));
((Text) partialResult[0]).set(jsonArrStrTemp);
((IntWritable) partialResult[1]).set(topKNum);
return partialResult;
}

/**
* map到reduce阶段的输入
* @param agg 中间结果
* @param partial terminate传过来的结构体struct
* The partial aggregation result.
* @throws HiveException
*/
public void merge(AggregationBuffer agg, Object partial) throws HiveException {
if (null == partial) {
return;
}
MergeArrAccountStatTopNAggregationBuffer mergeArrAccountStatTopNAggregationBuffer=(MergeArrAccountStatTopNAggregationBuffer) agg;
Object partialJsonArrString = mergeInputStructOI.getStructFieldData(partial, jsonArrString);
Object partialTopNum = mergeInputStructOI.getStructFieldData(partial, topNum);
mergeArrAccountStatTopNAggregationBuffer.setTopKNum(topNumFieldOI.get(partialTopNum));
mergeArrAccountStatTopNAggregationBuffer.addJsonArr(jsonArrStringFieldOI.getPrimitiveJavaObject(partialJsonArrString));
}

/**
* 再merge阶段以后
* @param agg
* @return
* @throws HiveException
*/
public Object terminate(AggregationBuffer agg) throws HiveException {
HashMap<String, JSONObject> map = ((MergeArrAccountStatTopNAggregationBuffer) agg).getMap();
Collection<JSONObject> values = map.values();
List<JSONObject> sorted = values.stream().sorted((o1, o2) -> o2.getInteger("interactive").compareTo(o1.getInteger("interactive"))).collect(Collectors.toList());
List<JSONObject> listResult = sorted.subList(0, Math.min(((MergeArrAccountStatTopNAggregationBuffer) agg).getTopKNum(), sorted.size()));
result.set(JSON.toJSONString(listResult));
return result;
}
}

案例二(一个参数)

public class MergeArrBrandStatUDAF extends AbstractGenericUDAFResolver {
@Override
public GenericUDAFEvaluator getEvaluator(GenericUDAFParameterInfo info) throws SemanticException {
return getEvaluator(info.getParameters());
}

@Override
public GenericUDAFEvaluator getEvaluator(TypeInfo[] parameters) throws SemanticException {
if (parameters.length != 1) {
throw new UDFArgumentTypeException(parameters.length - 1,
"Exactly one argument is expected.");
}

if (parameters[0].getCategory() != ObjectInspector.Category.PRIMITIVE) {
throw new UDFArgumentTypeException(0,
"Only primitive type arguments are accepted but "
+ parameters[0].getTypeName() + " is passed.");
}
switch (((PrimitiveTypeInfo) parameters[0]).getPrimitiveCategory()) {
case STRING:
// 在这做类型检查以及选择指定的 Evaluator
return new MergeAarrBrandStatUDAFEvaluator();
default:
throw new UDFArgumentTypeException(0,
"Only string type arguments are accepted but " + parameters[0].getTypeName() + " is passed.");
}
}
}

中间结果

public class MergeArrBrandStatUDAFAggregationBuffer extends GenericUDAFEvaluator.AbstractAggregationBuffer {
private HashMap<String, JSONObject> map = Maps.newHashMap();

public void addJsonArr(String jsonArr) {
synchronized (jsonArr) {
if (!Strings.isNullOrEmpty(jsonArr)) {
JSONArray jsonArray = JSON.parseArray(jsonArr);
for (Object o : jsonArray) {
JSONObject jsonObject = JSON.parseObject(o.toString());
String brandId = jsonObject.getString("brand_id");
if(map.containsKey(brandId)) {
JSONObject existsObj = map.get(brandId);
existsObj.put("price", existsObj.getDoubleValue("price") + jsonObject.getDouble("price"));
existsObj.put("works_num", existsObj.getIntValue("works_num") + jsonObject.getIntValue("works_num"));
existsObj.put("interactive", existsObj.getIntValue("interactive") + jsonObject.getIntValue("interactive"));
map.put(brandId, existsObj);
} else {
map.put(brandId, jsonObject);
}
}
}
}
}

public HashMap<String, JSONObject> getMap() {
return map;
}

public void setMap(HashMap<String, JSONObject> map) {
this.map = map;
}
}

执行的函数

public class MergeAarrBrandStatUDAFEvaluator extends GenericUDAFEvaluator {
static final Log LOG = LogFactory.getLog(MergeArrAccountStatTopNEvaluator.class.getName());

// Iterate 输入
private PrimitiveObjectInspector iterateInputJsonArrStringInputIO;
// Merge 输入
private PrimitiveObjectInspector mergejsonArrStringFieldInputIO;
// TerminatePartial 输出
private StringObjectInspector terminatePartialOutput;
// Terminate 输出
Text result;

/**
* 这个方法主要是指定的时候,每一个阶段的类型声明,包括输入和输出的类型
*
* @param mode 执行mapreduce的不同的阶段
* @param parameters
* @return
* @throws HiveException
*/
@Override
public ObjectInspector init(Mode mode, ObjectInspector[] parameters) throws HiveException {
super.init(mode, parameters);
// map输入数据类型说明
// 初始化输入参数(map阶段)这里先声明输入map端的数据类型
if (mode == Mode.PARTIAL1 || mode == Mode.COMPLETE) {
// 原始数据 (输入的第一个参数)
iterateInputJsonArrStringInputIO = (PrimitiveObjectInspector) parameters[0];
} else {
// Merge 输入类型说明(struct的数据也就是第一个参数,是下面terminatePartial的输出)
// 部分聚合数据(这里是map到reduce的merger阶段的数据类型说明)
mergejsonArrStringFieldInputIO = (PrimitiveObjectInspector) parameters[0];
}
// TerminatePartial 输出的数据类型声明,这里要声明可以序列化的类型(由于这里声明的是struct类型的输出,那么上面的Merge就是struct的输入)
// 初始化输出
if (mode == Mode.PARTIAL1 || mode == Mode.PARTIAL2) {
// // 最终结果
result = new Text("");
return PrimitiveObjectInspectorFactory.writableStringObjectInspector;
} else {
// Terminate 输出,由于结果是字符串,所以这里用到了text类型
result = new Text("");
return PrimitiveObjectInspectorFactory.writableStringObjectInspector;
}
}

public AggregationBuffer getNewAggregationBuffer() throws HiveException {
MergeArrBrandStatUDAFAggregationBuffer mergeArrBrandStatUDAFAggregationBuffer = new MergeArrBrandStatUDAFAggregationBuffer();
reset(mergeArrBrandStatUDAFAggregationBuffer);
return mergeArrBrandStatUDAFAggregationBuffer;
}

@Override
public void reset(AggregationBuffer agg) throws HiveException {
((MergeArrBrandStatUDAFAggregationBuffer) agg).setMap(Maps.newHashMap());
}

/**
* map 阶段的数据输入
*
* @param agg udaf的中间结果
* @param parameters 执行group by 以后的每一条数据
* @throws HiveException
*/
public void iterate(AggregationBuffer agg, Object[] parameters) throws HiveException {
Object jsonArr = parameters[0];
if (jsonArr != null) {
try {
//得到第一个参数
String jsonStrArr = PrimitiveObjectInspectorUtils.getString(jsonArr, iterateInputJsonArrStringInputIO);
if (!Strings.isNullOrEmpty(jsonStrArr)) {
((MergeArrBrandStatUDAFAggregationBuffer) agg).addJsonArr(jsonStrArr);
}
} catch (RuntimeException e) {
LOG.warn(getClass().getSimpleName() + " " + StringUtils.stringifyException(e));
LOG.warn(getClass().getSimpleName() + " ignoring similar exceptions.");
}
}
}

/**
* map 阶段过去以后在这个阶段就要输出了
*
* @param agg 声明的中间聚合对象,map阶段执行完以后的输入
* @return partialResult terminate的序列化输出
* @throws HiveException
*/
public Object terminatePartial(AggregationBuffer agg) throws HiveException {
MergeArrBrandStatUDAFAggregationBuffer mergeArrBrandStatUDAFAggregationBuffer = (MergeArrBrandStatUDAFAggregationBuffer) agg;
HashMap<String, JSONObject> map = mergeArrBrandStatUDAFAggregationBuffer.getMap();
result.set(JSON.toJSONString(map.values()));
return result;
}

/**
* map到reduce阶段的输入
*
* @param agg 中间结果
* @param partial terminate传过来的结构体struct
* The partial aggregation result.
* @throws HiveException
*/
public void merge(AggregationBuffer agg, Object partial) throws HiveException {
if (null == partial) {
return;
}
MergeArrBrandStatUDAFAggregationBuffer mergeArrBrandStatUDAFAggregationBuffer = (MergeArrBrandStatUDAFAggregationBuffer) agg;
String v = PrimitiveObjectInspectorUtils.getString(partial, mergejsonArrStringFieldInputIO);
if (!Strings.isNullOrEmpty(v)) {
mergeArrBrandStatUDAFAggregationBuffer.addJsonArr(v);
}

}

/**
* 再merge阶段以后
*
* @param agg
* @return
* @throws HiveException
*/
public Object terminate(AggregationBuffer agg) throws HiveException {
HashMap<String, JSONObject> map = ((MergeArrBrandStatUDAFAggregationBuffer) agg).getMap();
Collection<JSONObject> values = map.values();
result.set(JSON.toJSONString(values));
return result;
}
}