在并行环境中运行purrr的函数式工具:imap
在一个更大的脚本里,我有一个函数用来做一些聚类重采样,并且它使用 imap 来运行,如下方代码所示:
library(tibble)
library(purrr)
library(dplyr)
n_ids <- 500
# Total number of observations
n_obs <- 2000
set.seed(2025)
# Create vector of individual IDs with repeated observations
ids <- sample(1:n_ids, size = n_obs, replace = TRUE)
# Generate binomial variables (0/1)
outcome <- rbinom(n_obs, size = 1, prob = 0.4)
predictor <- rnorm(n_obs, mean = 35.5, sd = 2.5)
# Build the dataframe
df <- tibble(
id = ids,
outcome = outcome,
predictor = round(predictor, digits = 1)
)
ids <- unique(df$id)
sampled_ids <- sample(ids, length(ids), replace = TRUE)
# clustered resampling
resamp_df <- function(ids, i, data){
df1 <- data[data$id == ids, ]
df1$id2 <- rep(i)
return(df1)
}
d1 <- imap(sampled_ids, resamp_df, df)
d1 <- bind_rows(d1)
我想能够使用 purrr 的 in parallel 选项来运行这段代码。我查看了文档,文档涵盖了 map 的用法,虽然我原以为理解了如何为并行执行传递参数,但似乎并非如此。我知道 future_map,但出于特定原因,我想继续使用相同的功能。
解决方案
分为两个部分:
- 使用
mirai实现in_parallel(..)的基础用法 - 将本地对象传递给其他进程。
基本的 mirai
in_parallel() 使用一组 mirai 节点的配置。
# simple, single-process
imap(5:8, ~ runif(.y, max = .x))
# potentially parallelized
imap(5:8, in_parallel(~ runif(.y, max = .x)))
我之所以说“潜在地”,是因为无论是否存在 mirai 的进程设置,它都能工作。在这种情况下,仍然是同一个进程,因此这里的 in_parallel(.) 是顺序执行,并没有真正的并行化。而且它既不发出警告、也不抱怨,甚至也不对这一点发表评论,因此看起来没有明显的改进。
我们可以用下面这个更直观地演示:
map(1:3, ~ Sys.getpid())
# [[1]]
# [1] 48159
# [[2]]
# [1] 48159
# [[3]]
# [1] 48159
map(1:3, in_parallel(~ Sys.getpid()))
# [[1]]
# [1] 48159
# [[2]]
# [1] 48159
# [[3]]
# [1] 48159
没有差异。然而,如果我们设置3 个进程,我们可以看到它们之间的进程号可能不同。
mirai::daemons(3)
map(1:5, ~ Sys.getpid())
# [[1]]
# [1] 48159
# [[2]]
# [1] 48159
# [[3]]
# [1] 48159
# [[4]]
# [1] 48159
# [[5]]
# [1] 48159
map(1:5, in_parallel(~ Sys.getpid()))
# [[1]]
# [1] 71513
# [[2]]
# [1] 71520
# [[3]]
# [1] 71533
# [[4]]
# [1] 71520
# [[5]]
# [1] 71513
一些有趣的点需要了解:
### number of cores your particular computer has;
### can use `logical=FALSE` to differentiate between
### logical and physical cores, varies by chip
parallel::detectCores()
# [1] 16
### current state of your mirai persistent processes
mirai::status()
# $connections
# [1] 3
# $daemons
# [1] "ipc:///tmp/5e75569e6c116e9605b79e24"
# $mirai
# awaiting executing completed
# 0 0 5
一个计时的示例,也许更直观:
mirai::status() # default when nothing previously setup
# $connections
# [1] 0
# $daemons
# [1] 0
system.time(map(1:3, ~ Sys.sleep(1)))
# user system elapsed
# 0.003 0.001 3.015
system.time(map(1:3, in_parallel(~ Sys.sleep(1))))
# user system elapsed
# 0.003 0.001 3.016
mirai::daemons(3)
system.time(map(1:3, ~ Sys.sleep(1)))
# user system elapsed
# 0.000 0.000 3.014
system.time(map(1:3, in_parallel(~ Sys.sleep(1))))
# user system elapsed
# 0.001 0.000 1.008
本地对象
如果直接把这应用到你的示例,你会发现另一个问题:
imap(sampled_ids, in_parallel(~ resamp_df(.x, .y, df)))
# Error in `map2()`:
# ℹ In index: 1.
# Caused by error in `resamp_df()`:
# ! could not find function "resamp_df"
# Run `rlang::last_trace()` to see where the error occurred.
该函数已被明确定义,但所有的 mirai 工作进程都会以干净/空的R 环境启动,因此你定义的任何内容(即 resamp_df 和 df 本身)都不可用。我们可以通过多种方式传递它:
- 全局地,你可以使用
mirai::everywhere将对象传递给所有对象的全局环境:
r
mirai::everywhere({}, resamp_df=resamp_df, df=df)
imap(sampled_ids, in_parallel(~ resamp_df(.x, .y, df))) |> bind_rows()
第一个参数必须是某种形式的表达式,所有 ... 参数都是要存储在所有工作进程全局环境中的对象的命名参数。通常会把它们命名为与在该环境中的对象同名,但这并非必须。即使我们没有使用它,第一个参数也不是可选的,因此我们传递一个空的 {} 作为占位;也可以很容易地传递 pi、c 或其他无害的东西。
作为更多演示(也是为第二种选项做准备),我将整理工作进程的环境:
r
mirai::everywhere(ls(envir = .GlobalEnv))[]
# [[1]]
# [1] "df" "resamp_df"
# [[2]]
# [1] "df" "resamp_df"
mirai::everywhere(rm(resamp_df, df, envir=.GlobalEnv))[]
# [[1]]
# NULL
# [[2]]
# NULL
mirai::everywhere(ls(envir = .GlobalEnv))[]
# [[1]]
# character(0)
# [[2]]
# character(0)
2. ... 的参数 in_parallel 的工作方式与 everywhere 相似,但有一个小差别:发送给mirai的每个表达式在非全局环境中执行。虽然 .GlobalEnv 在其搜索路径中可见,但在一个临时/瞬态环境中工作意味着清理是自然的,不会有“过期对象”之类的意外。因为这一点,当表达式完成后,表达式结束时,我们将不再看到 resamp_df 和 df。
``r
imap(sampled_ids, in_parallel(~ resamp_df(.x, .y, data), resamp_df = resamp_df, data = df)) |>
bind_rows()
# # A tibble: 2,059 × 4
# id outcome predictor id2
# <int> <int> <dbl> <int>
# 1 406 1 35.3 1
# 2 406 0 38.2 1
# 3 406 1 32.4 1
# 4 406 0 34 1
# 5 406 1 36.2 1
# 6 406 0 34 1
# 7 42 0 36.3 2
# 8 42 1 39.3 2
# 9 8 1 28.3 3
# 10 8 1 36.3 3
# # ℹ 2,049 more rows
# # ℹ Useprint(n = ...)` to see more rows
mirai::everywhere(ls(envir = .GlobalEnv))[] # [[1]] # character(0) # [[2]] # character(0) ```
(请注意,我为该调用命名了 data = df,这纯粹是为了演示在调用中重命名。通常,大多数人很可能会写成 df=df,并将其用作 df。)