diff --git a/nimbleModel/R/all_utils.R b/nimbleModel/R/all_utils.R index fd3e369..1aabc54 100644 --- a/nimbleModel/R/all_utils.R +++ b/nimbleModel/R/all_utils.R @@ -202,3 +202,14 @@ evalNumeric <- function(expr) { } return(expr) } + +# Flatten nested lists. +flatten <- function(x) { + result <- do.call(c, x) + names(result) <- NULL + if (identical(result, list(NULL))) { + return(NULL) + } + result <- result[!sapply(result, is.null)] + return(result) +} diff --git a/nimbleModel/R/indexRange.R b/nimbleModel/R/indexRange.R index c547159..cba42f0 100644 --- a/nimbleModel/R/indexRange.R +++ b/nimbleModel/R/indexRange.R @@ -416,3 +416,65 @@ matrixExpandGrid <- function(matrixList) { ) return(do.call("cbind", unfoldedMatrices)) } + +combine_indexRanges <- function(range1, range2) { + if(is(range1, 'indexRangeMatrixClass') && is(range2, 'indexRangeMatrixClass')) + return(newIndexRange(unique(rbind(range1$values, range2$values)))) + if(is(range1, 'indexRangeMatrixClass') && is(range2, 'indexRangeSequenceClass')) { + tmp <- range1; range1 <- range2; range2 <- tmp + } + if(is(range1, 'indexRangeSequenceClass') && is(range2, 'indexRangeMatrixClass')) { + keep <- range2$values < range1$start | range2$values >= range1$start + range1$numElements + newValues <- range2$values[keep] + if(length(newValues)) { + range2 <- newIndexRange(newValues) + return(c(range1, range2)) + } else return(range1) + } + if(is(range1, 'indexRangeMatrixClass') && is(range2, 'indexRangeScalarClass')) { + tmp <- range1; range1 <- range2; range2 <- tmp + } + if(is(range1, 'indexRangeScalarClass') && is(range2, 'indexRangeMatrixClass')) { + if(range1$value %in% range2$values) + return(range2) + return(newIndexRange(c(range1$value, range2$values))) + } + if(is(range1, 'indexRangeSequenceClass') && is(range2, 'indexRangeScalarClass')) { + tmp <- range1; range1 <- range2; range2 <- tmp + } + if(is(range1, 'indexRangeScalarClass') && is(range2, 'indexRangeSequenceClass')) { + if(range1$value >= range2$start && range1$value < range2$start + range2$numElements) + return(range2) + if(range1$value == range2$start - 1) + return(newIndexRange(substitute(START:END, list(START=range2$start - 1, END=range2$start+range2$numElements-1)))) + if(range1$value == range2$start + range2$numElements) + return(newIndexRange(substitute(START:END, list(START=range2$start, END=range2$start+range2$numElements)))) + return(list(range2, range1)) + } + if(is(range1, 'indexRangeScalarClass') && is(range2, 'indexRangeScalarClass')) { + if(range1$value == range2$value) + return(range1) + if(abs(range1$value - range2$value) == 1) { + vals <- sort(c(range1$value, range2$value)) + return(newIndexRange(substitute(START:END, list(START=vals[1],END=vals[1]+1)))) + } + return(newIndexRange(c(range1$value, range2$value))) + } + if(is(range1, 'indexRangeSequenceClass') && is(range2, 'indexRangeSequenceClass')) { + if(range2$start < range1$start) { + tmp <- range1; range1 <- range2; range2 <- tmp + } + if(range1$start == range2$start) + return(newIndexRange(substitute(START:END, list(START=range1$start, + END=range1$start+max(c(range1$numElements,range2$numElements))-1)))) + if(range2$start <= range1$start+range1$numElements) + return(newIndexRange(substitute(START:END, list(START=range1$start, + END=max(c(range1$start+range1$numElements-1), + range2$start+range2$numElements-1))))) + return(c(range1, range2)) + } + stop("Unexpected input ranges") +} + + + diff --git a/nimbleModel/R/instructions.R b/nimbleModel/R/instructions.R index 0f8b402..4e454ec 100644 --- a/nimbleModel/R/instructions.R +++ b/nimbleModel/R/instructions.R @@ -193,6 +193,7 @@ makeInstrList <- function(model, input, includeData = TRUE, use_vec = FALSE) { rule$makeCalcRange(rule$apply(vr)) }) })) + ranges <- aggregate_calcRanges(ranges) sortIDs <- lapply(ranges, \(x) x$sortID) sortIDranges <- sapply(sortIDs, \(x) range(x, na.rm = TRUE)) diff --git a/nimbleModel/R/nodeRules.R b/nimbleModel/R/nodeRules.R index 6b5af93..f75bbd6 100644 --- a/nimbleModel/R/nodeRules.R +++ b/nimbleModel/R/nodeRules.R @@ -581,6 +581,31 @@ calcRangeClass <- R6Class( ) ) +aggregate_calcRanges <- function(rangeSet) { + if (!length(rangeSet)) { + return(rangeSet) + } + if (is.character(rangeSet) || !is.list(rangeSet) || + !all(sapply(rangeSet, \(x) inherits(x, 'calcRangeClass')))) { + stop("`rangeSet` must be a list of calcRanges") + } + declIDs <- sapply(rangeSet, \(x) x$declID) + rangesByDecl <- split(rangeSet, declIDs) + lens <- sapply(rangesByDecl, length) + if(exists('paciorek')) browser() + if(any(lens > 1)) { + for(i in which(lens > 1)) { + result <- combine_indexingRanges(rangesByDecl[[i]][[1]], rangesByDecl[[i]][[2]]) + idx <- 3 + while(idx <= lens[[i]]) { + result <- combine_indexingRanges(result, rangesByDecl[[i]][[idx]]) + idx <- idx+1 + } + rangesByDecl[[i]] <- result + } + } + return(flatten(rangesByDecl)) +} # Class for managing a set of like nodes (same declaration, but not necessarily same graph role or same sort ID). # Basically a `varRange` but with indication of which indexRanges relate to node indexing (external indexRanges) diff --git a/nimbleModel/R/varRange.R b/nimbleModel/R/varRange.R index 3a032e3..10e8191 100644 --- a/nimbleModel/R/varRange.R +++ b/nimbleModel/R/varRange.R @@ -403,19 +403,19 @@ removeDuplicateVarRangesOne <- function(varRanges) { return(varRanges[!dups]) } -# Flatten nested lists. -flatten <- function(x) { - result <- do.call(c, x) - names(result) <- NULL - if (identical(result, list(NULL))) { - return(NULL) - } - result <- result[!sapply(result, is.null)] - return(result) + +combine_indexingRanges <- function(range1, range2) { + if(!identical(range1$rangeToIndexSlot, range2$rangeToIndexSlot)) + stop("unexpected incompatibility between indexing ranges when combining calcRanges") + range <- range1$clone() + for(i in seq_along(range$indexingRange$indexRanges)) + range$indexingRange$indexRanges[[i]] <- combine_indexRanges(range1$indexingRange$indexRanges[[i]], range2$indexingRange$indexRanges[[i]]) + return(range) } # TODO: need combine() that combines "adjacent" varRanges +# the a # scalar+seq = seq # seq + seq = seq diff --git a/nimbleModel/tests/testthat/test-nimbleModel.R b/nimbleModel/tests/testthat/test-nimbleModel.R index 02ba833..b3020ed 100644 --- a/nimbleModel/tests/testthat/test-nimbleModel.R +++ b/nimbleModel/tests/testthat/test-nimbleModel.R @@ -1697,4 +1697,89 @@ test_that("duplication cases", { expect_identical(m$getNodes(c('z[1]','z[2]'), nodesAsChars = TRUE), c('z[1:3]')) + code <- nimbleCode({ + for(i in 1:3) + for(j in 1:2) + y[i,j,i+1] ~ dnorm(0,1) + }) + m <- nimbleModel(code) + rule <- m$modelDef$calcRules[['y']]$rules[[1]] + ranges <- c(rule$makeCalcRange(rule$apply('y')), + rule$makeCalcRange(rule$apply('y[2,1:2,3]'))) + newRanges <- nimbleModel:::aggregate_calcRanges(ranges) + expect_identical(length(newRanges), 1L) + expect_equal(ranges[[1]],newRanges[[1]]) + + code <- nimbleCode({ + for(i in 1:3) + y[i] ~ dnorm(0,1) + y[4] ~ dnorm(0,1) + }) + set.seed(99) + mclass <- nimbleModel(code, data = list(y=rnorm(4)), returnClass = TRUE) + cmclass <- nCompile(mclass) + m <- mclass$new() + cm <- cmclass$new() + + result <- sum(dnorm(m$y[c(1,2,4)], log=TRUE)) + expect_identical(m$calculate(c('y[1]','y[4]','y[1]','y[2]')), + result) + expect_identical(cm$calculate(c('y[1]','y[4]','y[1]','y[2]')), + result) + + set.seed(1) + vals <- rnorm(3) + set.seed(1) + m$simulate(c('y[1]','y[4]','y[1]','y[2]'), includeData = TRUE) + expect_identical(vals, m$y[c(1,2,4)]) + set.seed(1) + cm$simulate(c('y[1]','y[4]','y[1]','y[2]'), includeData = TRUE) + expect_identical(vals, cm$y[c(1,2,4)]) + + code <- nimbleCode({ + for(i in 1:3) + for(j in 1:2) + y[i,j] ~ dnorm(0,1) + }) + set.seed(99) + mclass <- nimbleModel(code, data = list(y=matrix(rnorm(6),3)), returnClass = TRUE) + cmclass <- nCompile(mclass) + m <- mclass$new() + cm <- cmclass$new() + + nodes <- c('y[1,2]','y[1,2]','y[1:2,2]') + result <- sum(dnorm(c(m$y[1,2],m$y[2,2]),log=TRUE)) + expect_identical(m$calculate(nodes), result) + expect_identical(cm$calculate(nodes), result) + + set.seed(1) + vals <- rnorm(2) + set.seed(1) + m$simulate(nodes, includeData = TRUE) + expect_identical(vals, c(m$y[1,2],m$y[2,2])) + set.seed(1) + cm$simulate(nodes, includeData = TRUE) + expect_identical(vals, c(m$y[1,2],m$y[2,2])) + + code <- nimbleCode({ + y[1:3] ~ dmnorm(z[1:3],pr[1:3,1:3]) + }) + set.seed(99) + mclass <- nimbleModel(code, data = list(y=rnorm(3)), inits = list(z=rep(0,3),pr=diag(3)), returnClass = TRUE) + # cmclass <- nCompile(mclass) + m <- mclass$new() + # cm <- cmclass$new() + + ## TODO: add compiled simulate when {d,r}mnorm_chol is resolved. + m$calculate() + nodes <- c('y[1]','y[2]') + result <- sum(dnorm(m$y,log=TRUE)) + expect_equal(m$calculate(nodes), result) + + set.seed(1) + vals <- rnorm(3) + set.seed(1) + m$simulate(nodes, includeData = TRUE) + expect_identical(vals, m$y) + })