【发布时间】:2018-07-02 15:27:08
【问题描述】:
我确信这里有一个快速修复,但我在为 pyspark DF 上的基本矢量操作创建 udf 时遇到问题。
我有:
300 维的密集向量
具有 500K 密集向量列的 Pyspark DF,每个向量 300 维
简而言之,我想找出 DF 中的 (2) 行中的哪一行与所讨论的 (1) Dense 向量具有最高的余弦相似度。 在对所有向量进行归一化之后,我'将能够通过对每一行执行相关向量的点积然后返回最大值来实现这一点。
我的代码:
value = df_other.select('vec_norm').collect()[0][0] #Pulling from another DF
def dot_product(vec):
dot_value = value.dot(DenseVector(vec[3]))
return dot_value
dot_product_udf = udf(dot_product, FloatType())
df_dot = df.withColumn('cos_dis',dot_product_udf(df['vec_norm']))
print df_dot.rdd.max(key=lambda x: x["cos_dis"])[0]
错误:
Py4JJavaError: An error occurred while calling z:org.apache.spark.api.python.PythonRDD.collectAndServe.
...
File "/usr/hdp/current/spark2-client/python/lib/pyspark.zip/pyspark/ml/linalg/__init__.py", line 402, in __len__
return len(self.array)
TypeError: len() of unsized object
如果我尝试使用 numpy 进行计算,我会遇到类似的问题:
...
def dot_product(vec):
#dot_value = value.dot(DenseVector(vec[3]))
dot_value = sum(value * DenseVector(vec[3]))
return dot_value
dot_product_udf = udf(dot_product, FloatType())
...
错误:
Py4JJavaError: An error occurred while calling z:org.apache.spark.api.python.PythonRDD.collectAndServe.
: org.apache.spark.SparkException: Job aborted due to stage failure: Task 2 in stage 72.0 failed 1 times, most recent failure: Lost task 2.0 in stage 72.0 (TID 632, localhost, executor driver): net.razorvine.pickle.PickleException: expected zero arguments for construction of ClassDict (for numpy.dtype)
到目前为止,我已使用以下问题进行故障排除,但无法解决问题(我猜测它是矢量类型问题):
- Issue with UDF on a column of Vectors in PySpark DataFrame
- Spark Error:expected zero arguments for construction of ClassDict (for numpy.core.multiarray._reconstruct)
- Spark __getnewargs__ error
- Error with "len() of unsized object"
任何帮助/建议都非常感谢!
编辑
样本数据:
> print type(value), len(value), value
<class 'pyspark.ml.linalg.DenseVector'> 300 [0.0667470050056,0.0439160518808...]
> df_value = df.select('vec_norm').collect()[0][0]
> print len(df_value)
> df.select('vec_norm').show(truncate=100)
300 <class 'pyspark.ml.linalg.DenseVector'>
+----------------------------------------------------------------------------------------------------+
| vec_norm|
+----------------------------------------------------------------------------------------------------+
|[-0.033044380089015266,0.09674768943906177,0.08259697668541087,0.04247286602516604,0.037005449248...|
|[-0.06890003507034705,0.06019625255379143,0.04288672222032615,-2.714061064477613E-4,0.02868655951...|
【问题讨论】:
标签: python apache-spark vector pyspark spark-dataframe