本篇博客将介绍Spark RDD的Map系算子的基本用法。
1、map
map将RDD的元素一个个传入call方法,经过call方法的计算之后,逐个返回,生成新的RDD,计算之后,记录数不会缩减。示例代码,将每个数字加10之后再打印出来, 代码如下
1import java.util.Arrays; 2 3import org.apache.spark.SparkConf; 4import org.apache.spark.api.java.JavaRDD; 5import org.apache.spark.api.java.JavaSparkContext; 6import org.apache.spark.api.java.function.Function; 7import org.apache.spark.api.java.function.VoidFunction; 8 9public class Map { 10 public static void main(String[] args) { 11 SparkConf conf = new SparkConf().setAppName("spark map").setMaster("local[*]"); 12 JavaSparkContext javaSparkContext = new JavaSparkContext(conf); 13 JavaRDD<Integer> listRDD = javaSparkContext.parallelize(Arrays.asList(1, 2, 3, 4)); 14 15 JavaRDD<Integer> numRDD = listRDD.map(new Function<Integer, Integer>() { 16 @Override 17 public Integer call(Integer num) throws Exception { 18 return num + 10; 19 } 20 }); 21 numRDD.foreach(new VoidFunction<Integer>() { 22 @Override 23 public void call(Integer num) throws Exception { 24 System.out.println(num); 25 } 26 }); 27 } 28 29}
执行结果:

2、flatMap
flatMap和map的处理方式一样,都是把原RDD的元素逐个传入进行计算,但是与之不同的是,flatMap返回值是一个Iterator,也就是会一生多,超生
1import java.util.Arrays; 2import java.util.Iterator; 3 4import org.apache.spark.SparkConf; 5import org.apache.spark.api.java.JavaRDD; 6import org.apache.spark.api.java.JavaSparkContext; 7import org.apache.spark.api.java.function.FlatMapFunction; 8import org.apache.spark.api.java.function.VoidFunction; 9 10public class FlatMap { 11 public static void main(String[] args) { 12 SparkConf conf = new SparkConf().setAppName("spark map").setMaster("local[*]"); 13 JavaSparkContext javaSparkContext = new JavaSparkContext(conf); 14 JavaRDD<String> listRDD = javaSparkContext 15 .parallelize(Arrays.asList("hello wold", "hello java", "hello spark")); 16 JavaRDD<String> rdd = listRDD.flatMap(new FlatMapFunction<String, String>() { 17 private static final long serialVersionUID = 1L; 18 19 @Override 20 public Iterator<String> call(String input) throws Exception { 21 return Arrays.asList(input.split(" ")).iterator(); 22 } 23 }); 24 rdd.foreach(new VoidFunction<String>() { 25 private static final long serialVersionUID = 1L; 26 @Override 27 public void call(String num) throws Exception { 28 System.out.println(num); 29 } 30 }); 31 } 32 33}
执行结果:

3、mapPartitions
mapPartitions一次性将整个分区的数据传入函数进行计算,适用于一次性聚会整个分区的场景
1public class MapPartitions { 2 public static void main(String[] args) { 3 SparkConf conf = new SparkConf().setAppName("spark map").setMaster("local[*]"); 4 JavaSparkContext javaSparkContext = new JavaSparkContext(conf); 5 JavaRDD<String> listRDD = javaSparkContext.parallelize(Arrays.asList("hello", "java", "wold", "spark"), 2); 6 7 /** 8 * mapPartitions回调的接口也是FlatMapFunction,FlatMapFunction的第一个泛型是Iterator表示传入的数据, 9 * 第二个泛型表示返回数据的类型 10 * 11 * mapPartitions传入FlatMapFunction接口处理的数据是一个分区的数据,所以,如果一个分区数据过大,会导致内存溢出 12 * 13 */ 14 JavaRDD<String> javaRDD = listRDD.mapPartitions(new FlatMapFunction<Iterator<String>, String>() { 15 int i = 0; 16 17 @Override 18 public Iterator<String> call(Iterator<String> input) throws Exception { 19 List<String> list = new ArrayList<String>(); 20 while (input.hasNext()) { 21 list.add(input.next() + i); 22 ++i; 23 } 24 return list.iterator(); 25 } 26 }); 27 28 javaRDD.foreach(new VoidFunction<String>() { 29 @Override 30 public void call(String t) throws Exception { 31 System.out.println(t); 32 } 33 }); 34 } 35 36}
运行结果:

上面的运算结果,后面的尾标只有0和1,说明FlatMapFunction被调用了两次,与MapPartitions的功能吻合。
4、mapPartitionsWithIndex
mapPartitionsWithIndex和mapPartitions一样,一次性传入整个分区的数据进行处理,但是不同的是,这里会传入分区编号进来
1public class mapPartitionsWithIndex { 2 public static void main(String[] args) { 3 SparkConf conf = new SparkConf().setAppName("spark map").setMaster("local[*]"); 4 JavaSparkContext javaSparkContext = new JavaSparkContext(conf); 5 JavaRDD<String> listRDD = javaSparkContext.parallelize(Arrays.asList("hello", "java", "wold", "spark"), 2); 6 7 /** 8 *和mapPartitions一样,一次性传入整个分区的数据进行处理,但是不同的是,这里会传入分区编号进来 9 * 10 */ 11 JavaRDD<String> javaRDD = listRDD.mapPartitionsWithIndex(new Function2<Integer, Iterator<String>, Iterator<String>>() { 12 13 @Override 14 public Iterator<String> call(Integer v1, Iterator<String> v2) throws Exception { 15 List<String> list = new ArrayList<String>(); 16 while (v2.hasNext()) { 17 list.add(v2.next() + "====分区编号:"+v1); 18 } 19 return list.iterator(); 20 } 21 22 },true); 23 24 javaRDD.foreach(new VoidFunction<String>() { 25 @Override 26 public void call(String t) throws Exception { 27 System.out.println(t); 28 } 29 }); 30 } 31 32}
执行结果:
