【问题标题】:Create a custom Transformer In Java spark ml在 Java spark ml 中创建自定义 Transformer
【发布时间】:2017-11-09 22:09:39
【问题描述】:

我想用 Java 创建一个自定义 Spark Transformer。

Transformer 是文本预处理器,其作用类似于 Tokenizer。它接受一个输入列和一个输出列作为参数。

我环顾四周,发现了 2 个 Scala Traits HasInputCol 和 HasOutputCol。

如何创建一个扩展 Transformer 并实现 HasInputCol 和 OutputCol 的类?

我的目标是拥有这样的东西。

   // Dataset that have a String column named "text"
   DataSet<Row> dataset;

   CustomTransformer customTransformer = new CustomTransformer();
   customTransformer.setInputCol("text");
   customTransformer.setOutputCol("result");

   // result that have 2 String columns named "text" and "result"
   DataSet<Row> result = customTransformer.transform(dataset);

【问题讨论】:

    标签: java scala apache-spark apache-spark-mllib transformer


    【解决方案1】:

    正如SergGr 建议的那样,您可以扩展UnaryTransformer。然而,这是相当棘手的。

    注意:以下所有 cmets 适用于 Spark 版本 2.2.0。

    要解决SPARK-12606 中描述的问题,他们得到"...Param null__inputCol does not belong to...",您应该像这样实现String uid()

    @Override
    public String uid() {
        return getUid();
    }
    
    private String getUid() {
    
        if (uid == null) {
            uid = Identifiable$.MODULE$.randomUID("mycustom");
        }
        return uid;
    }
    

    显然他们正在构造函数中初始化 uid。但问题是 UnaryTransformer 的 inputCol(和 outputCol)在 uid 在继承类中初始化之前被初始化。见HasInputCol

    final val inputCol: Param[String] = new Param[String](this, "inputCol", "input column name")
    

    Param 是这样构造的:

    def this(parent: Identifiable, name: String, doc: String) = this(parent.uid, name, doc)
    

    因此,当评估 parent.uid 时,将调用自定义的 uid() 实现,此时 uid 仍然为空。通过使用惰性求值实现 uid(),您可以确保 uid() 永远不会返回 null。

    在你的情况下:

    Param d7ac3108-799c-4aed-a093-c85d12833a4e__inputCol does not belong to fe3d99ba-e4eb-4e95-9412-f84188d936e3
    

    好像有点不一样。因为"d7ac3108-799c-4aed-a093-c85d12833a4e" != "fe3d99ba-e4eb-4e95-9412-f84188d936e3",看起来您对uid() 方法的实现在每次调用时都会返回一个新值。也许在你的情况下它是这样实现的:

    @Override
    public String uid() {
        return Identifiable$.MODULE$.randomUID("mycustom");
    }
    

    顺便说一下,在扩展UnaryTransformer时,请确保转换函数为Serializable

    【讨论】:

      【解决方案2】:

      您可能希望从org.apache.spark.ml.UnaryTransformer 继承您的CustomTransformer。你可以试试这样的:

      import org.apache.spark.ml.UnaryTransformer;
      import org.apache.spark.ml.util.Identifiable$;
      import org.apache.spark.sql.types.DataType;
      import org.apache.spark.sql.types.DataTypes;
      import scala.Function1;
      import scala.collection.JavaConversions$;
      import scala.collection.immutable.Seq;
      
      import java.util.Arrays;
      
      public class MyCustomTransformer extends UnaryTransformer<String, scala.collection.immutable.Seq<String>, MyCustomTransformer>
      {
          private final String uid = Identifiable$.MODULE$.randomUID("mycustom");
      
          @Override
          public String uid()
          {
              return uid;
          }
      
      
          @Override
          public Function1<String, scala.collection.immutable.Seq<String>> createTransformFunc()
          {
              // can't use labmda syntax :(
              return new scala.runtime.AbstractFunction1<String, Seq<String>>()
              {
                  @Override
                  public Seq<String> apply(String s)
                  {
                      // do the logic
                      String[] split = s.toLowerCase().split("\\s");
                      // convert to Scala type
                      return JavaConversions$.MODULE$.iterableAsScalaIterable(Arrays.asList(split)).toList();
                  }
              };
          }
      
      
          @Override
          public void validateInputType(DataType inputType)
          {
              super.validateInputType(inputType);
              if (inputType != DataTypes.StringType)
                  throw new IllegalArgumentException("Input type must be string type but got " + inputType + ".");
          }
      
          @Override
          public DataType outputDataType()
          {
              return DataTypes.createArrayType(DataTypes.StringType, true); // or false? depends on your data
          }
      }
      

      【讨论】:

      • 这不起作用。我想这是因为一个错误。我得到java.lang.IllegalArgumentException: requirement failed: Param d7ac3108-799c-4aed-a093-c85d12833a4e__inputCol does not belong to fe3d99ba-e4eb-4e95-9412-f84188d936e3.
      • @LonsomeHell,请仔细检查一下,您确定为它配置了有效的输入列吗?
      • 是的,我使用了带有有效列名的 setInput。
      • 我认为它与这个错误有关issues.apache.org/jira/browse/SPARK-12606
      • 我喜欢使用 AbstractFunction1 的技巧,所以你不需要实现所有的方法。
      【解决方案3】:

      我参加聚会有点晚了,但这里有一些自定义 Java Spark 转换的示例:https://github.com/dafrenchyman/spark/tree/master/src/main/java/com/mrsharky/spark/ml/feature

      这是一个只有一个输入列的示例,但您可以按照相同的模式轻松添加一个输出列。但是,这并没有实现读取器和写入器。您需要查看上面的链接以了解如何操作。

      public class DropColumns extends Transformer implements Serializable, 
      DefaultParamsWritable {
      
          private StringArrayParam _inputCols;
          private final String _uid;
      
          public DropColumns(String uid) {
              _uid = uid;
          }
      
          public DropColumns() {
              _uid = DropColumns.class.getName() + "_" + 
      UUID.randomUUID().toString();
          }
      
          // Getters
          public String[] getInputCols() { return get(_inputCols).get(); }
      
         // Setters
         public DropColumns setInputCols(String[] columns) {
             _inputCols = inputCols();
             set(_inputCols, columns);
             return this;
         }
      
      public DropColumns setInputCols(List<String> columns) {
          String[] columnsString = columns.toArray(new String[columns.size()]);
          return setInputCols(columnsString);
      }
      
      public DropColumns setInputCols(String column) {
          String[] columns = new String[]{column};
          return setInputCols(columns);
      }
      
      // Overrides
      @Override
      public Dataset<Row> transform(Dataset<?> data) {
          List<String> dropCol = new ArrayList<String>();
          Dataset<Row> newData = null;
          try {
              for (String currColumn : this.get(_inputCols).get() ) {
                  dropCol.add(currColumn);
              }
              Seq<String> seqCol = JavaConverters.asScalaIteratorConverter(dropCol.iterator()).asScala().toSeq();      
              newData = data.drop(seqCol);
          } catch (Exception ex) {
              ex.printStackTrace();
          }
          return newData;
      }
      
      @Override
      public Transformer copy(ParamMap extra) {
          DropColumns copied = new DropColumns();
          copied.setInputCols(this.getInputCols());
          return copied;
      }
      
      @Override
      public StructType transformSchema(StructType oldSchema) {
          StructField[] fields = oldSchema.fields();  
          List<StructField> newFields = new ArrayList<StructField>();
          List<String> columnsToRemove = Arrays.asList( get(_inputCols).get() );
          for (StructField currField : fields) {
              String fieldName = currField.name();
              if (!columnsToRemove.contains(fieldName)) {
                  newFields.add(currField);
              }
          }
          StructType schema = DataTypes.createStructType(newFields);
          return schema;
      }
      
      @Override
      public String uid() {
          return _uid;
      }
      
      @Override
      public MLWriter write() {
          return new DropColumnsWriter(this);
      }
      
      @Override
      public void save(String path) throws IOException {
          write().saveImpl(path);
      }
      
      public static MLReader<DropColumns> read() {
          return new DropColumnsReader();
      }
      
      public StringArrayParam inputCols() {
          return new StringArrayParam(this, "inputCols", "Columns to be dropped");
      }
      
      public DropColumns load(String path) {
          return ( (DropColumnsReader) read()).load(path);
      }
      }
      

      【讨论】:

        【解决方案4】:

        晚会,我还有另一个更新。我很难找到有关将 Spark Transformers 扩展到 Java 的信息,因此我将我的发现发布在这里。

        我也一直在研究 Java 中的自定义转换器。在撰写本文时,包含保存/加载功能要容易一些。可以通过实现 DefaultParamsWritable 创建可写参数。但是,实现 DefaultParamsReadable 仍然会导致我出现异常,但有一个简单的解决方法。

        下面是列重命名器的基本实现:

        public class ColumnRenamer extends Transformer implements DefaultParamsWritable {
            /**
             * A custom Spark transformer that renames the inputCols to the outputCols.
             * 
             * We would also like to implement DefaultParamsReadable<ColumnRenamer>, but
             * there appears to be a bug in DefaultParamsReadable when used in Java, see:
             * https://issues.apache.org/jira/browse/SPARK-17048
             **/
            private final String uid_;
            private StringArrayParam inputCols_;
            private StringArrayParam outputCols_;
            private HashMap<String, String> renameMap;
        
            public ColumnRenamer() {
                this(Identifiable.randomUID("ColumnRenamer"));
            }
        
            public ColumnRenamer(String uid) {
                this.uid_ = uid;
                init();
            }
        
            @Override
            public String uid() {
                return uid_;
            }
        
            @Override
            public Transformer copy(ParamMap extra) {
                return defaultCopy(extra);
            }
        
            /**
             * The below method is a work around, see:
             * https://issues.apache.org/jira/browse/SPARK-17048
             **/
            public static MLReader<ColumnRenamer> read() {
                return new DefaultParamsReader<>();
            }
        
            public Dataset<Row> transform(Dataset<?> dataset) {
                Dataset<Row> transformedDataset = dataset.toDF();
                // Check schema.
                transformSchema(transformedDataset.schema(), true); // logging = true
                // Rename columns.
                for (Map.Entry<String, String> entry: renameMap.entrySet()) {
                    String inputColName = entry.getKey();
                    String outputColName = entry.getValue();
                    transformedDataset = transformedDataset
                        .withColumnRenamed(inputColName, outputColName);
                }
                return transformedDataset;
            }
        
            @Override
            public StructType transformSchema(StructType schema) {
        
                // Validate the parameters here...
                
                String[] inputCols = getInputCols();
                String[] outputCols = getOutputCols();
                // Create rename mapping.
                renameMap = new HashMap<> ();
                for (int i = 0; i < inputCols.length; i++) {
                    renameMap.put(inputCols[i], outputCols[i]);
                }
                // Rename columns.
                ArrayList<StructField> fields = new ArrayList<> ();
                for (StructField field: schema.fields()) {
                    String columnName = field.name();
                    if (renameMap.containsKey(columnName)) {
                        columnName = renameMap.get(columnName);
                    }
                    fields.add(new StructField(
                        columnName, field.dataType(), field.nullable(), field.metadata()
                    ));
                }
                // Return as StructType.
                return new StructType(fields.toArray(new StructField[0]));
            }
        
            private void init() {
                inputCols_ = new StringArrayParam(this, "inputCols", "input column names");
                outputCols_ = new StringArrayParam(this, "outputCols", "output column names");
            }
        
            public StringArrayParam inputCols() {
                return inputCols_;
            }
            public ColumnRenamer setInputCols(String[] value) {
                set(inputCols_, value);
                return this;
            }
            public String[] getInputCols() {
                return getOrDefault(inputCols_);
            }
        
            public StringArrayParam outputCols() {
                return outputCols_;
            }
            public ColumnRenamer setOutputCols(String[] value) {
                set(outputCols_, value);
                return this;
            }
            public String[] getOutputCols() {
                return getOrDefault(outputCols_);
            }
        }
        

        【讨论】:

          猜你喜欢
          • 2015-11-26
          • 1970-01-01
          • 2016-02-04
          • 2017-03-17
          • 1970-01-01
          • 1970-01-01
          • 2016-05-12
          • 1970-01-01
          • 1970-01-01
          相关资源
          最近更新 更多