【问题标题】:Random sample of character vector, without elements prefixing one another字符向量的随机样本,没有元素彼此前缀
【发布时间】:2015-06-10 23:42:17
【问题描述】:

考虑一个字符向量pool,其元素是(填充零的)二进制数,最多有max_len 个数字。

max_len <- 4
pool <- unlist(lapply(seq_len(max_len), function(x) 
  do.call(paste0, expand.grid(rep(list(c('0', '1')), x)))))

pool
##  [1] "0"    "1"    "00"   "10"   "01"   "11"   "000"  "100"  "010"  "110" 
## [11] "001"  "101"  "011"  "111"  "0000" "1000" "0100" "1100" "0010" "1010"
## [21] "0110" "1110" "0001" "1001" "0101" "1101" "0011" "1011" "0111" "1111"

我想对这些元素中的n 进行抽样,但条件是所有抽样元素都不是任何其他抽样元素的prefixes(即,如果我们抽样1101,我们将禁止111110,而如果我们对 1 进行采样,我们会禁止那些以 1 开头的元素,例如 1011100 等)。

以下是我使用while 的尝试,但是当n 很大时(或接近2^max_len)当然会很慢。

set.seed(1)
n <- 10
chosen <- sample(pool, n)
while(any(rowSums(outer(paste0('^', chosen), chosen, Vectorize(grepl))) > 1)) {
  prefixes <- rowSums(outer(paste0('^', chosen), chosen, Vectorize(grepl))) > 1
  pool <- pool[rowSums(Vectorize(grepl, 'pattern')(
    paste0('^', chosen[!prefixes]), pool)) == 0]
  chosen <- c(chosen[!prefixes], sample(pool, sum(prefixes)))
}

chosen
## [1] "0100" "0101" "0001" "0011" "1000" "111"  "0000" "0110" "1100" "0111"

这可以通过最初从pool 中删除那些包含意味着pool 中剩余的元素不足以获取大小为n 的总样本的元素来稍微改进。例如,当max_len = 4n &gt; 9 时,我们可以立即从pool 中删除01,因为包含其中任何一个,最大样本将为9(0 和八个4 字符元素以11 开头以及以0 开头的八个4 字符元素)。

基于此逻辑,我们可以在获取初始样本之前省略 pool 中的元素,例如:

pool <- pool[
  nchar(pool) > tail(which(n > (2^max_len - rev(2^(0:max_len))[-1] + 1)), 1)]

谁能想到更好的方法?我觉得我忽略了一些更简单的东西。


编辑

为了阐明我的意图,我将池描绘成一组分支,其中的交汇点和尖端是节点(pool 的元素)。假设绘制了下图中的黄色节点(即010)。现在,由节点 0、01 和 010 组成的整个红色“分支”从池中移除。这就是我的意思是禁止对我们样本中已经存在“前缀”节点的节点进行采样(以及那些已经在我们的样本中的节点前缀的节点)。

如果采样的节点在分支的中途,比如下图中的01,那么所有的红色节点(0、01、010、011)都是不允许的,因为0前缀01,01前缀都是010和 011。

我的意思不是要在每个交界处采样 1 0(即沿着树枝走在叉子上翻转硬币) - 两者都在采样,只要: (1) 节点的父母(或祖父母等)或孩子(孙子女等)尚未采样; (2) 在对节点进行采样时,将有足够的节点剩余来实现所需的大小为n 的样本。

在上面的第二张图中,如果 010 是第一个选择,那么黑色节点上的所有节点仍然(当前)有效,假设 n &lt;= 4。例如,如果n==4 并且我们接下来对节点 1 进行采样(因此我们的选择现在包括 01 和 1),我们随后将禁止节点 00(由于上面的规则 2)但仍然可以选择 000 和 001,给我们我们的 4 -元素样本。另一方面,如果n==5,则在此阶段将不允许节点 1。

【问题讨论】:

  • 我对如何改进您的特定算法没有任何见解,但我们知道 R 的字符串操作速度像糖蜜一样慢。也许基于数值向量比较的解决方案会更快。像检查(a,b)一样,将 a 分解为所有可能长度的向量(a'),然后对所有对进行检查 a' == b。
  • 让我知道我的答案是否正确,如果是,我会清理/评论更多。具有 4 个样本的 max_len == 9 运行时间为 8 毫秒。
  • @jbaums 你的更新看起来很像我在链接中指向的霍夫曼树(例如here
  • @jbaums 我是该主题专家的 faaar(那么我会建议一个答案;))。我刚刚发现您的问题与“无前缀代码/霍夫曼代码”主题的相似之处令人震惊。在我天真的眼里,它似乎是解决您的问题的一种富有成效的方法,而不是最不重要的启发式方法(例如,正如您在编辑中添加的树所建议的那样)。我不知道您想要的采样的适当、更具体的术语。我不知道“那种方法”是否给定一个池,而不是有一种算法可以最小化结果的熵(位/符号)。
  • 对于任何想知道我的赏金在这里发生了什么的人......发布现有答案的人都投入了大量时间和精力来制作高质量的答案。我从这些中学到了很多,并计划奖励一堆赏金。奖励赏金有 24 小时的延迟,所以这需要一段时间。也就是说,如果有人想破解它,也欢迎新的答案;)

标签: r performance combinatorics


【解决方案1】:

简介

这是我们在另一个答案中实现的字符串算法的数字变体。它速度更快,并且不需要创建或排序池。

算法大纲

我们可以使用整数来表示你的二进制字符串,这大大简化了池生成和值顺序消除的问题。例如,对于max_len==3,我们可以将数字1--(其中- 表示填充)表示十进制的4。此外,如果我们选择这个数字,我们可以确定需要消除的数字是44 + 2 ^ x - 1 之间的数字。这里x是填充元素的个数(本例中为2个),所以要消除的数字在44 + 2 ^ 2 - 1之间(或者在47之间,表示为100,@987654337 @ 和 111)。

为了准确地匹配您的问题,我们需要处理一些问题,因为您将二进制中可能相同的数字视为算法某些部分的不同数字。例如,10010-1-- 都是相同的数字,但在您的方案中需要区别对待。在max_len==3 世界中,我们有 8 个可能的数字,但有 14 种可能的表示:

0 - 000: 0--, 00-
1 - 001:
2 - 010: 01-
3 - 011:
4 - 100: 1--, 10-
5 - 101:
6 - 110: 11-
7 - 111:

所以 0 和 4 有三种可能的编码,2 和 6 有两种,其他的只有一种。我们需要生成一个整数池,以表示具有多个表示的数字的选择概率更高,以及用于跟踪该数字包含多少空白的机制。我们可以通过在数字末尾附加几位来表示我们想要的权重来做到这一点。所以我们的数字变成了(我们在这里使用两位):

jbaum | int | bin | bin.enc | int.enc    
  0-- |   0 | 000 |   00000 |       0
  00- |   0 | 000 |   00001 |       1      
  000 |   0 | 000 |   00010 |       2      
  001 |   1 | 001 |   00100 |       3      
  01- |   2 | 010 |   01000 |       4  
  010 |   2 | 010 |   01001 |       5  
  011 |   3 | 011 |   01101 |       6  
  1-- |   4 | 100 |   10000 |       7  
  10- |   4 | 100 |   10001 |       8  
  100 |   4 | 100 |   10010 |       9  
  101 |   5 | 101 |   10100 |      10  
  11- |   6 | 110 |   11000 |      11   
  110 |   6 | 110 |   11001 |      12   
  111 |   7 | 111 |   11100 |      13

一些有用的属性:

  • enc.bits 表示我们需要多少位进行编码(本例中为两位)
  • int.enc %% enc.bits 告诉我们明确指定了多少个数字
  • int.enc %/% enc.bits 返回int
  • int * 2 ^ enc.bits + explicitly.specified 返回int.enc

请注意,这里的explicitly.specified 在我们的实现中介于0max_len - 1 之间,因为总是至少指定一位数字。我们现在有一个仅使用整数完全表示您的数据结构的编码。我们可以从整数中采样并使用正确的权重等重现您想要的结果。这种方法的一个限制是我们在 R 中使用 32 位整数,并且我们必须为编码保留一些位,因此我们将自己限制在max_len==25 左右。如果你使用双精度浮点数指定的整数,你可以做得更大,但我们在这里没有这样做。

