## ----setup, include = FALSE---------------------------------------------------
fixture_dir <- "embeddings"
recording <- nzchar(Sys.getenv("FOUNDRY_RECORD_DOCS"))
have_fixtures <- dir.exists(fixture_dir) && length(list.files(fixture_dir)) > 0
run_api <- requireNamespace("httptest2", quietly = TRUE) &&
  (recording || have_fixtures)
library(foundryR)
if (run_api) {
  httptest2::start_vignette(fixture_dir)
}
knitr::opts_chunk$set(collapse = TRUE, comment = "#>", eval = run_api,
  fig.width = 7, fig.height = 4.5, out.width = "100%")

## ----libraries, message = FALSE-----------------------------------------------
library(foundryR)
library(dplyr)

## ----austen-lines-------------------------------------------------------------
austen_lines <- c(
  "It is a truth universally acknowledged, that a single man in possession of a good fortune, must be in want of a wife.",
  "However little known the feelings or views of such a man may be on his first entering a neighbourhood.",
  "Mr. Bennet was so odd a mixture of quick parts, sarcastic humour, reserve, and caprice."
)

embedding <- foundry_embed(austen_lines[1], model = "text-embedding-3-small")
embedding[, c("text", "n_dims", ".input_idx", ".error", ".error_msg")]

## ----multiple-embed-----------------------------------------------------------
doc_embeddings <- foundry_embed(austen_lines, model = "text-embedding-3-small")
doc_embeddings[, c("text", "n_dims", ".input_idx", ".error")]

## ----reduced-dims-------------------------------------------------------------
compact <- foundry_embed(
  austen_lines[1],
  model = "text-embedding-3-small",
  dimensions = 256
)
compact[, c("text", "n_dims")]

## ----similarity---------------------------------------------------------------
mixed <- c(
  "It is a truth universally acknowledged, that a single man in possession of a good fortune, must be in want of a wife.",
  "Mr. Bennet was so odd a mixture of quick parts, sarcastic humour, reserve, and caprice.",
  "The quarterly revenue report showed a sharp rise in cloud subscriptions.",
  "Analysts raised their earnings forecast after the strong cloud numbers."
)

similarities <- foundry_embed(mixed, model = "text-embedding-3-small") |>
  foundry_similarity()
similarities

## ----similarity-summary, echo = FALSE, results = "asis"-----------------------
pair_score <- function(a, b) {
  hit <- (similarities$text_1 == a & similarities$text_2 == b) |
    (similarities$text_1 == b & similarities$text_2 == a)
  similarities$similarity[hit][[1]]
}
finance_pair <- pair_score(mixed[3], mixed[4])
austen_pair <- pair_score(mixed[1], mixed[2])
cross_max <- max(setdiff(similarities$similarity, c(finance_pair, austen_pair)))
cat(sprintf(
  "In this recording the two finance sentences score %.2f and the two Austen lines %.2f, while no cross-source pair scores above %.2f. The Austen lines have no content words in common, so their score comes from meaning and style rather than shared vocabulary.\n",
  finance_pair, austen_pair, cross_max
))

## ----similarity-pairs-data, include = FALSE, eval = run_api && requireNamespace("ggplot2", quietly = TRUE)----
label_for <- function(x) {
  dplyr::case_when(
    startsWith(x, "It is a truth") ~ "Austen 1",
    startsWith(x, "Mr. Bennet") ~ "Austen 2",
    startsWith(x, "The quarterly") ~ "Finance 1",
    TRUE ~ "Finance 2"
  )
}

pair_data <- similarities |>
  dplyr::mutate(
    first = pmin(label_for(text_1), label_for(text_2)),
    second = pmax(label_for(text_1), label_for(text_2)),
    pair = paste(first, "and", second),
    source_match = ifelse(
      sub(" .*$", "", first) == sub(" .*$", "", second),
      "same",
      "different"
    )
  ) |>
  dplyr::arrange(similarity)
stopifnot(!anyDuplicated(pair_data$pair))
pair_data$pair <- factor(pair_data$pair, levels = pair_data$pair)

same_scores <- pair_data$similarity[pair_data$source_match == "same"]
cross_scores <- pair_data$similarity[pair_data$source_match == "different"]
similarity_chart_title <- if (min(same_scores) > max(cross_scores)) {
  "Same-source pairs score above all cross-source pairs"
} else {
  "Same-source pairs tend to score above cross-source pairs"
}
wrap_text <- function(lines, width) {
  wrapped <- vapply(
    lines,
    function(line) paste(strwrap(line, width), collapse = "\n"),
    character(1)
  )
  paste(wrapped, collapse = "\n")
}
similarity_chart_alt <- paste0(
  similarity_chart_title,
  ". Horizontal bar chart of cosine similarity for six sentence pairs: ",
  paste(
    sprintf("%s %.2f", rev(as.character(pair_data$pair)), rev(pair_data$similarity)),
    collapse = "; "
  ),
  "."
)

