tests/testthat/test-rose.R

test_that("minority_prop value", {
  rec <- recipe(class ~ x + y, data = circle_example)
  rec21 <- rec |>
    step_rose(class, minority_prop = 0.1)

  rec22 <- rec |>
    step_rose(class, minority_prop = 0.2)

  rec21_p <- prep(rec21)
  rec22_p <- prep(rec22)

  tr_xtab1 <- table(bake(rec21_p, new_data = NULL)$class, useNA = "no")
  tr_xtab2 <- table(bake(rec22_p, new_data = NULL)$class, useNA = "no")

  expect_equal(sum(tr_xtab1), sum(tr_xtab2))

  expect_lt(tr_xtab1[["Circle"]], tr_xtab2[["Circle"]])
})

test_that("row matching works correctly #36", {
  expect_no_error(
    recipe(class ~ ., data = circle_example) |>
      step_rose(class, over_ratio = 1.2) |>
      prep()
  )

  expect_no_error(
    recipe(class ~ ., data = circle_example) |>
      step_rose(class, over_ratio = 0.8) |>
      prep()
  )

  expect_no_error(
    recipe(class ~ ., data = circle_example) |>
      step_rose(class, over_ratio = 1.7) |>
      prep()
  )
})

test_that("basic usage", {
  rec1 <- recipe(class ~ x + y, data = circle_example) |>
    step_rose(class)

  rec1_p <- prep(rec1)

  te_xtab <- table(bake(rec1_p, new_data = circle_example)$class, useNA = "no")
  og_xtab <- table(circle_example$class, useNA = "no")

  expect_equal(sort(te_xtab), sort(og_xtab))

  expect_no_warning(prep(rec1))
})

test_that("works with a single predictor", {
  skip_if_not_installed("modeldata")

  data("credit_data", package = "modeldata")

  expect_no_error(
    recipe(Status ~ Age, data = credit_data) |>
      step_rose(all_outcomes()) |>
      prep() |>
      bake(NULL)
  )
})

test_that("bad data", {
  rec <- recipe(~., data = circle_example)
  # numeric check
  expect_snapshot(
    error = TRUE,
    rec |>
      step_rose(x) |>
      prep()
  )
  # Multiple variable check
  expect_snapshot(
    error = TRUE,
    rec |>
      step_rose(class, id) |>
      prep()
  )
})

test_that("errors on unsupported predictor types", {
  df_date <- data.frame(
    x = factor(rep(c("a", "b"), c(2, 8))),
    y = as.Date("2020-01-01") + 1:10
  )

  expect_snapshot(
    error = TRUE,
    recipe(x ~ y, data = df_date) |>
      step_rose(x) |>
      prep()
  )
})

test_that("NA in response", {
  skip_if_not_installed("modeldata")

  data("credit_data", package = "modeldata")
  credit_data0 <- credit_data
  credit_data0[1, 1] <- NA

  expect_snapshot(
    error = TRUE,
    recipe(Status ~ Age, data = credit_data0) |>
      step_rose(Status) |>
      prep()
  )
})

test_that("`seed` produces identical sampling", {
  step_with_seed <- function(seed = sample.int(10^5, 1)) {
    recipe(class ~ x + y, data = circle_example) |>
      step_rose(class, seed = seed) |>
      prep() |>
      bake(new_data = NULL) |>
      pull(x)
  }

  run_1 <- step_with_seed(seed = 1234)
  run_2 <- step_with_seed(seed = 1234)
  run_3 <- step_with_seed(seed = 12345)

  expect_equal(run_1, run_2)
  expect_false(identical(run_1, run_3))
})

test_that("test tidy()", {
  rec <- recipe(class ~ x + y, data = circle_example) |>
    step_rose(class, id = "")

  rec_p <- prep(rec)

  untrained <- tibble(
    terms = "class",
    id = ""
  )

  trained <- tibble(
    terms = "class",
    id = ""
  )

  expect_equal(untrained, tidy(rec, number = 1))
  expect_equal(trained, tidy(rec_p, number = 1))
})

test_that("only except 2 classes", {
  df_char <- data.frame(
    x = factor(1:3),
    stringsAsFactors = FALSE
  )

  expect_snapshot(
    error = TRUE,
    recipe(~., data = df_char) |>
      step_rose(x) |>
      prep()
  )
})

test_that("factor levels are not affected by alphabet ordering or class sizes", {
  circle_example_alt_levels <- list()
  for (i in 1:4) {
    circle_example_alt_levels[[i]] <- circle_example
  }

  # Checking for forgetting levels by majority/minor switching
  for (i in c(2, 4)) {
    levels(circle_example_alt_levels[[i]]$class) <-
      rev(levels(circle_example_alt_levels[[i]]$class))
  }

  # Checking for forgetting levels by alphabetical switching
  for (i in c(3, 4)) {
    circle_example_alt_levels[[i]]$class <-
      factor(
        x = circle_example_alt_levels[[i]]$class,
        levels = rev(levels(circle_example_alt_levels[[i]]$class))
      )
  }

  for (i in 1:4) {
    rec_p <- recipe(class ~ x + y, data = circle_example_alt_levels[[i]]) |>
      step_rose(class) |>
      prep()

    expect_equal(
      levels(circle_example_alt_levels[[i]]$class), # Original levels
      rec_p$levels$class$values # New levels
    )
    expect_equal(
      levels(circle_example_alt_levels[[i]]$class), # Original levels
      levels(bake(rec_p, new_data = NULL)$class) # New levels
    )
  }
})