避免重复选择

有两种粗略的方法可以确保我们不会两次选择相同的值

  1. 跟踪哪些值可供选择,并从中随机抽样
  2. 从所有可能的值中随机抽样,然后检查该值是否之前被抽取过,如果有,再次抽样

虽然第一个选项看起来最干净,但实际上计算成本非常高。它需要对每个选取的所有可能值进行矢量扫描以预先取消选取的值,或者创建一个包含非取消资格值的收缩向量。如果通过 C 代码通过引用来缩小矢量,则缩小选项仅比矢量扫描更有效,但即便如此,它也需要对矢量的潜在大部分进行重复翻译,并且需要 C。

这里我们使用方法#2。这允许我们随机打乱可能值的宇宙一次,然后依次选择每个值,检查它是否没有被取消资格,如果有,则选择另一个,等等。这很有效,因为检查一个由于我们的值编码,值已被选中; 我们可以仅根据值在排序表中推断值的位​​置。因此,我们将每个值的状态记录在排序表中,并且可以通过直接索引访问(无需扫描)更新或查找该状态。

示例

此算法在基础 R 中的实现可通过 a gist 获得。这个特定的实现只拉完整的平局。以下是来自max_len==4 池的 8 个元素的 10 次抽取示例:

# each column represents a draw from a `max_len==4` pool

set.seed(6); replicate(10, sample0110b(4, 8))
     [,1]   [,2]   [,3]   [,4]   [,5]   [,6]   [,7]   [,8]   [,9]   [,10] 
[1,] "1000" "1"    "0011" "0010" "100"  "0011" "0"    "011"  "0100" "1011"
[2,] "111"  "0000" "1101" "0000" "0110" "0100" "1000" "00"   "0101" "1001"
[3,] "0011" "0110" "1001" "0100" "0000" "0101" "1101" "1111" "10"   "1100"
[4,] "0100" "0010" "0000" "0101" "1101" "101"  "1011" "1101" "0110" "1101"
[5,] "101"  "0100" "1100" "1100" "0101" "1001" "1001" "1000" "1111" "1111"
[6,] "110"  "0111" "1011" "111"  "1011" "110"  "1111" "0100" "0011" "000" 
[7,] "0101" "0101" "111"  "011"  "1010" "1000" "1100" "101"  "0001" "0101"
[8,] "011"  "0001" "01"   "1010" "0011" "1110" "1110" "1001" "110"  "1000"

我们最初也有两个实现,它们依赖于方法 #1 来避免重复,一个在基础 R 中,一个在 C 中,但是当 n 为大。这些函数确实实现了绘制不完整的绘图的能力,所以我们在这里提供它们以供参考:

比较基准

以下是一组基准,比较了本 Q/A 中显示的几个函数。以毫秒为单位的时间。 brodie.b 版本是此答案中描述的版本。 brodie 是原始实现,brodie.C 是带有一些 C 的原始实现。所有这些都强制执行完整示例的要求。 brodie.str 是另一个答案中基于字符串的版本。

   size    n  jbaum josilber  frank tensibai brodie.b brodie brodie.C brodie.str
1     4   10     11        1      3        1        1      1        1          0
2     4   50      -        -      -        1        -      -        -          1
3     4  100      -        -      -        1        -      -        -          0
4     4  256      -        -      -        1        -      -        -          1
5     4 1000      -        -      -        1        -      -        -          1
6     8   10      1      290      6        3        2      2        1          1
7     8   50    388        -      8        8        3      4        3          4
8     8  100  2,506        -     13       18        6      7        5          5
9     8  256      -        -     22       27       13     14       12          6
10    8 1000      -        -      -       27        -      -        -          7
11   16   10      -        -    615      688       31     61       19        424
12   16   50      -        -  2,123    2,497       28    276       19      1,764
13   16  100      -        -  4,202    4,807       30    451       23      3,166
14   16  256      -        - 11,822   11,942       40  1,077       43      8,717
15   16 1000      -        - 38,132   44,591       83  3,345      130     27,768

这可以相对较好地扩展到更大的池

system.time(sample0110b(18, 100000))
   user  system elapsed 
  8.441   0.079   8.527 

基准说明:

  • frank 和 brodie(减去 brodie.str)不需要任何预生成池,这会影响比较(见下文)
  • Josilber 是 LP 版本
  • jbaum 是 OP 示例
  • tensibai 稍作修改以在池为空时退出而不是失败
  • 未设置为运行 python,因此无法完全比较/考虑缓冲
  • - 代表不可行的选项或时间太慢而无法合理安排

时间不包括绘制池(0.82.5401 毫秒,大小分别为 4816),这是 jbaum 所必需的, josilberbrodie.str 运行或排序它们(0.12.73700 毫秒对于大小 4816),这是 brodie.str 所必需的除了平局。是否要包含这些取决于您为特定池运行该函数的次数。此外,几乎可以肯定有更好的方法来生成/排序池。

这是使用microbenchmark 的三个运行的中间时间。代码是 available as a gist,但请注意,您必须事先加载 sample0110bsample0110sample01101sample01 函数。

【讨论】:

  • 酷;因为我不知道二进制,所以很多都在我头上。顺便说一句,似乎"0" 有时被排除在外,就像我尝试z = replicate(1e4,sample011(2,4)); table(z) 时一样,我不知道为什么; size=1...
  • 感谢您为此付出的巨大努力(两个答案),Brodie。这是一个聪明的方法!不过,我注意到它正在生成包含相互前缀元素的样本。例如。 set.seed(1); sample011(16, 100)[c(26, 34, 38)] 都共享前缀01111010。还没有仔细检查你的函数,但我会挖掘一下,看看我能不能找出原因。
  • 我认为我们有误会。你写了如果我们选择“010”......那么我们将消除......“010”和“011”,但实际上我的意图是如果我们选择010,那么我们取出 0 和 01,以及任何以 010 开头的节点(例如 0100,尽管由于示例中有 max_len==3,因此这些节点不存在)。我会看看你的方法是否可以适应我的需要。
  • @jbaums,对不起,我想我理解正确,但我在那里留下了一段旧段落。回复:前缀,我意识到昨晚上床睡觉时(添加了一个更新来说明它),但我认为有一种方法可以修复它,我在早期的实现中就有,但由于速度而放弃了。
  • @jbaums,感谢您的赏金。这道题一定是最高悬赏题之一。我用盲采样版本更新了答案。事实证明,在大多数情况下,它甚至比 C 版本更快,并且完全用基础 R 编写。
【解决方案2】:

我发现这个问题很有趣,所以我尝试用非常低的 R 技能来解决这个问题(所以这可能会得到改进):

更新的编辑版本,感谢@Franck 建议:

library(microbenchmark)
library(lineprof)

max_len <- 16
pool <- unlist(lapply(seq_len(max_len), function(x) 
  do.call(paste0, expand.grid(rep(list(c('0', '1')), x)))))
n<-100

library(stringr)
tree_sample <- function(samples,pool) {
  results <- vector("integer",samples)
  # Will be used on a regular basis, compute it in advance
  PoolLen <- str_length(pool)
  # Make a mask vector based on the length of each entry of the pool
  masks <- strtoi(str_pad(str_pad("1",PoolLen,"right","1"),max_len,"right","0"),base=2)

  # Make an integer vector from "0" right padded orignal: for max_len=4 and pool entry "1" we get "1000" => 8
  # This will allow to find this entry as parent of 10 and 11 which become "1000" and "1100", as integer 8 and 12 respectively
  # once bitwise "anded" with the repective mask "1000" the first bit is striclty the same, so it's a parent.
  integerPool <- strtoi(str_pad(pool,max_len,"right","0"),base=2)

  # Create a vector to filter the available value to sample
  ok <- rep(TRUE,length(pool))

  #Precompute the result of the bitwise and betwwen our integer pool and the masks   
  MaskedPool <- bitwAnd(integerPool,masks)

  while(samples) {
    samp <- sample(pool[ok],1) # Get a sample
    results[samples] <- samp # Store it as result
    ok[pool == samp] <- FALSE # Remove it from available entries

    vsamp <- strtoi(str_pad(samp,max_len,"right","0"),base=2) # Get the integer value of the "0" right padded sample
    mlen <- str_length(samp) # Get sample len

    #Creation of unitary mask to remove childs of sample
    mask <- strtoi(paste0(rep(1:0,c(mlen,max_len-mlen)),collapse=""),base=2)

    # Get the result of bitwise And between the integerPool and the sample mask 
    FilterVec <- bitwAnd(integerPool,mask)

    # Get the bitwise and result of the sample and it's mask
    Childm <- bitwAnd(vsamp,mask)

    ok[FilterVec == Childm] <- FALSE  # Remove from available entries the childs of the sample
    ok[MaskedPool == bitwAnd(vsamp,masks)] <- FALSE # compare the sample with all the masks to remove parents matching

    samples <- samples -1
  }
  print(results)
}
microbenchmark(tree_sample(n,pool),times=10L)

