使用R的dplyr生成指定列最大值列名新列的问题排查
问题:获取指定列每行最大值对应列名的错误分析
问题背景
现有如下数据框,目标是新增一列var.max,记录每行中var2、var3、var4三列里数值最大的列名:
dat <- data.frame(var1 = rnorm(10), var2 = rnorm(10), var3 = rnorm(10), var4 = rnorm(10)) dat
输出:
var1 var2 var3 var4 1 -1.3784414 1.06816022 1.46578217 -0.4141153 2 -0.3272332 -0.69470574 0.02220395 -0.5502878 3 0.2559891 -0.06964848 -0.34745180 0.6399705 4 0.6029044 1.23680560 -0.72392358 -0.1990832 5 1.3097174 -0.58028595 -0.01487186 -0.8765290 6 -1.2356668 0.41330063 -1.00375989 -1.1974204 7 -0.4126320 3.83320678 -1.42059022 -0.6747575 8 1.7339653 0.58610348 0.40200428 1.4582103 9 1.2994859 1.65355306 0.75985071 0.6455882 10 -0.2353356 2.04468739 -0.11521602 0.3251901
错误代码及结果
使用以下代码无法得到正确结果:
library(dplyr) dat %>% rowwise() %>% mutate(var.max = colnames(.)[which.max(c_across(var2:var4))])
输出:
# A tibble: 10 x 5 # Rowwise: var1 var2 var3 var4 var.max <dbl> <dbl> <dbl> <dbl> <chr> 1 -1.38 1.07 1.47 -0.414 var2 2 -0.327 -0.695 0.0222 -0.550 var2 3 0.256 -0.0696 -0.347 0.640 var3 4 0.603 1.24 -0.724 -0.199 var1 5 1.31 -0.580 -0.0149 -0.877 var2 6 -1.24 0.413 -1.00 -1.20 var1 7 -0.413 3.83 -1.42 -0.675 var1 8 1.73 0.586 0.402 1.46 var3 9 1.30 1.65 0.760 0.646 var1 10 -0.235 2.04 -0.115 0.325 var1
可行的两种情况
情况1:移除var1后代码正常运行
dat %>% select(-var1) %>% rowwise() %>% mutate(var.max = colnames(.)[which.max(c_across(var2:var4))])
输出:
# A tibble: 10 x 4 # Rowwise: var2 var3 var4 var.max <dbl> <dbl> <dbl> <chr> 1 1.07 1.47 -0.414 var3 2 -0.695 0.0222 -0.550 var3 3 -0.0696 -0.347 0.640 var4 4 1.24 -0.724 -0.199 var2 5 -0.580 -0.0149 -0.877 var3 6 0.413 -1.00 -1.20 var2 7 3.83 -1.42 -0.675 var2 8 0.586 0.402 1.46 var4 9 1.65 0.760 0.646 var2 10 2.04 -0.115 0.325 var2
情况2:将var1移到最后一列后代码正常运行
dat %>% select(var2, var3, var4, var1) %>% rowwise() %>% mutate(var.max = colnames(.)[which.max(c_across(var2:var4))])
输出:
# A tibble: 10 x 5 # Rowwise: var2 var3 var4 var1 var.max <dbl> <dbl> <dbl> <dbl> <chr> 1 1.07 1.47 -0.414 -1.38 var3 2 -0.695 0.0222 -0.550 -0.327 var3 3 -0.0696 -0.347 0.640 0.256 var4 4 1.24 -0.724 -0.199 0.603 var2 5 -0.580 -0.0149 -0.877 1.31 var3 6 0.413 -1.00 -1.20 -1.24 var2 7 3.83 -1.42 -0.675 -0.413 var2 8 0.586 0.402 1.46 1.73 var4 9 1.65 0.760 0.646 1.30 var2 10 2.04 -0.115 0.325 -0.235 var2
问题根源
错误的核心是索引错位:
c_across(var2:var4)返回当前行这三列的数值向量,它的索引是1(对应var2)、2(对应var3)、3(对应var4)。colnames(.)取的是整个数据框的列名,原数据框列顺序为var1, var2, var3, var4,所以colnames(.)[1]是var1,colnames(.)[2]是var2,以此类推。- 当
which.max返回1时,代码会取colnames(.)[1]即var1,而非我们需要的var2,导致列名匹配错误。
当移除var1或把var1移到最后时,数据框列顺序变为var2, var3, var4(或var2, var3, var4, var1),此时colnames(.)[1]对应var2,索引匹配正确,结果自然正常。
解决方法
提供几种简洁的修正方案:
方案1:直接指定目标列名向量
dat %>% rowwise() %>% mutate(var.max = c("var2", "var3", "var4")[which.max(c_across(var2:var4))])
方案2:从c_across结果中取列名
dat %>% rowwise() %>% mutate(var.max = names(c_across(var2:var4))[which.max(c_across(var2:var4))])
方案3:使用pick()更清晰地选择列
dat %>% rowwise() %>% mutate(var.max = names(pick(var2:var4))[which.max(pick(var2:var4))])
方案4:向量化方法(无需rowwise,效率更高)
dat %>% mutate(var.max = colnames(select(., var2:var4))[max.col(select(., var2:var4))])
内容的提问来源于stack exchange,提问作者erc
相关产品推荐
相关产品推荐