test_that("non-predictor variables are ignored", {
  circle_example2 <- circle_example |>
    mutate(id = as.character(row_number())) |>
    as_tibble()

  res <- recipe(class ~ ., data = circle_example2) |>
    update_role(id, new_role = "id") |>
    step_rose(class) |>
    prep() |>
    bake(new_data = NULL)

  expect_equal(
    c(circle_example2$id, rep(NA, nrow(res) - nrow(circle_example2))),
    as.character(res$id)
  )
})


test_that("id variables don't turn predictors to factors", {
  # https://github.com/tidymodels/themis/issues/56
  rec_id <- recipe(class ~ ., data = circle_example) |>
    update_role(id, new_role = "id") |>
    step_rose(class) |>
    prep() |>
    bake(new_data = NULL)

  expect_equal(is.double(rec_id$x), TRUE)
  expect_equal(is.double(rec_id$y), TRUE)
})

test_that("tunable", {
  rec <- recipe(~., data = mtcars) |>
    step_rose(all_predictors())
  rec_param <- tunable.step_rose(rec$steps[[1]])
  expect_equal(rec_param$name, c("over_ratio"))
  expect_true(all(rec_param$source == "recipe"))
  expect_true(is.list(rec_param$call_info))
  expect_equal(nrow(rec_param), 1)
  expect_equal(
    names(rec_param),
    c("name", "call_info", "source", "component", "component_id")
  )
})

test_that("indicator_column marks all rows TRUE (ROSE generates fully synthetic data)", {
  rec <- recipe(class ~ x + y, data = circle_example) |>
    step_rose(class, indicator_column = ".new_row") |>
    prep()

  res <- bake(rec, new_data = NULL)

  expect_true(".new_row" %in% names(res))
  expect_type(res$.new_row, "logical")
  expect_true(all(res$.new_row))
})

test_that("indicator_column bad args", {
  expect_snapshot(
    error = TRUE,
    recipe(class ~ x + y, data = circle_example) |>
      step_rose(class, indicator_column = 1)
  )
  expect_snapshot(
    error = TRUE,
    recipe(class ~ x + y, data = circle_example) |>
      step_rose(class, indicator_column = "") |>
      prep()
  )
  expect_snapshot(
    error = TRUE,
    recipe(class ~ x + y, data = circle_example) |>
      step_rose(class, indicator_column = "x") |>
      prep()
  )
})

test_that("bad args", {
  expect_snapshot(
    error = TRUE,
    recipe(~., data = mtcars) |>
      step_rose(over_ratio = "yes") |>
      prep()
  )
  expect_snapshot(
    error = TRUE,
    recipe(~., data = mtcars) |>
      step_rose(minority_prop = TRUE)
  )
  expect_snapshot(
    error = TRUE,
    recipe(~., data = mtcars) |>
      step_rose(minority_prop = 1.5)
  )
  expect_snapshot(
    error = TRUE,
    recipe(~., data = mtcars) |>
      step_rose(minority_smoothness = TRUE)
  )
  expect_snapshot(
    error = TRUE,
    recipe(~., data = mtcars) |>
      step_rose(majority_smoothness = TRUE)
  )
  expect_snapshot(
    error = TRUE,
    recipe(~., data = mtcars) |>
      step_rose(seed = TRUE)
  )
})


test_that("tunable is setup to works with extract_parameter_set_dials", {
  skip_if_not_installed("dials")
  rec <- recipe(~., data = mtcars) |>
    step_rose(
      all_predictors(),
      over_ratio = hardhat::tune()
    )

  params <- extract_parameter_set_dials(rec)

  expect_s3_class(params, "parameters")
  expect_identical(nrow(params), 1L)
})

test_that("rose() basic usage", {
  circle_numeric <- circle_example[, c("x", "y", "class")]

  res <- rose(circle_numeric, var = "class")
  expect_s3_class(res, "data.frame")
  expect_named(res, c("x", "y", "class"))
  expect_s3_class(res$class, "factor")
})

test_that("rose() preserves factor levels", {
  circle_numeric <- circle_example[, c("x", "y", "class")]
  original_levels <- levels(circle_numeric$class)

  res <- rose(circle_numeric, var = "class")
  expect_equal(levels(res$class), original_levels)
})

test_that("rose() returns tibble when given tibble", {
  circle_tbl <- as_tibble(circle_example[, c("x", "y", "class")])

  res <- rose(circle_tbl, var = "class")
  expect_s3_class(res, "tbl_df")
})