## ----similarity-pairs, echo = FALSE, eval = run_api && requireNamespace("ggplot2", quietly = TRUE), fig.alt = get0("similarity_chart_alt", ifnotfound = "Bar chart of cosine similarity for sentence pairs from two sources.")----
ggplot2::ggplot(
  pair_data,
  ggplot2::aes(x = similarity, y = pair, fill = source_match)
) +
  ggplot2::geom_col(width = 0.66) +
  ggplot2::geom_text(
    ggplot2::aes(label = sprintf("%.2f", similarity), color = source_match),
    hjust = -0.25,
    size = 3.6
  ) +
  ggplot2::scale_fill_manual(
    values = c(same = "#0072B2", different = "#A6A6A6"),
    guide = "none"
  ) +
  ggplot2::scale_color_manual(
    values = c(same = "#005A8C", different = "#595959"),
    guide = "none"
  ) +
  ggplot2::scale_x_continuous(
    breaks = seq(0, 1, by = 0.25),
    expand = ggplot2::expansion(mult = c(0, 0.02))
  ) +
  ggplot2::coord_cartesian(xlim = c(min(0, pair_data$similarity), 1)) +
  ggplot2::labs(
    title = wrap_text(similarity_chart_title, 62),
    subtitle = wrap_text(
      "Cosine similarity for all six pairs of two Austen lines and two finance sentences",
      84
    ),
    x = "Cosine similarity",
    y = NULL,
    caption = wrap_text(
      c(
        "Blue bars pair sentences from the same source.",
        "Scores come from foundry_similarity() on text-embedding-3-small embeddings."
      ),
      101
    )
  ) +
  ggplot2::theme_minimal(base_size = 12) +
  ggplot2::theme(
    plot.title = ggplot2::element_text(face = "bold"),
    plot.title.position = "plot",
    plot.subtitle = ggplot2::element_text(color = "grey30"),
    plot.caption = ggplot2::element_text(color = "grey40", hjust = 0),
    plot.caption.position = "plot",
    panel.grid.major.y = ggplot2::element_blank(),
    panel.grid.minor = ggplot2::element_blank(),
    axis.text.y = ggplot2::element_text(color = "grey15")
  )

## ----semantic-search----------------------------------------------------------
documents <- c(
  "How to install R packages using install.packages()",
  "Data visualization with ggplot2 in R",
  "Introduction to machine learning with Python",
  "Statistical hypothesis testing explained",
  "Building web applications with Shiny",
  "Deep learning with TensorFlow and Keras"
)
query <- "How do I create charts and graphs in R?"

search_embeddings <- foundry_embed_batch(
  c(query, documents),
  model = "text-embedding-3-small",
  batch_size = 4,
  max_active = 2
) |>
  mutate(label = c("query", paste0("doc_", seq_along(documents))))

search_pairs <- foundry_similarity(search_embeddings, text_col = "label")

search_pairs |>
  filter(text_1 == "query" | text_2 == "query") |>
  mutate(
    label = if_else(text_1 == "query", text_2, text_1),
    text = documents[match(label, paste0("doc_", seq_along(documents)))]
  ) |>
  select(text, similarity) |>
  arrange(desc(similarity)) |>
  slice_head(n = 3)

## ----clustering---------------------------------------------------------------
texts <- c(
  "Python is great for machine learning",
  "R excels at statistical analysis",
  "JavaScript powers modern web applications",
  "Italian pasta with tomato sauce",
  "Sushi is a popular Japanese dish",
  "French croissants are flaky and buttery",
  "Soccer is the world's most popular sport",
  "Basketball requires speed and agility",
  "Tennis matches can last for hours"
)

cluster_embeddings <- foundry_embed(texts, model = "text-embedding-3-small")
embedding_matrix <- do.call(rbind, cluster_embeddings$embedding)

set.seed(42)
clusters <- kmeans(embedding_matrix, centers = 3, nstart = 10)

cluster_embeddings |>
  mutate(cluster = clusters$cluster) |>
  select(text, cluster) |>
  arrange(cluster)

## ----projection-data, include = FALSE, eval = run_api && requireNamespace("ggplot2", quietly = TRUE)----
topics <- rep(c("Programming", "Food", "Sports"), each = 3)
stopifnot(length(topics) == length(texts))

