## ----setup, include = FALSE---------------------------------------------------
knitr::opts_chunk$set(collapse = TRUE, comment = "#>")

## -----------------------------------------------------------------------------
library(textclassificationtutorial)

text <- c(
  "analyze customer data", "build statistical models",
  "create predictive analytics", "report business metrics",
  "provide patient care", "support clinical treatment",
  "coordinate nursing care", "assist hospital patients",
  "analyze experimental results", "develop data dashboard",
  "monitor patient health", "coordinate clinical team"
)
label <- rep(c("data", "care"), each = 4)
label <- c(label, "data", "data", "care", "care")

## -----------------------------------------------------------------------------
folds <- stratified_folds(label, k = 3, seed = 2026)
folds

## -----------------------------------------------------------------------------
fold_results <- lapply(folds, function(test_index) {
  train_index <- setdiff(seq_along(text), test_index)

  train_clean <- preprocess_text(text[train_index])
  test_clean <- preprocess_text(text[test_index])

  train_dtm <- document_term_matrix(train_clean)
  test_dtm <- document_term_matrix(test_clean)

  model <- fit_naive_bayes(train_dtm, label[train_index])
  estimate <- predict(model, test_dtm)

  classification_metrics(
    truth = label[test_index],
    estimate = estimate,
    positive = "data"
  )
})

results <- do.call(rbind, fold_results)
results
colMeans(results[c(
  "accuracy", "balanced_accuracy", "precision", "recall", "f1"
)], na.rm = TRUE)

