代码之家  ›  专栏  ›  技术社区  ›  Xavier

变长矢量的矢量化和

  •  1
  • Xavier  · 技术社区  · 2 年前

    我试图通过以下方程模拟y的观测结果:

    对于X的所有可能值和参数贝塔的随机值以及参数德尔塔的矢量。

    使用来自tidyverse的R和包,这是我迄今为止得到的:

    library(tidyverse)
    library(extraDistr)
    n <- 100
    x_max <- 4
    
    set.seed(1)
    data <- 
        tibble(
            id = 1:100,
            beta = rnorm(100, mean = 0, sd = 1/(x_max - 1)),
            delta = cbind(0, rdirichlet(n, rep(1, x_max - 1)))
        ) %>% 
        expand_grid(x = 1:x_max)
    
    data
    
    # A tibble: 400 × 4
          id    beta delta[,1]   [,2]    [,3]   [,4]     x
       <int>   <dbl>     <dbl>  <dbl>   <dbl>  <dbl> <int>
     1     1 -0.209          0 0.169  0.170   0.661      1
     2     1 -0.209          0 0.169  0.170   0.661      2
     3     1 -0.209          0 0.169  0.170   0.661      3
     4     1 -0.209          0 0.169  0.170   0.661      4
     5     2  0.0612         0 0.0962 0.00297 0.901      1
     6     2  0.0612         0 0.0962 0.00297 0.901      2
     7     2  0.0612         0 0.0962 0.00297 0.901      3
     8     2  0.0612         0 0.0962 0.00297 0.901      4
     9     3 -0.279          0 0.241  0.714   0.0451     1
    10     3 -0.279          0 0.241  0.714   0.0451     2
    # ℹ 390 more rows
    # ℹ Use `print(n = ...)` to see more rows
    

    我被困在从上面的公式计算y值的部分。挑战在于在计算增量参数的总和时使用x作为索引变量。我要找的是这样的东西:

    data %>% 
        rowwise() %>% 
        mutate(
            y = beta * sum(c_across(delta[,1]:delta[,x]))
        )
    

    但是,此代码不起作用,我收到以下错误消息:

    Error in `mutate()`:
    ℹ In argument: `y = beta * sum(c_across(delta[, 1]:delta[, x]))`.
    ℹ In row 1.
    Caused by error in `c_across()`:
    ! Problem while evaluating `delta[, 1]`.
    Caused by error:
    ! object 'delta' not found
    

    我知道嵌套矩阵一定有问题 delta 。然而,我仍然觉得 dplyr 应该能够支持矢量化索引操作吗?

    从上面的例子中,我希望得到的结果是:

    # A tibble: 400 × 5
    # Rowwise: 
          id    beta delta[,1]   [,2]    [,3]   [,4]     x        y
       <int>   <dbl>     <dbl>  <dbl>   <dbl>  <dbl> <int>    <dbl>
     1     1 -0.209          0 0.169  0.170   0.661      1  0      
     2     1 -0.209          0 0.169  0.170   0.661      2 -0.0352 
     3     1 -0.209          0 0.169  0.170   0.661      3 -0.0708 
     4     1 -0.209          0 0.169  0.170   0.661      4 -0.209  
     5     2  0.0612         0 0.0962 0.00297 0.901      1  0      
     6     2  0.0612         0 0.0962 0.00297 0.901      2  0.00589
     7     2  0.0612         0 0.0962 0.00297 0.901      3  0.00607
     8     2  0.0612         0 0.0962 0.00297 0.901      4  0.0612 
     9     3 -0.279          0 0.241  0.714   0.0451     1  0      
    10     3 -0.279          0 0.241  0.714   0.0451     2 -0.0671 
    # ℹ 390 more rows
    # ℹ Use `print(n = ...)` to see more rows
    
    1 回复  |  直到 2 年前
        1
  •  1
  •   Andy Baxter    2 年前

    这接近你想要的吗?

    library(tidyverse)
    library(extraDistr)
    
    set.seed(1)
    
    n <- 100
    x_max <- 4
    
    data <- 
      tibble(
        id = 1:n,
        beta = rnorm(n, mean = 0, sd = 1/(x_max - 1)),
        delta = cbind(0, rdirichlet(n, rep(1, x_max - 1)))
      ) %>% 
      expand_grid(x = 1:x_max)
    
    data %>% 
      rowwise() %>% 
      mutate(
        y = beta * sum(delta[,1:x])
      )
    #> # A tibble: 400 × 5
    #> # Rowwise: 
    #>       id    beta delta[,1]   [,2]    [,3]   [,4]     x        y
    #>    <int>   <dbl>     <dbl>  <dbl>   <dbl>  <dbl> <int>    <dbl>
    #>  1     1 -0.209          0 0.169  0.170   0.661      1  0      
    #>  2     1 -0.209          0 0.169  0.170   0.661      2 -0.0352 
    #>  3     1 -0.209          0 0.169  0.170   0.661      3 -0.0708 
    #>  4     1 -0.209          0 0.169  0.170   0.661      4 -0.209  
    #>  5     2  0.0612         0 0.0962 0.00297 0.901      1  0      
    #>  6     2  0.0612         0 0.0962 0.00297 0.901      2  0.00589
    #>  7     2  0.0612         0 0.0962 0.00297 0.901      3  0.00607
    #>  8     2  0.0612         0 0.0962 0.00297 0.901      4  0.0612 
    #>  9     3 -0.279          0 0.241  0.714   0.0451     1  0      
    #> 10     3 -0.279          0 0.241  0.714   0.0451     2 -0.0671 
    #> # ℹ 390 more rows
    

    我认为混乱 c_across 是的 delta 从技术上讲,表现为一列(矩阵),因此不能“跨”求和:

    str(data)
    #> tibble [400 × 4] (S3: tbl_df/tbl/data.frame)
    #>  $ id   : int [1:400] 1 1 1 1 2 2 2 2 3 3 ...
    #>  $ beta : num [1:400] -0.2088 -0.2088 -0.2088 -0.2088 0.0612 ...
    #>  $ delta: num [1:400, 1:4] 0 0 0 0 0 0 0 0 0 0 ...
    #>  $ x    : int [1:400] 1 2 3 4 1 2 3 4 1 2 ...