pca <- stats::prcomp(embedding_matrix, rank. = 2)
projection <- tibble::tibble(
  pc1 = pca$x[, 1],
  pc2 = pca$x[, 2],
  cluster = factor(clusters$cluster),
  topic = topics
)

cluster_labels <- projection |>
  dplyr::group_by(cluster) |>
  dplyr::summarise(
    n_points = dplyr::n(),
    top_count = max(table(topic)),
    label_text = names(which.max(table(topic))),
    label_x = mean(pc1),
    label_y = max(pc2),
    .groups = "drop"
  )
label_offset <- 0.08 * diff(range(projection$pc2))

cluster_colors <- c("#0072B2", "#D55E00", "#009E73")
cluster_text <- c("#005A8C", "#A84A0A", "#00735A")
cluster_shapes <- c(16, 17, 15)

clusters_are_pure <- all(cluster_labels$top_count == cluster_labels$n_points) &&
  dplyr::n_distinct(cluster_labels$label_text) == nrow(cluster_labels)
projection_title <- if (clusters_are_pure) {
  "k-means separates the three topics in embedding space"
} else {
  "k-means clusters mostly follow the three topics"
}
projection_alt <- paste0(
  projection_title,
  ". Scatter plot of nine sentence embeddings on their first two principal components; ",
  paste(
    sprintf(
      "cluster %s holds %d of %d %s sentences",
      cluster_labels$cluster,
      cluster_labels$top_count,
      cluster_labels$n_points,
      tolower(cluster_labels$label_text)
    ),
    collapse = "; "
  ),
  "."
)

## ----projection, echo = FALSE, eval = run_api && requireNamespace("ggplot2", quietly = TRUE), fig.alt = get0("projection_alt", ifnotfound = "Scatter plot of sentence embeddings projected onto two principal components.")----
ggplot2::ggplot(
  projection,
  ggplot2::aes(pc1, pc2, color = cluster, shape = cluster)
) +
  ggplot2::geom_point(size = 3.4) +
  ggplot2::geom_text(
    data = cluster_labels,
    ggplot2::aes(x = label_x, y = label_y + label_offset, label = label_text),
    inherit.aes = FALSE,
    color = cluster_text[as.integer(cluster_labels$cluster)],
    fontface = "bold",
    size = 4.2,
    vjust = 0
  ) +
  ggplot2::scale_color_manual(values = cluster_colors, guide = "none") +
  ggplot2::scale_shape_manual(values = cluster_shapes, guide = "none") +
  ggplot2::scale_x_continuous(expand = ggplot2::expansion(mult = 0.12)) +
  ggplot2::scale_y_continuous(expand = ggplot2::expansion(mult = c(0.08, 0.22))) +
  ggplot2::coord_cartesian(clip = "off") +
  ggplot2::labs(
    title = wrap_text(projection_title, 62),
    subtitle = wrap_text(
      c(
        "Nine sentences on the first two principal components of their embeddings",
        "Color and shape show the k-means cluster with k = 3"
      ),
      84
    ),
    x = "Principal component 1",
    y = "Principal component 2",
    caption = wrap_text(
      c(
        "Embeddings come from text-embedding-3-small via foundry_embed().",
        "Axis positions are relative, so tick values are omitted."
      ),
      101
    )
  ) +
  ggplot2::theme_minimal(base_size = 12) +
  ggplot2::theme(
    plot.title = ggplot2::element_text(face = "bold"),
    plot.title.position = "plot",
    plot.subtitle = ggplot2::element_text(color = "grey30"),
    plot.caption = ggplot2::element_text(color = "grey40", hjust = 0),
    plot.caption.position = "plot",
    axis.text = ggplot2::element_blank(),
    panel.grid.minor = ggplot2::element_blank()
  )

## ----batch-example------------------------------------------------------------
many_texts <- c(
  "The login page rejects my password reset link.",
  "The invoice total does not match the purchase order.",
  "The chart export button is missing from the report.",
  "The password reset email arrived after it expired.",
  "The billing address changed but the invoice did not update.",
  "The dashboard chart uses the wrong date range."
)

many_embeddings <- foundry_embed_batch(
  many_texts,
  model = "text-embedding-3-small",
  batch_size = 3,
  max_active = 2
)

many_embeddings[, c("text", "n_dims", ".error", ".error_msg")]

## ----cleanup, include = FALSE-------------------------------------------------
if (run_api) {
  httptest2::end_vignette()
}

