如何从Spark Dataset<Row>创建按Location分组的Map<String, List<Row>> Dataset
问题:基于Spark Dataset按Location分组生成Map<String, List>
原始Dataset结构
userid username location 1 ram karnataka 2 shyam rajasthan 3 hemant himachal 4 raju chitoor 5 ravi rajasthan 6 titoo himachal 7 shukhi chitoor 8 raju chennai
需求说明
需要基于location列创建一个Map<String, List<Row>>类型的结果,其中key为location列的值,value为对应分组的Row列表,预期输出格式如下:
user_location userid username location karnataka 1 ram karnataka rajasthan 2 shyam rajasthan 5 ravi rajasthan himachal 3 hemant himachal 6 titoo himachal chitoor 4 raju chitoor 7 shukhi chitoor
已尝试方法及问题
RDD方式
已实现分组计数,但无法直接获取对应location的Row列表,代码如下:
System.out.println("reading user records"); Dataset<Row> user_df = MongoSpark.read(my_spark) .option("collection", "user") .option("readPreference.name", "secondaryPreferred") .load(); strLogMsg = "printing user_df records"; System.out.println(strLogMsg); user_df.show(2); strLogMsg = "creating user_df groups by location"; System.out.println(strLogMsg); JavaPairRDD<String, Row> jpRDD = user_df.toJavaRDD().mapToPair(new PairFunction<Row, String, Row>() { public Tuple2<String, Row> call(Row row) throws Exception { return new Tuple2<String, Row>((String) row.getAs("location"), row); } }); JavaPairRDD<String, Iterable<Row>> user_location_list = jpRDD.groupByKey(); strLogMsg = "printing user_location_list records count"; System.out.println(strLogMsg); System.out.println(user_location_list.count());
Spark SQL方式
尝试使用Spark SQL分组时抛出org.apache.spark.sql.AnalysisException: [UNRESOLVED_COLUMN.WITH_SUGGESTION]错误,代码如下:
String[] arColNames = temp_df.columns(); String strCols =""; for(String strcol : arColNames){ if(strCols.isEmpty()){ strCols = strcol; }else{ strCols += ","+strcol; } } System.out.println("strCols - " + strCols); Dataset<Row> user_by_location_ds = temp_df.groupBy("location").agg(functions.collect_list(functions.struct(strCols)).as("loc_list"));
正确实现方案
方案1:RDD方式修正
groupByKey返回的Iterable<Row>可以直接转为List<Row>,之后可按需转为Map<String, List<Row>>或Dataset:
// 将Iterable<Row>转为List<Row> JavaPairRDD<String, List<Row>> userLocationListRDD = user_location_list.mapValues(iter -> { List<Row> rowList = new ArrayList<>(); iter.forEach(rowList::add); return rowList; }); // 转为Map<String, List<Row>> Map<String, List<Row>> locationToRowsMap = userLocationListRDD.collectAsMap();
方案2:Spark SQL方式修正
问题出在functions.struct(strCols)的参数传递,struct需要传入Column对象数组而非拼接的字符串。正确写法如下:
import org.apache.spark.sql.Column; import static org.apache.spark.sql.functions.*; import java.util.Arrays; // 将列名转为Column对象数组 Column[] columns = Arrays.stream(temp_df.columns()) .map(col) .toArray(Column[]::new); // 按location分组,收集对应行的结构体列表 Dataset<Row> user_by_location_ds = temp_df.groupBy("location") .agg(collect_list(struct(columns)).as("loc_list")); // 若需要转为Map<String, List<Row>>,可进一步处理 Map<String, List<Row>> resultMap = user_by_location_ds.javaRDD().mapToPair(row -> { String location = row.getString(0); List<Row> rowList = row.getList(1); return new Tuple2<>(location, rowList); }).collectAsMap();
内容的提问来源于stack exchange,提问作者Chandeshwar Prasad sit
相关产品推荐
相关产品推荐

