【发布时间】:2019-05-16 14:44:53
【问题描述】:
在此查询中,我得到一个数据框,其中包含一列 5d 欧几里得点(存储为双精度数组)。我需要找到所有可用的平均距离。也就是说,对于每个点 a,我计算数据框中到其他点 b 的距离,并找到这些距离的平均值。请注意,我不希望对这个问题进行任何数学方法或简化。数据框有两列,unique_id 和 vector。
我可以进行查询,但仅针对以下方式的 1 点。 UDF 距离计算存储数组(即包装数组)和给定数组之间的距离。然而,很明显,这种方法只适用于一点。另外,我尝试将数据集传递给静态函数。但是每次我这样做我都会得到一个“Invalid Tree: null”,即对象一进入函数就变为null......最后,我想到了做一个UDAF,但我意识到这不是一个适当的聚合函数。对此的任何帮助将不胜感激!
(注意:这段代码是java,但应该和其他语言没有太大区别)
long equal = 2;
WrappedArray<Double> num = (WrappedArray<Double> spo.select("vectors")
.filter(col("unique_id").equalTo(equal)).first().get(0);
List<Double> frameList = scala.collection.JavaConverters.seqAsJavaList(num);
double[] array_answer = frameList.stream().mapToDouble(Double::doubleValue).toArray();
UserDefinedFunction compare = udf(
(WrappedArray<Double> array) -> cosine_distance(array, array_answer), DataTypes.DoubleType
);
double answer = (double) spo.select("vectors").filter(col("unique_id").notEqual(equal))
.withColumn("calc", compare.apply(col("vectors")))
.select(avg("calc")).first().get(0);
System.out.println(answer);
【问题讨论】:
标签: apache-spark apache-spark-sql