Time and memory efficient way to find intersection between pairs from a large list (R)

Viewed 88

I have a list of length N, and each list element is 10 character strings sampled from some bigger group. I want to find every pair of elements that are very similar, say with >=8 strings in common. I can do this alright for N<10,000 using crossprod, but I run out of memory for larger N.

Here's an example where N=1000, resulting in 6 very similar pairs:

# make list x
set.seed(1)
N = 1000
x = lapply(1:N, function(x) sample(letters, 10))
names(x) = as.character(1:length(x))

# find pairwise intersect, store in square matrix N_intersect
N_intersect = x %>% stack() %>% table() %>% crossprod()
diag(N_intersect) = 0
N_intersect[lower.tri(N_intersect)] = 0

# find when N_intersect is over some threshold
result = which(N_intersect > 8, arr.ind = T)

result
    ind ind
111 111 375
705 705 708
48   48 771
317 317 797
566 566 883
705 705 958

However since the output of crossprod is an NxN matrix, it quickly goes out of memory for large N. I know there are sparse methods for crossprod, but that only seems to reduce the memory by around ~30%.

The thing is, the number of highly similar elements is usually super small, like 1/1000 pairs. So I don't need to store the big square matrix from crossprod, but I can't think of a memory-efficient method that can do this quickly. I could just check each pair of elements in a for loop, but that takes a few hours.

1 Answers

Answering my own question in case it helps anyone.

I found a method to detect similar pairs in my list x that runs in approximately the same time as crossprod, but most important it doesn't go out of memory. I did this by taking the input to crossprod, x %>% stack() %>% table(), an Nxm matrix where m is the number of possible characters. m=26 in the example above, since I'm using letters. Each row represents a list element in x, with 1s denoting letters in the element and 0s letters not in it.

My question is then "which pairs of rows have at least 9 shared columns with 1s?" I answered this for each row by 1) subsetting columns and 2) rowSums. This method is f.new and I compare it to the old method f.crossprod:

f.crossprod = function(x, threshold) {
  # find intersection of each pair elements with crossprod
  # store in upper triangular square matrix
  N_intersect = x %>% stack() %>% table() %>% crossprod()
  diag(N_intersect) = 0
  N_intersect[lower.tri(N_intersect)] = 0
  
  # find when N_intersect is over some threshold
  result = which(N_intersect >= threshold, arr.ind = T) %>%
    as.data.frame() %>%
    `colnames<-`(c("row", "col"))
  rownames(result) = NULL
  return(result)
}

f.new = function(x, threshold) {
  x2 = t(table(stack(x)))
  result = list()
  for (ii in 3:nrow(x2)) {
    ia = which(x2[ii,]>0)
    cands = which(rowSums(x2[1:(ii-1),ia]) >= threshold)
    result[[ii]] = data.frame(row=cands, col=rep(ii, length(cands)))
  }
  result = do.call(rbind, result) %>% filter(!row==col)
  rownames(result) = NULL
  return(result)
}

I can see they both return the same results:

set.seed(1)
N = 1000
x = lapply(1:N, function(x) sample(letters, 10))
names(x) = as.character(1:length(x))

r1 = f.new(x, 9)
r2 = f.crossprod(x, 9)
identical(r1, r2)

> TRUE

The timing is similar, although f.new is slightly faster for N>5000. Most importantly f.new has no problem with N>20000:

enter image description here

Related