Skip to content

Commit 97ae005

Browse files
add hdbscan engine (#247)
* add hdbscan engine to `db_clust()` (#238) * test hdbscan engine * document hdbscan engine * add NEWS bullet for hdbscan engine (#238)
1 parent 2fa6d73 commit 97ae005

16 files changed

Lines changed: 611 additions & 0 deletions

‎NAMESPACE‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@ S3method(extract_cluster_assignment,cluster_fit)
1414
S3method(extract_cluster_assignment,cluster_spec)
1515
S3method(extract_cluster_assignment,dbscan)
1616
S3method(extract_cluster_assignment,hclust)
17+
S3method(extract_cluster_assignment,hdbscan)
1718
S3method(extract_cluster_assignment,kmeans)
1819
S3method(extract_cluster_assignment,kmodes)
1920
S3method(extract_cluster_assignment,kproto)
@@ -27,6 +28,7 @@ S3method(extract_fit_summary,cluster_fit)
2728
S3method(extract_fit_summary,cluster_spec)
2829
S3method(extract_fit_summary,dbscan)
2930
S3method(extract_fit_summary,hclust)
31+
S3method(extract_fit_summary,hdbscan)
3032
S3method(extract_fit_summary,kmeans)
3133
S3method(extract_fit_summary,kmodes)
3234
S3method(extract_fit_summary,kproto)
@@ -93,6 +95,7 @@ S3method(update,k_means)
9395
S3method(update,mean_shift)
9496
export("%>%")
9597
export(.db_clust_fit_dbscan)
98+
export(.db_clust_fit_hdbscan)
9699
export(.gm_clust_fit_mclust)
97100
export(.hier_clust_fit_stats)
98101
export(.k_means_fit_ClusterR)

‎NEWS.md‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,8 @@
2828

2929
* `mean_shift()` gains a new engine with `meanShiftR`. (#244)
3030

31+
* `db_clust()` gains a new engine with `hdbscan`, using `dbscan::hdbscan()` to fit HDBSCAN models. (#238)
32+
3133
* The `.config` column produced by `tune_cluster()` has changed from the
3234
`Preprocessor{num}_Model{num}` pattern to `pre{num}_mod{num}_post{num}` to
3335
align with updates in the tune package. (#220)

‎R/db_clust.R‎

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
#' are listed below.
1111
#'
1212
#' - \link[=details_db_clust_dbscan]{dbscan}
13+
#' - \link[=details_db_clust_hdbscan]{hdbscan}
1314
#'
1415
#' @param mode A single character string for the type of model. The only
1516
#' possible value for this model is `"partition"`.
@@ -263,3 +264,39 @@ dbscan_helper <- function(object, ...) {
263264
)
264265
training_data$cluster[order(training_data$overall_order)]
265266
}
267+
268+
#' Simple Wrapper around hdbscan function
269+
#'
270+
#' This wrapper passes the data to `dbscan::hdbscan()` and stashes the training
271+
#' data on the result so it can be reused for prediction and extraction.
272+
#'
273+
#' @param x matrix or data frame.
274+
#' @param min_points Minimum cluster size used as the `minPts` argument of
275+
#' `dbscan::hdbscan()`.
276+
#' @param min_cluster_size Engine-specific override for `minPts`. When supplied,
277+
#' it is used in place of `min_points`.
278+
#'
279+
#' @return hdbscan object
280+
#' @keywords internal
281+
#' @export
282+
.db_clust_fit_hdbscan <- function(
283+
x,
284+
min_points = NULL,
285+
min_cluster_size = NULL,
286+
...
287+
) {
288+
min_pts <- min_cluster_size %||% min_points
289+
290+
if (is.null(min_pts)) {
291+
cli::cli_abort(
292+
"Please specify `min_points` to be able to fit specification.",
293+
call = call("fit")
294+
)
295+
}
296+
297+
res <- dbscan::hdbscan(x, minPts = min_pts, ...)
298+
attr(res, "min_points") <- min_pts
299+
attr(res, "training_data") <- x
300+
301+
res
302+
}

‎R/db_clust_data.R‎

Lines changed: 65 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -78,6 +78,71 @@ make_db_clust <- function() {
7878
)
7979
)
8080
)
81+
82+
# ----------------------------------------------------------------------------
83+
84+
modelenv::set_model_engine("db_clust", "partition", "hdbscan")
85+
modelenv::set_dependency(
86+
model = "db_clust",
87+
mode = "partition",
88+
eng = "hdbscan",
89+
pkg = "dbscan"
90+
)
91+
modelenv::set_dependency(
92+
model = "db_clust",
93+
mode = "partition",
94+
eng = "hdbscan",
95+
pkg = "tidyclust"
96+
)
97+
98+
modelenv::set_fit(
99+
model = "db_clust",
100+
eng = "hdbscan",
101+
mode = "partition",
102+
value = list(
103+
interface = "matrix",
104+
protect = c("x", "min_points"),
105+
func = c(pkg = "tidyclust", fun = ".db_clust_fit_hdbscan"),
106+
defaults = list()
107+
)
108+
)
109+
110+
modelenv::set_encoding(
111+
model = "db_clust",
112+
eng = "hdbscan",
113+
mode = "partition",
114+
options = list(
115+
predictor_indicators = "traditional",
116+
compute_intercept = TRUE,
117+
remove_intercept = TRUE,
118+
allow_sparse_x = FALSE
119+
)
120+
)
121+
122+
modelenv::set_model_arg(
123+
model = "db_clust",
124+
eng = "hdbscan",
125+
exposed = "min_points",
126+
original = "min_points",
127+
func = list(pkg = "dials", fun = "min_points"),
128+
has_submodel = TRUE
129+
)
130+
131+
modelenv::set_pred(
132+
model = "db_clust",
133+
eng = "hdbscan",
134+
mode = "partition",
135+
type = "cluster",
136+
value = list(
137+
pre = NULL,
138+
post = NULL,
139+
func = c(fun = ".db_clust_predict_hdbscan"),
140+
args = list(
141+
object = rlang::expr(object$fit),
142+
new_data = rlang::expr(new_data)
143+
)
144+
)
145+
)
81146
}
82147

83148
# nocov end

‎R/db_clust_hdbscan.R‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,11 @@
1+
#' Hierarchical Density-Based Spatial Clustering (HDBSCAN) via dbscan
2+
#'
3+
#' [db_clust()] creates an HDBSCAN model.
4+
#'
5+
#' @includeRmd man/rmd/db_clust_hdbscan.md details
6+
#'
7+
#' @name details_db_clust_hdbscan
8+
#' @keywords internal
9+
NULL
10+
11+
# See inst/README-DOCS.md for a description of how these files are processed

‎R/extract_cluster_assignment.R‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -172,6 +172,13 @@ extract_cluster_assignment.dbscan <- function(object, ...) {
172172
cluster_assignment_tibble_w_outliers(clusters, n_clusters, ...)
173173
}
174174

175+
#' @export
176+
extract_cluster_assignment.hdbscan <- function(object, ...) {
177+
clusters <- object$cluster
178+
n_clusters <- length(unique(clusters[clusters != 0])) + 1
179+
cluster_assignment_tibble_w_outliers(clusters, n_clusters, ...)
180+
}
181+
175182
#' @export
176183
extract_cluster_assignment.ms <- function(object, ...) {
177184
n_clusters <- nrow(object$cluster.center)

‎R/extract_fit_summary.R‎

Lines changed: 64 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -275,6 +275,70 @@ extract_fit_summary.dbscan <- function(object, ...) {
275275
summary
276276
}
277277

278+
#' @export
279+
extract_fit_summary.hdbscan <- function(object, ...) {
280+
clusts <- extract_cluster_assignment(object, ...)$.cluster
281+
n_clust <- dplyr::n_distinct(clusts)
282+
training_data <- attr(object, "training_data")
283+
284+
overall_centroid <- colMeans(training_data)
285+
286+
by_clust <- training_data %>%
287+
tibble::as_tibble() %>%
288+
dplyr::mutate(
289+
.cluster = clusts
290+
) %>%
291+
dplyr::group_by(.cluster) %>%
292+
tidyr::nest()
293+
294+
centroids <- by_clust$data %>%
295+
map(dplyr::summarize_all, mean) %>%
296+
dplyr::bind_rows()
297+
298+
outlier_idx <- which(unique(clusts) == "Outlier")
299+
300+
sse_within_total_total <- map2_dbl(
301+
by_clust$data,
302+
seq_len(n_clust),
303+
~ sum(
304+
philentropy::dist_many_many(
305+
as.matrix(centroids[.y, ]),
306+
as.matrix(.x),
307+
method = "euclidean"
308+
)
309+
)
310+
)
311+
312+
clust_names <- unique(clusts)[c(outlier_idx, setdiff(1:n_clust, outlier_idx))]
313+
summary <- list(
314+
cluster_names = clust_names,
315+
centroids = centroids,
316+
n_members = unname(as.integer(table(clusts))),
317+
sse_within_total_total = sse_within_total_total,
318+
sse_total = sum(
319+
philentropy::dist_many_many(
320+
t(overall_centroid),
321+
as.matrix(training_data),
322+
method = "euclidean"
323+
)
324+
),
325+
orig_labels = NULL,
326+
cluster_assignments = clusts
327+
)
328+
329+
if (length(outlier_idx) > 0) {
330+
summary$centroids[outlier_idx, ] <- rep(NA, ncol(summary$centroids))
331+
summary$sse_within_total_total[outlier_idx] <- NA
332+
}
333+
334+
# reorder centroids
335+
summary$centroids <- summary$centroids[
336+
c(outlier_idx, setdiff(1:n_clust, outlier_idx)),
337+
]
338+
339+
summary
340+
}
341+
278342
#' @export
279343
extract_fit_summary.ms <- function(
280344
object,

‎R/predict_helpers.R‎

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -251,6 +251,21 @@ make_predictions_w_outliers <- function(x, prefix, n_clusters, labels = NULL) {
251251
make_predictions_w_outliers(clusters, prefix, n_clusters, labels)
252252
}
253253

254+
.db_clust_predict_hdbscan <- function(
255+
object,
256+
new_data,
257+
prefix = "Cluster_",
258+
labels = NULL
259+
) {
260+
training_data <- attr(object, "training_data")
261+
clusters <- object$cluster
262+
n_clusters <- length(unique(clusters[clusters != 0])) + 1
263+
264+
preds <- stats::predict(object, newdata = new_data, data = training_data)
265+
266+
make_predictions_w_outliers(preds, prefix, n_clusters, labels)
267+
}
268+
254269
.mean_shift_predict_LPCM <- function(
255270
object,
256271
new_data,

‎R/tunable.R‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -100,6 +100,9 @@ tunable.db_clust <- function(x, ...) {
100100
if (x$engine == "dbscan") {
101101
res <- add_engine_parameters(res, dbscan_db_clust_engine_args)
102102
}
103+
if (x$engine == "hdbscan") {
104+
res <- add_engine_parameters(res, hdbscan_db_clust_engine_args)
105+
}
103106
res
104107
}
105108

@@ -118,6 +121,17 @@ dbscan_db_clust_engine_args <-
118121
component_id = "engine"
119122
)
120123

124+
hdbscan_db_clust_engine_args <-
125+
tibble::tibble(
126+
name = "min_cluster_size",
127+
call_info = list(
128+
list(pkg = "dials", fun = "min_points")
129+
),
130+
source = "cluster_spec",
131+
component = "hdbscan",
132+
component_id = "engine"
133+
)
134+
121135

122136
#' @rdname tunable.cluster_spec
123137
#' @export

‎man/db_clust.Rd‎

Lines changed: 1 addition & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

0 commit comments

Comments
 (0)