主要思想是使用bitmask comparison 来了解一个样本是否是另一个样本的父级(公共位部分),如果是,则从池中抑制该元素。

现在在我的机器上从长度为 16 的池中抽取 100 个样本需要 1.4 秒。

【讨论】:

    【解决方案3】:

    您可以对池进行排序以帮助决定取消哪些元素。例如,查看一个三元素排序池:

     [1] "0"   "00"  "000" "001" "01"  "010" "011" "1"   "10"  "100" "101" "11" 
    [13] "110" "111"
    

    我可以告诉我,我可以取消在我选择的项目之后的任何字符比我的项目更多的任何内容,直到第一个具有相同数量或更少字符的项目。例如,如果我选择“01”,我可以立即看到接下来的两项(“010”、“011”)需要删除,但不是后面的一项,因为“1”的字符较少。之后删除“0”很容易。这是一个实现:

    library(fastmatch)  # could use `match`, but we repeatedly search against same hash
    
    # `pool` must be sorted!
    
    sample01 <- function(pool, n) {
      picked <- logical(length(pool))
      chrs <- nchar(pool)
      pick.list <- character(n)
      pool.seq <- seq_along(pool)
    
      for(i in seq(n)) {
        # Make sure pool not exhausted
    
        left <- which(!picked)
        left.len <- length(left)
        if(!length(left)) break
    
        # Sample from pool
    
        seq.left <- seq.int(left)
        pool.left <- pool[left]
        chrs.left <- chrs[left]
        pick <- sample(length(pool.left), 1L)
    
        # Find all the elements with more characters that are disqualified
        # and store their indices in `valid` (bad name...)
    
        valid.tmp <- chrs.left > chrs.left[[pick]] & seq.left > pick
        first.invalid <- which(!valid.tmp & seq.left > pick)
        valid <- if(length(first.invalid)) {
          pick:(first.invalid[[1L]] - 1L)
        } else pick:left.len
    
        # Translate back to original pool indices since we're working on a 
        # subset in `pool.left`
    
        pool.seq.left <- pool.seq[left]
        pool.idx <- pool.seq.left[valid]
        val <- pool[[pool.idx[[1L]]]]
    
        # Record the picked value, and all the disqualifications
    
        pick.list[[i]] <- val
        picked[pool.idx] <- TRUE
    
        # Disqualify shorter matches
    
        to.rem <- vapply(
          seq.int(nchar(val) - 1), substr, character(1L), x=val, start=1L
        )
        to.rem.idx <- fmatch(to.rem, pool, nomatch=0)
        picked[to.rem.idx] <- TRUE  
      }
      pick.list  
    }
    

    还有一个创建排序池的函数(与您的代码完全一样,但返回排序后):

    make_pool <- function(size)
      sort(
        unlist(
          lapply(
            seq_len(size), 
            function(x) do.call(paste0, expand.grid(rep(list(c('0', '1')), x))) 
      ) ) )
    

    然后,使用max_len 3 池(用于目视检查事物是否符合预期):

    pool3 <- make_pool(3)
    set.seed(1)
    sample01(pool3, 8)
    # [1] "001" "1"   "010" "011" "000" ""    ""    ""   
    sample01(pool3, 8)
    # [1] "110" "111" "011" "10"  "00"  ""    ""    ""   
    sample01(pool3, 8)
    # [1] "000" "01"  "11"  "10"  "001" ""    ""    ""   
    sample01(pool3, 8)
    # [1] "011" "101" "111" "001" "110" "100" "000" "010"    
    

    请注意,在最后一种情况下,我们得到所有 3 位二进制组合 (2 ^ 3),因为偶然我们一直从 3 位数字中采样。此外,只有 3 个大小的池,有许多样本会阻止完整的 8 个平局;您可以通过消除阻止从池中完全抽牌的组合的建议来解决这个问题。

    这非常快。查看 max_len==9 示例,该示例使用替代解决方案耗时 2 秒:

    pool9 <- make_pool(9)
    microbenchmark(sample01(pool9, 4))
    # Unit: microseconds
    #                expr     min      lq  median      uq     max neval
    #  sample01(pool9, 4) 493.107 565.015 571.624 593.791 983.663   100    
    

    大约半毫秒。您也可以合理地尝试相当大的池:

    pool16 <- make_pool(16)  # 131K entries
    system.time(sample01(pool16, 100))
    #  user  system elapsed 
    # 3.407   0.146   3.552 
    

    这不是很快,但我们谈论的是一个包含 130K 项目的池。还有可能进行额外的优化。

    请注意,大型池的排序步骤相对较慢,但我没有计算它,因为您只需要执行一次,您可能会想出一个合理的算法来生成预先排序的池。

    我在一个现已删除的答案中探索了一种更快的整数到二进制方法的可能性,但这需要更多的工作才能准确地与您正在寻找的东西联系起来。

    【讨论】:

    • 你说“你可以用你的建议解决这个问题”,但实际上 OP 的建议(消除特定字符串)是不够的。像josilber's answers之类的东西是必要的,而不是顺序抽样。以 OP 为例,max_len=4n=10,我们知道 01 是不允许的,但我们的前四次抽奖可能是 000111 和 @987654336 @。排除这样的组合将是一个令人难以置信的麻烦。由于 OP 忽略了定义概率,我只选择tail(pool,n)
    • @Frank 如果要对特定节点进行采样,计算仍可用于采样的最大节点数并不难。当然,可能效率低下,但 Rcpp 可能使其成为顺序采样解决方案的一部分。
    • @jbaums,我认为特别是对于大型最大镜头,您将很少遇到不完整的情况,这意味着您可以以低成本重新采样。如果您可以确认此实现(不包括不完整的情况)是否符合您的要求,这将很有用,这样我就可以确保我们在同一页面上。这个应该处理前缀。
    • while 中包装sample01 以确保当n 相对于2^size 较小(例如s &lt;- sample01(pool &lt;- make_pool(6), 10); while(any(s=='')) s &lt;- sample01(pool, 10))时,完整样本工作得很好,但在n 接近时会出现问题or 等于2^size(例如s &lt;- sample01(pool &lt;- make_pool(6), 64); while(any(s=='')) s &lt;- sample01(pool, 64);一点也不奇怪,因为在这种特殊情况下,只有一个可能的完整样本存在)。理想情况下,我希望能够在这两种情况下有效地进行采样。这是一个很好的方法——我可能要求太多了!
    • @jbaums,这是有道理的。另外,仅供参考,我对数字版本有一个可行的解决方案,但仍在清理它。重新先发制人地选择无效的选择可能非常困难。例如,即使n 相对于size 小,选择01 作为前两个是无效的,因此您需要预先计算无效的计算,这可能不可行或只是昂贵。
    【解决方案4】:

    它是在 python 中而不是在 r 中,但 jbaums 说没关系。

    这是我的贡献,请参阅源代码中的 cmets 以了解关键部分。
    我仍在研究分析解决方案以确定深度树tS 样本的可能组合c 的数量,因此我可以改进函数combs。也许有人有它? 这确实是现在的瓶颈。

    在我的笔记本电脑上,从深度为 16 的树中采样 100 个节点大约需要 8 毫秒。 不是第一次,但是由于 combBuffer 被填满了,所以采样越多,速度就会越快。

    import random
    
    
    class Tree(object):
        """
        :param level: The distance of this node from the root.
        :type level: int
        :param parent: This trees parent node
        :type parent: Tree
        :param isleft: Determines if this is a left or a right child node. Can be
                       omitted if this is the root node.
        :type isleft: bool
    
        A binary tree representing possible strings which match r'[01]{1,n}'. Its
        purpose is to be able to sample n of its nodes where none of the sampled
        nodes' ids is a prefix for another one.
        It is possible to change Tree.maxdepth and then reuse the root. All
        children are created ON DEMAND, which means everything is lazily evaluated.
        If the Tree gets too big anyway, you can call 'prune' on any node to delete
        its children.
    
            >>> t = Tree()
            >>> t.sample(8, toString=True, depth=3)
            ['111', '110', '101', '100', '011', '010', '001', '000']
            >>> Tree.maxdepth = 2
            >>> t.sample(4, toString=True)
            ['11', '10', '01', '00']
        """
    
        maxdepth = 10
        _combBuffer = {}
    
        def __init__(self, level=0, parent=None, isleft=None):
            self.parent = parent
            self.level = level
            self.isleft = isleft
            self._left = None
            self._right = None
    
        @classmethod
        def setMaxdepth(cls, depth):
            """
            :param depth: The new depth
            :type depth: int
    
            Sets the maxdepth of the Trees. This basically is the depth of the root
            node.
            """
            if cls.maxdepth == depth:
                return
    
            cls.maxdepth = depth
    
        @property
        def left(self):
            """This tree's left child, 'None' if this is a leave node"""
            if self.depth == 0:
                return None
    
            if self._left is None:
                self._left = Tree(self.level+1, self, True)
            return self._left
    
        @property
        def right(self):
            """This tree's right child, 'None' if this is a leave node"""
            if self.depth == 0:
                return None
    
            if self._right is None:
                self._right = Tree(self.level+1, self, False)
            return self._right
    
        @property
        def depth(self):
            """
            This tree's depth. (maxdepth-level)
            """
            return self.maxdepth-self.level
    
        @property
        def id(self):
            """
            This tree's id, string of '0's and '1's equal to the path from the root
            to this subtree. Where '1' means going left and '0' means going right.
            """
            # level 0 is the root node, it has no id
            if self.level == 0:
                return ''
            # This takes at most Tree.maxdepth recursions. Therefore
            # it is save to do it this way. We could also save each nodes
            # id once it is created to avoid recreating it every time, however
            # this won't save much time but use quite some space.
            return self.parent.id + ('1' if self.isleft else '0')
    
        @property
        def leaves(self):
            """
            The amount of leave nodes, this tree has. (2**depth)
            """
            return 2**self.depth
    
        def __str__(self):
            return self.id
    
        def __len__(self):
            return 2*self.leaves-1
    
        def prune(self):
            """
            Recursively prune this tree's children.
            """
            if self._left is not None:
                self._left.prune()
                self._left.parent = None
                self._left = None
    
            if self._right is not None:
                self._right.prune()
                self._right.parent = None
                self._right = None
    
        def combs(self, n):
            """
            :param n: The amount of samples to be taken from this tree
            :type n: int
            :returns: The amount of possible combinations to choose n samples from
                      this tree
    
            Determines recursively the amount of combinations of n nodes to be
            sampled from this tree.
            Subsequent calls with same n on trees with same depth will return the
            result from the previous computation rather than computing it again.
    
                >>> t = Tree()
                >>> Tree.maxdepth = 4
                >>> t.combs(16)
                1
                >>> Tree.maxdepth = 3
                >>> t.combs(6)
                58
            """
    
            # important for the amount of combinations is only n and the depth of
            # this tree
            key = (self.depth, n)
    
            # We use the dict to save computation time. Calling the function with
            # equal values on equal nodes just returns the alrady computed value if
            # possible.
            if key not in Tree._combBuffer:
                leaves = self.leaves
    
                if n < 0:
                    N = 0
                elif n == 0 or self.depth == 0 or n == leaves:
                    N = 1
                elif n == 1:
                    return (2*leaves-1)
                else:
                    if n > leaves/2:
                        # if n > leaves/2, at least n-leaves/2 have to stay on
                        # either side, otherweise the other one would have to
                        # sample more nodes than possible.
                        nMin = n-leaves/2
                    else:
                        nMin = 0
    
                    # The rest n-2*nMin is the amount of samples that are free to
                    # fall on either side
                    free = n-2*nMin
    
                    N = 0
                    # sum up the combinations of all possible splits
                    for addLeft in range(0, free+1):
                        nLeft = nMin + addLeft
                        nRight = n - nLeft
                        N += self.left.combs(nLeft)*self.right.combs(nRight)
    
                Tree._combBuffer[key] = N
                return N
            return Tree._combBuffer[key]
    
        def sample(self, n, toString=False, depth=None):
            """
            :param n: How may samples to take from this tree
            :type n: int
            :param toString: If 'True' result will direclty be turned into a list
                             of strings
            :type toString: bool
            :param depth: If not None, will overwrite Tree.maxdepth
            :type depth: int or None
            :returns: List of n nodes sampled from this tree
            :throws ValueError: when n is invalid
    
            Takes n random samples from this tree where none of the sample's ids is
            a prefix for another one's.
    
            For an example see Tree's docstring.
            """
            if depth is not None:
                Tree.setMaxdepth(depth)
    
            if toString:
                return [str(e) for e in self.sample(n)]
    
            if n < 0:
                raise ValueError('Negative sample size is not possible!')
    
            if n == 0:
                return []
    
            leaves = self.leaves
            if n > leaves:
                raise ValueError(('Cannot sample {} nodes, with only {} ' +
                                  'leaves!').format(n, leaves))
    
            # Only one sample to choose, that is nice! We are free to take any node
            # from this tree, including this very node itself.
            if n == 1 and self.level > 0:
                # This tree has 2*leaves-1 nodes, therefore
                # the probability that we keep the root node has to be
                # 1/(2*leaves-1) = P_root. Lets create a random number from the
                # interval [0, 2*leaves-1).
                # It will be 0 with probability 1/(2*leaves-1)
                P_root = random.randint(0, len(self)-1)
                if P_root == 0:
                    return [self]
                else:
                    # The probability to land here is 1-P_root
    
                    # A child tree's size is (leaves-1) and since it obeys the same
                    # rule as above, the probability for each of its nodes to
                    # 'survive' is 1/(leaves-1) = P_child.
                    # However all nodes must have equal probability, therefore to
                    # make sure that their probability is also P_root we multiply
                    # them by 1/2*(1-P_root). The latter is already done, the
                    # former will be achieved by the next condition.
                    # If we do everything right, this should hold:
                    # 1/2 * (1-P_root) * P_child = P_root
    
                    # Lets see...
                    # 1/2 * (1-1/(2*leaves-1)) * (1/leaves-1)
                    # (1-1/(2*leaves-1)) * (1/(2*(leaves-1)))
                    # (1-1/(2*leaves-1)) * (1/(2*leaves-2))
                    # (1/(2*leaves-2)) - 1/((2*leaves-2) * (2*leaves-1))
                    # (2*leaves-1)/((2*leaves-2) * (2*leaves-1)) - 1/((2*leaves-2) * (2*leaves-1))
                    # (2*leaves-2)/((2*leaves-2) * (2*leaves-1))
                    # 1/(2*leaves-1)
                    # There we go!
                    if random.random() < 0.5:
                        return self.right.sample(1)
                    else:
                        return self.left.sample(1)
    
            # Now comes the tricky part... n > 1 therefore we are NOT going to
            # sample this node. Its probability to be chosen is 0!
            # It HAS to be 0 since we are definitely sampling from one of its
            # children which means that this node will be blocked by those samples.
            # The difficult part now is to prove that the sampling the way we do it
            # is really random.
    
            if n > leaves/2:
                # if n > leaves/2, at least n-leaves/2 have to stay on either
                # side, otherweise the other one would have to sample more
                # nodes than possible.
                nMin = n-leaves/2
            else:
                nMin = 0
            # The rest n-2*nMin is the amount of samples that are free to fall
            # on either side
            free = n-2*nMin
    
            # Let's have a look at an example, suppose we were to distribute 5
            # samples among two children which have 4 leaves each.
            # Each child has to get at least 1 sample, so the free samples are 3.
            # There are 4 different ways to split the samples among the
            # children (left, right):
            # (1, 4), (2, 3), (3, 2), (4, 1)
            # The amount of unique sample combinations per child are
            # (7, 1), (11, 6), (6, 11), (1, 7)
            # The amount of total unique samples per possible split are
            #   7   ,   66  ,   66  ,    7
            # In case of the first and last split, all samples have a probability
            # of 1/7, this was already proven above.
            # Lets suppose we are good to go and the per sample probabilities for
            # the other two cases are (1/11, 1/6) and (1/6, 1/11), this way the
            # overall per sample probabilities for the splits would be:
            #  1/7  ,  1/66 , 1/66 , 1/7
            # If we used uniform random to determine the split, all splits would be
            # equally probable and therefore be multiplied with the same value (1/4)
            # But this would mean that NOT every sample is equally probable!
            # We need to know in advance how many sample combinations there will be
            # for a given split in order to find out the probability to choose it.
            # In fact, due to the restrictions, this becomes very nasty to
            # determine. So instead of solving it analytically, I do it numerically
            # with the method 'combs'. It gives me the amount of possible sample
            # combinations for a certain amount of samples and a given tree depth.
            # It will return 146 for this node and 7 for the outer and 66 for the
            # inner splits.
            # What we now do is, we take a number from [0, 146).
            # if it is smaller than 7, we sample from the first split,
            # if it is smaller than 7+66, we sample from the second split,
            # ...
            # This way we get the probabilities we need.
    
            r = random.randint(0, self.combs(n)-1)
            p = 0
            for addLeft in xrange(0, free+1):
                nLeft = nMin + addLeft
                nRight = n - nLeft
    
                p += (self.left.combs(nLeft) * self.right.combs(nRight))
                if r < p:
                    return self.left.sample(nLeft) + self.right.sample(nRight)
            assert False, ('Something really strange happend, p did not sum up ' +
                           'to combs or r was too big')
    
    
    def main():
        """
        Do a microbenchmark.
        """
        import timeit
        i = 1
        main.t = Tree()
        template = ' {:>2}  {:>5} {:>4}  {:<5}'
        print(template.format('i', 'depth', 'n', 'time (ms)'))
        N = 100
        for depth in [4, 8, 15, 16, 17, 18]:
            for n in [10, 50, 100, 150]:
                if n > 2**depth:
                    time = '--'
                else:
                    time = timeit.timeit(
                        'main.t.sample({}, depth={})'.format(n, depth), setup=
                        'from __main__ import main', number=N)*1000./N
                print(template.format(i, depth, n, time))
                i += 1
    
    
    if __name__ == "__main__":
        main()
    

    基准输出:

      i  depth    n  time (ms)
      1      4   10  0.182511806488
      2      4   50  --   
      3      4  100  --   
      4      4  150  --   
      5      8   10  0.397620201111
      6      8   50  1.66054964066
      7      8  100  2.90236949921
      8      8  150  3.48146915436
      9     15   10  0.804011821747
     10     15   50  3.7428188324
     11     15  100  7.34910964966
     12     15  150  10.8230614662
     13     16   10  0.804491043091
     14     16   50  3.66818904877
     15     16  100  7.09567070007
     16     16  150  10.404779911
     17     17   10  0.865840911865
     18     17   50  3.9999294281
     19     17  100  7.70257949829
     20     17  150  11.3758206367
     21     18   10  0.915451049805
     22     18   50  4.22935962677
     23     18  100  8.22361946106
     24     18  150  12.2081303596
    

    来自深度为 10 的树的 10 个大小为 10 的样本:

    ['1111010111', '1110111010', '1010111010', '011110010', '0111100001', '011101110', '01110010', '01001111', '0001000100', '000001010']
    ['110', '0110101110', '0110001100', '0011110', '0001111011', '0001100010', '0001100001', '0001100000', '0000011010', '0000001111']
    ['11010000', '1011111101', '1010001101', '1001110001', '1001100110', '10001110', '011111110', '011001100', '0101110000', '001110101']
    ['11111101', '110111', '110110111', '1101010101', '1101001011', '1001001100', '100100010', '0100001010', '0100000111', '0010010110']
    ['111101000', '1110111101', '1101101', '1101000000', '1011110001', '0111111101', '01101011', '011010011', '01100010', '0101100110']
    ['1111110001', '11000110', '1100010100', '101010000', '1010010001', '100011001', '100000110', '0100001111', '001101100', '0001101101']
    ['111110010', '1110100', '1101000011', '101101', '101000101', '1000001010', '0111100', '0101010011', '0101000110', '000100111']
    ['111100111', '1110001110', '1100111111', '1100110010', '11000110', '1011111111', '0111111', '0110000100', '0100011', '0010110111']
    ['1101011010', '1011111', '1011100100', '1010000010', '10010', '1000010100', '0111011111', '01010101', '001101', '000101100']
    ['111111110', '111101001', '1110111011', '111011011', '1001011101', '1000010100', '0111010101', '010100110', '0100001101', '0010000000']
    

    【讨论】:

    • 有趣的是看到一个 python 替代品。看起来性能与我的基于整数的非完整方法相当。您知道这是如何缩放的(即,如果您绘制 1K 而不是 100,或者使用尺寸 8 而不是 16)?我实际上预计这会更快,因为我的方法中的主要减速之一是由于其矢量化性质,难以有效地使用 R 创建递减池。
    • @BrodieG 计算最密集的部分是数字生成。我发现random.randint(0,n) 对于非常大的 n 确实仍然“有效”,因此切换到那个。性能提升令人难以置信,现在我降到了 8 毫秒!采样时间随着n 的增加而增加,而不是随着深度增加。这是因为combs 必须将n 分解为所有可能的“拆分”,正如我在代码中所说的那样。然后下一个递归实例必须做同样的事情。分支因素是疯狂的。即使有缓冲区,从深度为 17 的树中提取 100k 个样本也需要很长时间。
    • 酷;你介意添加一些示例输出吗?
    • 另外,我可能遗漏了一些东西,但这是否表明存在非线性?您的时间表明应该接近 100k 样本。也许问题在于 2 ^ 17 ~= 100k?
    • 当然,我稍后会添加一些示例。时间确实表明n 线性增加,但是我可以向您保证,采样 100k 个项目需要的时间远远超过 8 秒。 5分钟后我停止了。我猜它只有在combBuffer 被充分填充时才是线性的。
    【解决方案5】:

    将 id 映射到字符串。您可以将数字映射到 0/1 向量,正如 @BrodieG 提到的:

    # some key objects
    
    n_pool      = sum(2^(1:max_len))      # total number of indices
    cuts        = cumsum(2^(1:max_len-1)) # new group starts
    inds_by_g   = mapply(seq,cuts,cuts*2) # indices grouped by length
    
    # the mapping to strings (one among many possibilities)
    
    library(data.table)
    get_01str <- function(id,max_len){
        cuts = cumsum(2^(1:max_len-1))
        g    = findInterval(id,cuts)
        gid  = id-cuts[g]+1
    
        data.table(g,gid)[,s:=
          do.call(paste,c(list(sep=""),lapply(
            seq(g[1]), 
            function(x) (gid-1) %/% 2^(x-1) %% 2
          )))
        ,by=g]$s      
    } 
    

    寻找要删除的 id。我们将从采样池中依次删除 ids:

     # the mapping from one index to indices of nixed strings
    
    get_nixstrs <- function(g,gid,max_len){
    
        cuts         = cumsum(2^(1:max_len-1))
        gids_child   = {
          x = gid%%2^sequence(g-1)
          ifelse(x,x,2^sequence(g-1))
        }
        ids_child    = gids_child+cuts[sequence(g-1)]-1
    
        ids_parent   = if (g==max_len) gid+cuts[g]-1 else {
    
          gids_par       = vector(mode="list",max_len)
          gids_par[[g]]  = gid
          for (gg in seq(g,max_len-1)) 
            gids_par[[gg+1]] = c(gids_par[[gg]],gids_par[[gg]]+2^gg)
    
          unlist(mapply(`+`,gids_par,cuts-1))
        }
    
        c(ids_child,ids_parent)
    }
    

    索引按g、字符数nchar(get_01str(id)) 分组。因为索引按g 排序,所以g=findInterval(id,cuts) 是一条更快的路线。

    g1 &lt; g &lt; max_len 组中的索引具有一个大小为 g-1 的“子”索引和两个大小为 g+1 的父索引。对于每个子节点,我们取它的子节点,直到我们点击g==1;对于每个父节点,我们取他们的一对父节点,直到我们点击g==max_len

    就组内的标识符gid 而言,树的结构是最简单的。 gid 映射到两个父母,gidgid+2^g;并反转此映射找到孩子。

    抽样

    drawem <- function(n,max_len){
        cuts        = cumsum(2^(1:max_len-1))
        inds_by_g   = mapply(seq,cuts,cuts*2)
    
        oklens = (1:max_len)[ n <= 2^max_len*(1-2^(-(1:max_len)))+1 ]
        okinds = unlist(inds_by_g[oklens])
    
        mysamp = rep(0,n)
        for (i in 1:n){
    
            id        = if (length(okinds)==1) okinds else sample(okinds,1)
            g         = findInterval(id,cuts)
            gid       = id-cuts[g]+1
            nixed     = get_nixstrs(g,gid,max_len)
    
            # print(id); print(okinds); print(nixed)
    
            mysamp[i] = id
            okinds    = setdiff(okinds,nixed)
            if (!length(okinds)) break
        }
    
        res <- rep("",n)
        res[seq.int(i)] <- get_01str(mysamp[seq.int(i)],max_len)
        res
    }
    

    oklens 部分集成了 OP 的想法,即省略保证使采样不可能的字符串。然而,即使这样做,我们也可能会遵循一条让我们别无选择的采样路径。以 OP 的 max_len=4n=10 为例,我们知道我们必须从考虑中删除 01,但是如果我们的前四次抽奖是 000111 和 @ 会发生什么987654347@?哦,好吧,我想我们运气不好。这就是为什么您应该实际定义采样概率的原因。 (OP 有另一个想法,用于在每一步确定哪些节点将导致不可能的状态,但这似乎是一项艰巨的任务。)

    插图

    # how the indices line up
    
    n_pool = sum(2^(1:max_len)) 
    pdt <- data.table(id=1:n_pool)
    pdt[,g:=findInterval(id,cuts)]
    pdt[,gid:=1:.N,by=g]
    pdt[,s:=get_01str(id,max_len)]
    
    # example run
    
    set.seed(4); drawem(5,5)
    # [1] "01100" "1"     "0001"  "0101"  "00101"
    
    set.seed(4); drawem(8,4)
    # [1] "1100" "0"    "111"  "101"  "1101" "100"  ""     ""  
    

    基准(比@BrodieG 的答案中的那些更早)

    require(rbenchmark)
    max_len = 8
    n = 8
    
    benchmark(
          jos_lp     = {
            pool <- unlist(lapply(seq_len(max_len),
              function(x) do.call(paste0, expand.grid(rep(list(c('0', '1')), x)))))
            sample.lp(pool, n)},
          bro_string = {pool <- make_pool(max_len);sample01(pool,n)},
          fra_num    = drawem(n,max_len),
          replications=5)[1:5]
    #         test replications elapsed relative user.self
    # 2 bro_string            5    0.05      2.5      0.05
    # 3    fra_num            5    0.02      1.0      0.02
    # 1     jos_lp            5    1.56     78.0      1.55
    
    n = 12
    max_len = 12
    benchmark(
      bro_string={pool <- make_pool(max_len);sample01(pool,n)},
      fra_num=drawem(n,max_len),
      replications=5)[1:5]
    #         test replications elapsed relative user.self
    # 1 bro_string            5    0.54     6.75      0.51
    # 2    fra_num            5    0.08     1.00      0.08
    

    其他答案。还有两个其他答案:

    jos_enum = {pool <- unlist(lapply(seq_len(max_len), 
        function(x) do.call(paste0, expand.grid(rep(list(c('0', '1')), x)))))
      get.template(pool, n)}
    bro_num  = sample011(max_len,n)    
    

    我省略了@josilber 的枚举方法,因为它花费的时间太长了;和@BrodieG 的数字/索引方法,因为它当时不起作用,但现在起作用了。有关更多基准测试,请参阅 @BrodieG 的更新答案。

    速度与正确性。虽然@josilber 的答案要慢得多(而且对于枚举方法,显然更占用内存),但他们保证会在第一次尝试。使用@BrodieG 的字符串方法或此答案,您将不得不一次又一次地重新采样,以期画出完整的n。使用大的max_len,我想这应该不是问题。

    这个答案比bro_string 扩展得更好,因为它不需要预先构造pool

    【讨论】:

    • 我认为(没有完全消化它)这与我正在考虑的整数二进制原理相似。请注意,在您的示例中,您同时拥有“111”和“1111”样本,我认为这不应该发生。同样,“101”和“1011”。
    • @BrodieG 好吧,我想我修好了。不太确定。我只是发布,因为您基于整数的答案的评论流表明它不能完全满足 OP 的要求。
    • 乍一看,答案似乎是有效的。我的整数方法的主要问题是需要采取一个额外的步骤(即,如果我对“110”进行采样,则也从池中删除“11”和“1”),看来你已经处理好了。一件奇怪的事情,尝试set.seed(1); replicate(10, drawem(8, 3)) 我总是得到三个字符的结果。不过我没有仔细看逻辑。
    • 实际上,3 char 结果可能是由于排除了会阻止完整绘制的值。
    • 太好了,弗兰克。谢谢你的修复。在系统允许的情况下,将奖励您的麻烦!
    【解决方案6】:

    如果您不想生成所有可能元组的集合然后随机采样(您注意到这对于大输入大小可能是不可行的),另一种选择是使用整数规划绘制单个样本。基本上,您可以为pool 中的每个元素分配一个随机值,然后选择具有最大值总和的可行元组。这应该使每个元组被选中的概率相等,因为它们的大小都相同,并且它们的值是随机选择的。模型的约束将确保不选择不允许的元组对,并且选择正确数量的元素。

    这是lpSolve 包的解决方案:

    library(lpSolve)
    sample.lp <- function(pool, max_len) {
      pool <- sort(pool)
      pml <- max(nchar(pool))
      runs <- c(rev(cumsum(2^(seq(pml-1)))), 0)
      banned.from <- rep(seq(pool), runs[nchar(pool)])
      banned.to <- banned.from + unlist(lapply(runs[nchar(pool)], seq_len))
      banned.constr <- matrix(0, nrow=length(banned.from), ncol=length(pool))
      banned.constr[cbind(seq(banned.from), banned.from)] <- 1
      banned.constr[cbind(seq(banned.to), banned.to)] <- 1
      mod <- lp(direction="max",
                objective.in=runif(length(pool)),
                const.mat=rbind(banned.constr, rep(1, length(pool))),
                const.dir=c(rep("<=", length(banned.from)), "=="),
                const.rhs=c(rep(1, length(banned.from)), max_len),
                all.bin=TRUE)
      pool[which(mod$solution == 1)]
    }
    set.seed(144)
    pool <- unlist(lapply(seq_len(4), function(x) do.call(paste0, expand.grid(rep(list(c('0', '1')), x)))))
    sample.lp(pool, 4)
    # [1] "0011" "010"  "1000" "1100"
    sample.lp(pool, 8)
    # [1] "0000" "0100" "0110" "1001" "1010" "1100" "1101" "1110"
    

    这似乎可以扩展到相当大的池。例如,从大小为 510 的池中获取长度为 20 的样本需要 2 秒多一点:

    pool <- unlist(lapply(seq_len(8), function(x) do.call(paste0, expand.grid(rep(list(c('0', '1')), x)))))
    length(pool)
    # [1] 510
    system.time(sample.lp(pool, 20))
    #    user  system elapsed 
    #   0.232   0.008   0.239 
    

    如果您需要解决非常非常庞大的问题,那么您可以从 lpSolve 附带的非开源求解器转移到 gurobi 或 cplex 等商业求解器(一般不是免费的,但可免费用于学术用途)。

    【讨论】:

    • 不错的方法。不幸的是,对于较大的pools(例如,11.5 秒从具有max_len = 9 的池中抽取 4 个),这似乎变得相当慢,而对于大型池和小型 n,我的 Q 中的 while 方法几乎是即时的.我想我可以根据 b/w nmax_len 的关系在这些方法之间切换。感谢您提供有关可用的各种求解器的提示。
    • @jbaums 您在 11.5 秒的计时中使用了多大的池?是的,我希望这对于max_len 的大值是最好的,在这种情况下拒绝抽样永远找不到任何东西并且枚举是棘手的。
    • 大小为 1022(元素的最大长度为 9 个字符),我正在绘制大小为 4 的样本。
    • @jbaums 大部分时间都用于计算不能一起使用的元素对。我将其更改为使用pool 结构的版本,您引用的示例现在运行时间为 2.2 秒(大部分用于解决优化问题)。
    【解决方案7】:

    一种方法是使用迭代方法简单地生成所有可能的适当大小的元组:

    1. 构建所有大小为 1 的元组(pool 中的所有元素)
    2. pool 中的元素进行叉积
    3. 多次删除使用pool 的相同元素的任何元组
    4. 删除任何与另一个元组完全相同的重复项
    5. 删除任何不能一起使用的元组
    6. 冲洗并重复,直到获得合适的元组大小

    这对于给定的大小是可运行的(pool 长度为 30,max_len 4):

    get.template <- function(pool, max_len) {
      banned <- which(outer(paste0('^', pool), pool, Vectorize(grepl)), arr.ind=T)
      banned <- banned[banned[,1] != banned[,2],]
      banned <- paste(banned[,1], banned[,2])
      vals <- matrix(seq(length(pool)))
      for (k in 2:max_len) {
        vals <- cbind(vals[rep(1:nrow(vals), each=length(pool)),],
                      rep(1:length(pool), nrow(vals)))
        # Can't sample same value more than once
        vals <- vals[apply(vals, 1, function(x) length(unique(x)) == length(x)),]
        # Sort rows to ensure unique only
        vals <- t(apply(vals, 1, sort))
        vals <- unique(vals)
        # Can't have banned pair
        combos <- combn(ncol(vals), 2)
        for (k in seq(ncol(combos))) {
            c1 <- combos[1,k]
            c2 <- combos[2,k]
            vals <- vals[!paste(vals[,c1], vals[,c2]) %in% banned,]
        }
      }
      return(matrix(pool[vals], nrow=nrow(vals)))
    }
    
    max_len <- 4
    pool <- unlist(lapply(seq_len(max_len), function(x) do.call(paste0, expand.grid(rep(list(c('0', '1')), x)))))
    system.time(template <- get.template(pool, 4))
    #   user  system elapsed 
    #  4.549   0.050   4.614 
    

    现在您可以从template 的行中任意多次采样(这将非常快),这与从定义的空间中随机采样相同。

    【讨论】:

      【解决方案8】:

      简介

      我发现这个问题非常有趣,以至于我不得不仔细考虑,并最终提供我自己的答案。由于我得出的算法并没有立即从问题描述中得出,所以我将首先解释我是如何得出这个解决方案的,然后提供一个 C++ 的示例实现(我从未写过 R)。

      解决方案的开发

      一目了然

      阅读问题描述最初令人困惑,但是当我看到带有树木图片的编辑时,我立即理解了问题,并且我的直觉表明二叉树也是一种解决方案:构建一棵树(一个集合大小为 1) 的树,并在进行选择时消除分支和祖先后将树分解为较小的树的集合。

      虽然这最初看起来不错,但收藏的选择过程和维护会很痛苦。尽管如此,这棵树似乎应该在任何解决方案中发挥重要作用。

      修订版 1

      不要破坏树。相反,在每个节点上都有一个布尔数据有效负载,指示它是否已经被消除。这样就只剩下一棵树保持形式了。

      但请注意,这不仅仅是任何二叉树,它实际上是深度为 max_len-1 的完整二叉树。

      修订版 2

      一个完整的二叉树可以很好地表示为一个数组。典型的数组表示使用树的广度优先搜索,具有以下性质:

      Let x be the array index.
      x = 0 is the root of the entire tree
      left_child(x) = 2x + 1
      right_child(x) = 2x + 2
      parent(x) = floor((n-1)/2)
      

      在下图中,每个节点都标有其数组索引:

      作为一个数组,这占用更少的内存(不再有指针),使用的内存是连续的(有利于缓存),并且可以完全放在堆栈上而不是堆上(假设你的语言给你一个选择)。当然,这里有一些条件适用,特别是数组的大小。稍后我会谈到这一点。

      就像在修订版 1 中一样,存储在数组中的数据将是布尔值:true 表示可用,false 表示不可用。由于根节点实际上不是一个有效的选择,索引 0 应该被初始化为 false。如何进行选择的问题仍然存在:

      由于指标已被限制,因此跟踪已消除的指标数量以及剩余的指标数量是微不足道的。在该范围内选择一个随机数,然后遍历数组,直到看到许多索引设置为 true(包括当前索引)。到达的索引就是要做出的选择。 选择直到选择了 n 个指标,或者没有任何内容可供选择。

      这是一个完整的算法,可以工作,但是在选择过程中还有改进的空间,还有一个实际的大小问题尚未解决:数组大小为 O(2^n )。随着 n 变大,首先缓存的好处消失了,然后数据开始被分页到磁盘,在某些时候它变得根本无法存储。

      修订版 3

      我决定先解决更简单的问题:改进选择过程。

      从左到右扫描数组很浪费。跟踪已消除的范围可能比连续检查和查找多个错误更有效。然而,我们的树表示并不理想,因为每轮将要消除的节点中很少有在数组中是连续的。

      通过重新排列数组映射到树的方式,可以更好地利用这一点。特别是,让我们使用前序深度优先搜索而不是广度优先搜索。为了做到这一点,树的大小需要固定,这就是这个问题的情况。子节点和父节点的索引在数学上是如何连接的也不太明显。

      通过使用这种安排,可以保证每个不是叶子的选择都消除一个连续的范围:它的子树。

      修订版 4

      通过跟踪消除的范围,不再需要真/假数据,因此根本不需要数组或树。 在每次随机抽取时,消除的范围可用于快速找到要选择的节点。所有祖先和整个子树都被消除,并且可以表示为可以轻松与其他合并的范围。

      最后的任务是将选定的节点转换为 OP 想要的字符串表示形式。这很容易,因为这棵二叉树仍然保持严格的顺序:从根开始遍历,所有 >= 右孩子的元素都在右边,其他元素在左边。因此搜索树将通过在向左遍历时附加“0”来提供祖先列表和二进制字符串;或右转时为“1”。

      示例实现

      #include <stdint.h>
      #include <algorithm>
      #include <cmath>
      #include <list>
      #include <deque>
      #include <ctime>
      #include <cstdlib>
      #include <iostream>
      
      /*
       * A range of values of the form (a, b), where a <= b, and is inclusive.
       * Ex (1,1) is the range from 1 to 1 (ie: just 1)
       */
      class Range
      {
      private:
          friend bool operator< (const Range& lhs, const Range& rhs);
          friend std::ostream& operator<<(std::ostream& os, const Range& obj);
      
          int64_t m_start;
          int64_t m_end;
      
      public:
          Range(int64_t start, int64_t end) : m_start(start), m_end(end) {}
          int64_t getStart() const { return m_start; }
          int64_t getEnd() const { return m_end; }
          int64_t size() const { return m_end - m_start + 1; }
          bool canMerge(const Range& other) const {
              return !((other.m_start > m_end + 1) || (m_start > other.m_end + 1));
          }
          int64_t merge(const Range& other) {
              int64_t change = 0;
              if (m_start > other.m_start) {
                  change += m_start - other.m_start;
                  m_start = other.m_start;
              }
              if (other.m_end > m_end) {
                  change += other.m_end - m_end;
                  m_end = other.m_end;
              }
              return change;
          }
      };
      
      inline bool operator< (const Range& lhs, const Range& rhs){return lhs.m_start < rhs.m_start;}
      std::ostream& operator<<(std::ostream& os, const Range& obj) {
          os << '(' << obj.m_start << ',' << obj.m_end << ')';
          return os;
      }
      
      /*
       * Stuct to allow returning of multiple values
       */
      struct NodeInfo {
          int64_t subTreeSize;
          int64_t depth;
          std::list<int64_t> ancestors;
          std::string representation;
      };
      
      /*
       * Collection of functions representing a complete binary tree
       * as an array created using pre-order depth-first search,
       * with 0 as the root.
       * Depth of the root is defined as 0.
       */
      class Tree
      {
      private:
          int64_t m_depth;
      public:
          Tree(int64_t depth) : m_depth(depth) {}
          int64_t size() const {
              return (int64_t(1) << (m_depth+1))-1;
          }
          int64_t getDepthOf(int64_t node) const{
              if (node == 0) { return 0; }
              int64_t searchDepth = m_depth;
              int64_t currentDepth = 1;
              while (true) {
                  int64_t rightChild = int64_t(1) << searchDepth;
                  if (node == 1 || node == rightChild) {
                      break;
                  } else if (node > rightChild) {
                      node -= rightChild;
                  } else {
                      node -= 1;
                  }
                  currentDepth += 1;
                  searchDepth -= 1;
              }
              return currentDepth;
          }
          int64_t getSubtreeSizeOf(int64_t node, int64_t nodeDepth = -1) const {
              if (node == 0) {
                  return size();
              }
              if (nodeDepth == -1) {
                  nodeDepth = getDepthOf(node);
              }
              return (int64_t(1) << (m_depth + 1 - nodeDepth)) - 1;
          }
          int64_t getLeftChildOf(int64_t node, int64_t nodeDepth = -1) const {
              if (nodeDepth == -1) {
                  nodeDepth = getDepthOf(node);
              }
              if (nodeDepth == m_depth) { return -1; }
              return node + 1;
          }
          int64_t getRightChildOf(int64_t node, int64_t nodeDepth = -1) const {
              if (nodeDepth == -1) {
                  nodeDepth = getDepthOf(node);
              }
              if (nodeDepth == m_depth) { return -1; }
              return node + 1 + ((getSubtreeSizeOf(node, nodeDepth) - 1) / 2);
          }
          NodeInfo getNodeInfo(int64_t node) const {
              NodeInfo info;
              int64_t depth = 0;
              int64_t currentNode = 0;
              while (currentNode != node) {
                  if (currentNode != 0) {
                      info.ancestors.push_back(currentNode);
                  }
                  int64_t rightChild = getRightChildOf(currentNode, depth);
                  if (rightChild == -1) {
                      break;
                  } else if (node >= rightChild) {
                      info.representation += '1';
                      currentNode = rightChild;
                  } else {
                      info.representation += '0';
                      currentNode = getLeftChildOf(currentNode, depth);
                  }
                  depth++;
              }
              info.depth = depth;
              info.subTreeSize = getSubtreeSizeOf(node, depth);
              return info;
          }
      };
      
      // random selection amongst remaining allowed nodes
      int64_t selectNode(const std::deque<Range>& eliminationList, int64_t poolSize, std::mt19937_64& randomGenerator)
      {
          std::uniform_int_distribution<> randomDistribution(1, poolSize);
          int64_t selection = randomDistribution(randomGenerator);
          for (auto const& range : eliminationList) {
              if (selection >= range.getStart()) { selection += range.size(); }
              else { break; }
          }
          return selection;
      }
      
      // determin how many nodes have been elimintated
      int64_t countEliminated(const std::deque<Range>& eliminationList)
      {
          int64_t count = 0;
          for (auto const& range : eliminationList) {
              count += range.size();
          }
          return count;
      }
      
      // merge all the elimination ranges to listA, and return the number of new elimintations
      int64_t mergeEliminations(std::deque<Range>& listA, std::deque<Range>& listB) {
          if(listB.empty()) { return 0; }
          if(listA.empty()) {
              listA.swap(listB);
              return countEliminated(listA);
          }
      
          int64_t newEliminations = 0;
          int64_t x = 0;
          auto listA_iter = listA.begin();
          auto listB_iter = listB.begin();
          while (listB_iter != listB.end()) {
              if (listA_iter == listA.end()) {
                  listA_iter = listA.insert(listA_iter, *listB_iter);
                  x = listB_iter->size();
                  assert(x >= 0);
                  newEliminations += x;
                  ++listB_iter;
              } else if (listA_iter->canMerge(*listB_iter)) {
                  x = listA_iter->merge(*listB_iter);
                  assert(x >= 0);
                  newEliminations += x;
                  ++listB_iter;
              } else if (*listB_iter < *listA_iter) {
                  listA_iter = listA.insert(listA_iter, *listB_iter) + 1;
                  x = listB_iter->size();
                  assert(x >= 0);
                  newEliminations += x;
                  ++listB_iter;
              } else if ((listA_iter+1) != listA.end() && listA_iter->canMerge(*(listA_iter+1))) {
                  listA_iter->merge(*(listA_iter+1));
                  listA_iter = listA.erase(listA_iter+1);
              } else {
                  ++listA_iter;
              }
          }
          while (listA_iter != listA.end()) {
              if ((listA_iter+1) != listA.end() && listA_iter->canMerge(*(listA_iter+1))) {
                  listA_iter->merge(*(listA_iter+1));
                  listA_iter = listA.erase(listA_iter+1);
              } else {
                  ++listA_iter;
              }
          }
          return newEliminations;
      }
      
      int main (int argc, char** argv)
      {
          std::random_device rd;
          std::mt19937_64 randomGenerator(rd());
      
          int64_t max_len = std::stoll(argv[1]);
          int64_t num_samples = std::stoll(argv[2]);
      
          int64_t samplesRemaining = num_samples;
          Tree tree(max_len);
          int64_t poolSize = tree.size() - 1;
          std::deque<Range> eliminationList;
          std::deque<Range> eliminated;
          std::list<std::string> foundList;
      
          while (samplesRemaining > 0 && poolSize > 0) {
              // find a valid node
              int64_t selectedNode = selectNode(eliminationList, poolSize, randomGenerator);
              NodeInfo info = tree.getNodeInfo(selectedNode);
              foundList.push_back(info.representation);
              samplesRemaining--;
      
              // determine which nodes this choice eliminates
              eliminated.clear();
              for( auto const& ancestor : info.ancestors) {
                  Range r(ancestor, ancestor);
                  if(eliminated.empty() || !eliminated.back().canMerge(r)) {
                      eliminated.push_back(r);
                  } else {
                      eliminated.back().merge(r);
                  }
              }
              Range r(selectedNode, selectedNode + info.subTreeSize - 1);
              if(eliminated.empty() || !eliminated.back().canMerge(r)) {
                  eliminated.push_back(r);
              } else {
                  eliminated.back().merge(r);
              }
      
              // add the eliminated nodes to the existing list
              poolSize -= mergeEliminations(eliminationList, eliminated);
          }
      
          // Print some stats
          // std::cout << "tree: " << tree.size() << " samplesRemaining: "
          //                       << samplesRemaining << " poolSize: "
          //                       << poolSize << " samples: " << foundList.size()
          //                       << " eliminated: "
          //                       << countEliminated(eliminationList) << std::endl;
      
          // Print list of binary strings
          // std::cout << "list:";
          // for (auto const& s : foundList) {
          //  std::cout << " " << s;
          // }
          // std::cout << std::endl;
      }
      

      其他想法

      该算法对于 max_len 的扩展性非常好。使用 n 缩放不是很好,但根据我自己的分析,它似乎比其他解决方案做得更好。

      可以轻松修改此算法,以允许包含不仅仅是“0”和“1”的字符串。 单词中更多可能的字母会增加树的扇出,并且每次选择都会消除更广泛的范围 - 每个子树中的所有节点仍然保持连续。

      【讨论】:

        猜你喜欢
        • 1970-01-01
        • 1970-01-01
        • 1970-01-01
        • 2015-08-20
        • 2011-12-06
        • 1970-01-01
        • 2018-04-28
        • 1970-01-01
        • 1970-01-01
        相关资源
        最近更新 更多