在R中如何用向量化方法实现整数(如年龄)的区间分类?
问题描述
需要写一个函数,接收两个参数:
- 参数1:整数向量(比如年龄数据)
- 参数2:形如
"上限-下限"的区间字符串向量(例如"1-2")
要给每个输入的整数返回对应的区间组。已经用嵌套循环写了个classifyAge函数,测试有效,但想改成向量化方法提速,试了cut没成功,求解决办法。
原循环代码和测试结果:
classifyAge <- function(ages, intervals) { result <- character(length(ages)) for (i in seq_along(ages)) { for (j in seq_along(intervals)) { range <- as.numeric(strsplit(intervals[j], "-")[[1]]) if (ages[i] >= range[1] & ages[i] <= range[2]) { result[i] <- intervals[j] break } } } return(result) } result <- classifyAge(c(1, 2, 3, 5, 5, 7,0), c("1-2", "3-4", "5-Inf")) print(result) # [1] "1-2" "1-2" "3-4" "5-Inf" "5-Inf" "5-Inf" ""
方案1:用outer实现向量化匹配
先把区间字符串拆成数值上下限,处理好"Inf"转成R自带的Inf,然后用outer一次性生成所有年龄和区间的匹配矩阵,再找每个年龄第一个匹配的区间就行,完全不用循环。
代码:
classifyAge_vec1 <- function(ages, intervals) { # 解析区间:拆分上下限,把字符Inf换成R的Inf值 interval_ranges <- strsplit(intervals, "-") |> lapply(\(x) { vals <- as.numeric(x) vals[is.na(vals)] <- Inf # 处理"Inf"转成numeric后变成NA的情况 vals }) |> do.call(rbind, args = _) # 生成匹配矩阵:每个年龄是否在对应区间里 match_matrix <- outer(ages, seq_len(nrow(interval_ranges)), \(a, idx) { a >= interval_ranges[idx, 1] & a <= interval_ranges[idx, 2] }) # 找每个年龄第一个匹配的区间索引,没匹配到就是NA match_idx <- apply(match_matrix, 1, \(x) which(x)[1]) # 映射回区间字符串,没匹配的就为空 result <- ifelse(is.na(match_idx), "", intervals[match_idx]) result } # 测试 result1 <- classifyAge_vec1(c(1, 2, 3, 5, 5, 7,0), c("1-2", "3-4", "5-Inf")) print(result1) # [1] "1-2" "1-2" "3-4" "5-Inf" "5-Inf" "5-Inf" ""
方案2:用findInterval高效匹配
findInterval是R专门用来做区间匹配的函数,效率比循环高一大截,适合区间有序的场景。
代码:
classifyAge_vec2 <- function(ages, intervals) { # 解析区间,处理Inf interval_ranges <- strsplit(intervals, "-") |> lapply(\(x) { vals <- as.numeric(x) vals[is.na(vals)] <- Inf vals }) |> do.call(rbind, args = _) # 提取区间左端点,先确保区间是按左端点递增的(原问题区间是有序的,无序的话先排序) breaks <- interval_ranges[, 1] if (!is.unsorted(breaks)) { sort_idx <- order(breaks) breaks <- breaks[sort_idx] interval_ranges <- interval_ranges[sort_idx, ] intervals <- intervals[sort_idx] } # 用findInterval定位每个年龄对应的区间索引 idx <- findInterval(ages, breaks, rightmost.closed = TRUE) # 验证是否在区间右端点内(findInterval只看左端点) valid <- ages <= interval_ranges[idx, 2] # 没匹配到的情况(idx为0或者不满足右端点)设为空字符串 result <- ifelse(idx == 0 | !valid, "", intervals[idx]) result } # 测试 result2 <- classifyAge_vec2(c(1, 2, 3, 5, 5, 7,0), c("1-2", "3-4", "5-Inf")) print(result2) # [1] "1-2" "1-2" "3-4" "5-Inf" "5-Inf" "5-Inf" ""
方案3:修正cut函数的用法
之前用cut失败大概率是没处理好区间格式和匹配逻辑,cut需要指定断点和标签,调整参数就能适配需求:
代码:
classifyAge_vec3 <- function(ages, intervals) { # 解析区间,处理Inf interval_ranges <- strsplit(intervals, "-") |> lapply(\(x) { vals <- as.numeric(x) vals[is.na(vals)] <- Inf vals }) |> do.call(rbind, args = _) # 生成cut需要的断点:所有左端点 + 最后一个区间的右端点 breaks <- c(interval_ranges[, 1], interval_ranges[nrow(interval_ranges), 2]) breaks <- unique(sort(breaks)) # 去重并排序 # 用cut匹配区间,设置include.lowest=TRUE确保左闭右闭,和原函数逻辑一致 result <- cut(ages, breaks = breaks, labels = intervals, include.lowest = TRUE, right = TRUE) # 把NA(不在任何区间的数值)转成空字符串 result <- as.character(result) result[is.na(result)] <- "" result } # 测试 result3 <- classifyAge_vec3(c(1, 2, 3, 5, 5, 7,0), c("1-2", "3-4", "5-Inf")) print(result3) # [1] "1-2" "1-2" "3-4" "5-Inf" "5-Inf" "5-Inf" ""
内容的提问来源于stack exchange,提问作者Aku-Ville Lehtimäki
相关产品推荐
相关产品推荐

