You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在dplyr中实现带条件的递归计算生成mv列?

问题描述

首先给出可复现的测试数据框:

dat <- structure(list(A = c(1.3, 1.5, 1.6, 1.2, 1.1, 1.2), 
                     B = c(0.25, 0.21, 0.21, 0.15, 0.26, 0.17),
                     sig = c(1, 0, 1, 1, 1, 1 ),
                     coef = c(1.25, 2.5, 3.3, 1.8, 2.25, 4.5)), 
                class = c("tbl_df", "tbl", "data.frame"), 
                row.names = c(NA, -6L))

需要创建新列mv,计算规则为递归依赖上一行的mv值,初始值设为1000:

  • 当A ≥ 1.2且sig = 1:
    • 第一行:mv = (1 - B)*1000 + B*1000*coef
    • 后续行:将公式中的1000替换为上一行的mv值
  • 当A ≥ 1.2且sig = 0:mv = 上一行mv值 - 上一行mv值*B
  • 当A < 1.2:mv保持上一行的值不变

期望输出:

ABsigcoefmv
1.30.2511.251062.5
1.50.2102.5839.4
1.60.2113.31244.79
1.20.1511.81394.17
1.10.2612.251394.17
1.20.1714.52223.70

原尝试代码未得到正确结果:

dat <- dat %>%
  mutate(mv = case_when(
    sig==1 ~ accumulate(
    B *(A>=1.2) * coef, .f = ~ .x * (1 + .y), .init = 1000)[-1],
    sig== 0 ~ accumulate(
      B *(A>=1.2), .f = ~ .x - (1 * .y), .init = 1000)[-1]))

解决方案

原代码的问题在于case_when无法结合accumulate实现动态递归逻辑——accumulate会对整列统一计算,无法根据每行条件调整规则。正确做法是用accumulate搭配自定义函数,逐行判断条件并计算。

代码如下:

library(dplyr)
library(purrr)

dat <- dat %>%
  mutate(mv = accumulate(
    1:nrow(.),
    .init = 1000,
    .f = function(prev, i) {
      current_A <- A[i]
      current_B <- B[i]
      current_sig <- sig[i]
      current_coef <- coef[i]
      
      if (current_A >= 1.2 && current_sig == 1) {
        prev * (1 - current_B + current_B * current_coef)
      } else if (current_A >= 1.2 && current_sig == 0) {
        prev * (1 - current_B)
      } else {
        prev
      }
    }
  )[-1])

运行后输出结果与预期一致:

print(dat)
#> # A tibble: 6 × 5
#>       A     B   sig  coef    mv
#>   <dbl> <dbl> <dbl> <dbl> <dbl>
#> 1   1.3  0.25     1  1.25 1062.
#> 2   1.5  0.21     0  2.5   839.
#> 3   1.6  0.21     1  3.3  1245.
#> 4   1.2  0.15     1  1.8  1394.
#> 5   1.1  0.26     1  2.25 1394.
#> 6   1.2  0.17     1  4.5  2224.

代码说明

  • 用accumulate遍历每行索引,初始值设为1000
  • 自定义函数中,根据当前行的A、sig值选择计算规则:
    • 满足A≥1.2且sig=1时,简化公式为prev*(1 - B + B*coef),避免重复计算prev
    • 满足A≥1.2且sig=0时,简化为prev*(1-B)
    • 其他情况直接返回上一行的mv值
  • 最后移除accumulate返回的初始值([-1]),得到与原数据行数匹配的mv列

内容的提问来源于stack exchange,提问作者M.O

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.14 11:22:04