【问题标题】:How to cast Dataset<Row> columns to Non- Primitive Data Type如何将 Dataset<Row> 列转换为非原始数据类型
【发布时间】:2019-03-27 11:15:40
【问题描述】:

我有一个Dataset&lt;Row&gt;,其中有四列,两列是非原始数据类型List&lt;Long&gt; and List&lt;String&gt;

  +------+---------------+---------------------------------------------+---------------+
  |    Id| value         |     time                                      |aggregateType  |
  +------+---------------+---------------------------------------------+---------------+
  |0001  |  [1.5,3.4,4.5]| [1551502200000,1551502200000,1551502200000] | Sum             |
  +------+---------------+---------------------------------------------+---------------+

我有一个 UDF3,它接受三个参数并返回一个 DoubleUDF3&lt;String,List&lt;Long&gt;,List&lt;String&gt;,Double&gt;

所以当我调用 UDF 时,它会抛出一个异常

错误

caused by java.lang.classcastexception scala.collection.mutable.wrappedarray$ofref cannot be cast to java.lang.List

但如果我将类型更改为 String 就像 UDF3&lt;String,String,String,Double&gt; 一样,它不会抱怨。

抛出异常的代码

 UDF3<String,List<Long>,List<String>,Double> getAggregate = new UDF3<String,List<Long>,List<String>,Double>() {

 public Double call(String t1,List<Long> t2,List<String> t3) throws Exception {

 //do some process to return double

  return double;
  }

  sparkSession.udf().register("getAggregate_UDF",getAggregate, DataTypes.DoubleType);

  inputDS = inputDs.withColumn("value_new",callUDF("getAggregate_UDF",col("aggregateType"),col("time"),col("value")));

将所有类型改为String后的代码

 UDF3<String,String,String,Double> getAggregate = new UDF3<String,String,String,Double>() {

 public Double call(String t1,String t2,String t3) throws Exception {

 //code to convert t2 and t3 to List<Long> and List<String> respectively

 //do some process to return double

  return double;
  }

  sparkSession.udf().register("getAggregate_UDF",getAggregate, DataTypes.DoubleType);

  inputDS = inputDs.withColumn("value_new",callUDF("getAggregate_UDF",col("aggregateType"),col("time").cast("String"),col("value").cast("String")));

上述代码有效,但需要手动转换String to List

需要帮助

I) 如何在数据集中转换非原始数据类型List&lt;Long&gt; and List&lt;String&gt; 以克服caused by java.lang.classcastexception scala.collection.mutable.wrappedarray$ofref cannot be cast to java.lang.List

II) 如果有任何解决方法,请建议我

谢谢。

【问题讨论】:

  • 您能否粘贴您的 printSchema 让我们知道实际的数据类型。

标签: apache-spark apache-spark-sql


【解决方案1】:

您的 UDF 将始终接收 WrappedArray 实例而不是 List,因为这是引擎存储它们的方式。

你需要这样写:

import scala.collection.mutable.WrappedArray;
import scala.collection.JavaConversions;

UDF3<String, WrappedArray<Long>, WrappedArray<String>, Double> myUDF = new UDF3<String, WrappedArray<Long>, WrappedArray<String>, Double> () {
      public Double call(String param1, WrappedArray<Long> param2, WrappedArray<String> param3) throws Exception {
        List<Long> param1AsList = JavaConversions.seqAsJavaList(param1);
        List<String> param2AsList = JavaConversions.seqAsJavaList(param2);

        ... do work ...

        return myDoubleResult;
    }
};

【讨论】:

  • 嗨@rluta,我已经尝试过像上面那样的UDF3,它抛出了上面提到的错误,所以我将所有类型都更改为String。请看一下我现在更新的问题。
  • 对不起,我以为你在使用 Scala。对于java,您需要专门使用签名中的 WrappedArray 类和JavaConversions 来重铸为列表。我会更新答案
  • 非常感谢。你的解决方案帮助了我:)
【解决方案2】:

这是我的例子,你必须使用 WrappedArray 来接收数组并转换为列表

 /*
     +------+---------------+---------------------------------------------+---------------+
     |    Id| value         |     time                                      |aggregateType  |
     +------+---------------+---------------------------------------------+---------------+
     |0001  |  [1.5,3.4,4.5]| [1551502200000,1551502200000,1551502200000] | Sum             |
     +------+---------------+---------------------------------------------+---------------+
     **/

    StructType dataSchema = new StructType(new StructField[] {createStructField("Id", DataTypes.StringType, true),
                                                              createStructField("value",
                                                                                DataTypes.createArrayType(DataTypes.DoubleType,
                                                                                                          false),
                                                                                false),

                                                              createStructField("time",
                                                                                DataTypes.createArrayType(DataTypes.LongType,
                                                                                                          false),
                                                                                false),
                                                              createStructField("aggregateType",
                                                                                DataTypes.StringType,
                                                                                true),});

    List<Row> data = new ArrayList<>();

    data.add(RowFactory.create("0001",
                               Arrays.asList(1.5, 3.4, 4.5),
                               Arrays.asList(1551502200000L, 1551502200000L, 1551502200000L),
                               "sum"));
    Dataset<Row> example = spark.createDataFrame(data, dataSchema);
    example.show(false);

    UDF3<String, WrappedArray<Long>, WrappedArray<Double>, Double> myUDF = (param1, param2, param3) -> {

        List<Long> param1AsList = JavaConversions.seqAsJavaList(param2);
        List<Double> param2AsList = JavaConversions.seqAsJavaList(param3);

        //Example
        double myDoubleResult = 0;
        if ("sum".equals(param1)) {

            myDoubleResult = param2AsList.stream()
                                         .mapToDouble(f -> f)
                                         .sum();
        }

        return myDoubleResult;
    };

    spark.udf()
         .register("myUDF", myUDF, DataTypes.DoubleType);

    example = example.withColumn("new", callUDF("myUDF", col("aggregateType"), col("time"), col("value")));
    example.show(false);

您可以从github获取它

【讨论】:

  • 嗨,howie,我对如何在下面的代码 inputDS = inputDs.withColumn("value_new",callUDF("getAggregate_UDF",//how to pass the entire Row here)); 中将整个 Row 传递给 UDF 函数有疑问
  • 这一行有多个标记 - 语法错误插入';'完成 LocalVariableDeclarationStatement - 语法错误插入 '[]' 来完成 Dimension -WrappedArray 是原始类型。应参数化对泛型类型 WrappedArray 的引用 - 标记“splitSeqUDF”上的语法错误,时间标记后预期 AnnotationName - 参数 splitSeqUDF 的非法修饰符,允许使用最终结果
猜你喜欢
  • 2022-11-07
  • 1970-01-01
  • 2019-04-20
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2020-05-20
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多