如何在SparklyR DataFrame中按规则生成自定义分组ID
SparklyR 按自定义规则生成 group_id 解决方案
问题场景
现有如下SparklyR DataFrame,可通过以下代码生成:
# 加载库 library(sparklyr) library(dplyr) library(magrittr) # 连接Spark sc <- spark_connect(master = "local") # 创建示例数据 df <- read.table(text = "firm sector chapter A x 11 B x 12 C y 21 D x 11 E z 31 F y 22 G z 32", header=TRUE) # 复制数据到Spark集群 sdf <- copy_to(sc, df, name = "sdf", overwrite=TRUE)
需要按以下规则生成group列:
- 属于
sector = 'y'的企业,每个单独成组 - 其他
sector的企业,同sector归为同一组
最终期望得到如下结果:
# Source: spark<sdf2> [?? x 4] firm sector chapter group <chr> <chr> <int> <int> 1 A x 11 1 2 B x 12 1 3 C y 21 2 4 D x 11 1 5 E z 31 4 6 F y 22 3 7 G z 32 4
尝试过的报错方法
以下方法在SparklyR中无法正常运行:
sdf %>% mutate(group_id = group_indices(sector))sdf %>% group_by(sector) %>% mutate(group_id = cur_group_id())sdf %>% group_by(sector) %>% mutate(group_id = seq_along(sector))
解决方案
SparklyR需使用分布式兼容的函数实现,不能直接用本地dplyr的序列/分组索引函数。以下是两种可行方案:
方案一:基于分组键的dense_rank实现
通过自定义分组键,结合dense_rank生成符合规则的group_id:
result <- sdf %>% mutate( # 非y组用sector作为分组键,y组用firm作为分组键(保证每个y企业单独成组) group_key = ifelse(sector == "y", firm, sector), # 设定排序优先级:非y组在前,保证其group_id先分配 priority = ifelse(sector == "y", 2, 1) ) %>% # 按优先级和分组键排序,确保rank顺序符合预期 arrange(priority, group_key) %>% # 基于优先级+分组键的组合生成唯一group_id mutate(group = dense_rank(concat(priority, group_key))) %>% # 移除辅助列 select(-group_key, -priority) %>% # 按原firm顺序排列(可选) arrange(firm) # 查看结果 result
方案二:拆分处理后合并(适合数据量较小场景)
先给非y的sector分配group_id,再给y的企业分配独立id,最后合并结果:
# 1. 给非y的sector分配唯一group_id sector_groups <- sdf %>% filter(sector != "y") %>% distinct(sector) %>% arrange(sector) %>% mutate(group = row_number()) # 2. 给y的企业分配从非y组最大id开始的连续id max_sector_group <- sector_groups %>% summarise(max_group = max(group)) %>% pull(max_group) y_groups <- sdf %>% filter(sector == "y") %>% arrange(firm) %>% mutate(group = max_sector_group + row_number()) # 3. 合并两部分结果 result <- sdf %>% filter(sector != "y") %>% left_join(sector_groups, by = "sector") %>% bind_rows(y_groups) %>% arrange(firm) # 查看结果 result
说明
之前尝试的方法报错原因:
group_indices、cur_group_id在Sparklyr的分布式环境中支持有限,部分版本无法直接使用seq_along是本地R函数,无法在Spark集群的分布式数据上执行
内容的提问来源于stack exchange,提问作者Adriana LE
相关产品推荐
相关产品推荐