test_that("rose() returns data.frame when given data.frame", {
  circle_df <- as.data.frame(circle_example[, c("x", "y", "class")])

  res <- rose(circle_df, var = "class")
  expect_s3_class(res, "data.frame")
  expect_false(inherits(res, "tbl_df"))
})

test_that("rose() bad args", {
  circle_numeric <- circle_example[, c("x", "y", "class")]

  expect_snapshot(
    error = TRUE,
    rose(matrix(), var = "class")
  )
  expect_snapshot(
    error = TRUE,
    rose(circle_numeric, var = c("class", "x"))
  )
  expect_snapshot(
    error = TRUE,
    rose(circle_numeric, var = "x")
  )
  expect_snapshot(
    error = TRUE,
    rose(circle_numeric, var = "class", over_ratio = TRUE)
  )
  expect_snapshot(
    error = TRUE,
    rose(circle_numeric, var = "class", minority_prop = TRUE)
  )
  expect_snapshot(
    error = TRUE,
    rose(circle_numeric, var = "class", minority_smoothness = TRUE)
  )
  expect_snapshot(
    error = TRUE,
    rose(circle_numeric, var = "class", majority_smoothness = TRUE)
  )
})

test_that("rose() errors on more than 2 class levels", {
  df <- data.frame(
    x = 1:9,
    class = factor(rep(c("a", "b", "c"), 3))
  )
  expect_snapshot(
    error = TRUE,
    rose(df, var = "class")
  )
})

test_that("unused outcome levels are skipped with a warning (#238)", {
  circle_example$class <- factor(
    circle_example$class,
    levels = c(levels(circle_example$class), "unused")
  )

  expect_snapshot(
    res <- recipe(class ~ x + y, data = circle_example) |>
      step_rose(class) |>
      prep() |>
      bake(new_data = NULL)
  )

  expect_gt(nrow(res), 0)
})

test_that("step_rose() rejects a named `over_ratio` vector (#323)", {
  expect_snapshot(
    error = TRUE,
    recipe(class ~ x + y, data = circle_example) |>
      step_rose(class, over_ratio = c(Circle = 1)) |>
      prep()
  )

  expect_snapshot(
    error = TRUE,
    rose(
      circle_example[c("x", "y", "class")],
      "class",
      over_ratio = c(Circle = 1)
    )
  )
})

test_that("backwards compatible for arguments added after 1.0.3", {
  rec <- recipe(class ~ x + y, data = circle_example) |>
    step_rose(class)

  exp <- bake(prep(rec), new_data = NULL)

  # simulates a recipe created by an older version of themis
  old <- rec
  old$steps[[1]]$indicator_column <- NULL

  expect_identical(bake(prep(old), new_data = NULL), exp)

  # simulates a recipe trained by an older version of themis
  old_trained <- prep(rec)
  old_trained$steps[[1]]$indicator_column <- NULL

  expect_identical(bake(old_trained, new_data = NULL), exp)
})

# Infrastructure ---------------------------------------------------------------

test_that("bake method errors when needed non-standard role columns are missing", {
  rec <- recipe(class ~ x + y, data = circle_example) |>
    step_rose(class, skip = FALSE) |>
    add_role(class, new_role = "potato") |>
    update_role_requirements(role = "potato", bake = FALSE)

  trained <- prep(rec, training = circle_example, verbose = FALSE)

  expect_snapshot(
    error = TRUE,
    bake(trained, new_data = circle_example[, -3])
  )
})

test_that("empty printing", {
  rec <- recipe(mpg ~ ., mtcars)
  rec <- step_rose(rec)

  expect_snapshot(rec)

  rec <- prep(rec, mtcars)

  expect_snapshot(rec)
})

test_that("empty selection prep/bake is a no-op", {
  rec1 <- recipe(mpg ~ ., mtcars)
  rec2 <- step_rose(rec1)

  rec1 <- prep(rec1, mtcars)
  rec2 <- prep(rec2, mtcars)

  baked1 <- bake(rec1, mtcars)
  baked2 <- bake(rec2, mtcars)

  expect_identical(baked1, baked2)
})

test_that("empty selection tidy method works", {
  rec <- recipe(mpg ~ ., mtcars)
  rec <- step_rose(rec)

  expect <- tibble(terms = character(), id = character())

  expect_identical(tidy(rec, number = 1), expect)

  rec <- prep(rec, mtcars)

  expect_identical(tidy(rec, number = 1), expect)
})

test_that("printing", {
  rec <- recipe(class ~ x + y, data = circle_example) |>
    step_rose(class)

  expect_snapshot(print(rec))
  expect_snapshot(prep(rec))
})

test_that("0 and 1 rows data work in bake method", {
  rec <- recipe(class ~ x + y, data = circle_example) |>
    step_rose(class, skip = FALSE) |>
    prep()

  expect_identical(nrow(bake(rec, new_data = slice(circle_example, 0))), 0L)
  expect_identical(nrow(bake(rec, new_data = slice(circle_example, 1))), 1L)
})

Try the themis package in your browser

Any scripts or data that you put into this service are public.

themis documentation built on Aug. 2, 2026, 9:07 a.m.