## ----setup, include=FALSE----------------------------------------------------- knitr::opts_chunk$set( collapse = TRUE, comment = "#>", fig.align = "center", fig.width = 6, fig.height = 5, message = FALSE, warning = FALSE ) library(FITclust) library(ggplot2) ## ----eval = FALSE------------------------------------------------------------- # library(devtools) # install_github("ghashti-j/FITclust") # library(FITclust) ## ----------------------------------------------------------------------------- set.seed(42) demoData <- rbind( data.frame(x1 = rnorm(100, -3, 1), x2 = rnorm(100, -3 - 0.25, 1), cluster = 1L, group = 0L), data.frame(x1 = rnorm(200, -3, 1), x2 = rnorm(200, -3 + 0.25, 1), cluster = 1L, group = 1L), data.frame(x1 = rnorm(200, 3, 1), x2 = rnorm(200, 3 - 0.25, 1), cluster = 2L, group = 0L), data.frame(x1 = rnorm(100, 3, 1), x2 = rnorm(100, 3 + 0.25, 1), cluster = 2L, group = 1L) ) dataMat <- as.matrix(demoData[, c("x1", "x2")]) groupVec <- demoData$group trueCluster <- demoData$cluster cat("n =", nrow(dataMat), " group counts =", paste(table(groupVec), collapse = "/"), " cluster counts =", paste(table(trueCluster), collapse = "/"), "\n") ## ----fig.align='center'------------------------------------------------------- ggplot(demoData, aes(x1, x2, shape = factor(group), fill = factor(group))) + geom_point(size = 2, colour = "black", stroke = 0.3, alpha = 0.7) + scale_shape_manual("Group", values = c("0" = 21, "1" = 24)) + scale_fill_manual("Group", values = c("0" = "#8ABF69", "1" = "#D08890")) + labs(x = expression(x[1]), y = expression(x[2])) + coord_fixed() + theme_bw() + theme(panel.grid = element_blank(), legend.position = "bottom") ## ----------------------------------------------------------------------------- alphaVec <- resolveAlpha("uniform", groupVec, sort(unique(groupVec))) transport <- buildTransport(dataMat, groupVec, alphaVec, verbose = FALSE) cat("barycenter atoms =", nrow(transport$barycenter), " converged =", transport$baryConverged, " iterations =", transport$baryIter, "\n") ## ----------------------------------------------------------------------------- baseFit <- fcm(dataMat, numClusters = 2, numStart = 5) fullFit <- fcm(transport$fn(1), numClusters = 2, numStart = 5) cat("Delta_soft at t = 0:", round(softViolation(baseFit$membership, groupVec), 3), "\n") cat("Delta_soft at t = 1:", round(softViolation(fullFit$membership, groupVec), 3), "\n") ## ----------------------------------------------------------------------------- set.seed(1) fitCentroid <- fitSKM(dataMat, groupVec, numClusters = 2, deltaFair = 0.05, tSeq = seq(0, 1, by = 0.02), verbose = FALSE) cat("t* =", fitCentroid$tOptimal, " Delta_soft:", round(fitCentroid$violationBaseline, 3), "->", round(fitCentroid$violation, 3), "\n") ## ----fig.align='center'------------------------------------------------------- hist <- fitCentroid$history ggplot(hist, aes(t, violationSoft)) + geom_line() + geom_point(size = 1) + geom_hline(yintercept = 0.05, linetype = "dashed", colour = "#E31A1C") + geom_vline(xintercept = fitCentroid$tOptimal, linetype = "dotted", colour = "grey30") + labs(x = expression(t), y = expression(Delta[soft](t))) + theme_bw() + theme(panel.grid = element_blank()) ## ----------------------------------------------------------------------------- alignLabels <- function(current, reference) { overlap <- table(current, reference) mapping <- apply(overlap, 1, which.max) as.integer(mapping[as.character(current)]) } baseLabels <- fitCentroid$clustersBaseline fairLabels <- alignLabels(fitCentroid$clusters, baseLabels) cat("reassigned:", sum(fairLabels != baseLabels), "of", length(baseLabels), sprintf("(%.1f%%)", 100 * mean(fairLabels != baseLabels)), "\n") ## ----fig.align='center'------------------------------------------------------- plotDF <- data.frame(x1 = dataMat[, 1], x2 = dataMat[, 2], cluster = factor(fairLabels), group = factor(groupVec)) ggplot(plotDF, aes(x1, x2, shape = group, fill = cluster)) + geom_point(size = 2, colour = "black", stroke = 0.3) + scale_shape_manual("Group", values = c("0" = 21, "1" = 24)) + scale_fill_manual("Cluster", values = c("1" = "#4E9BC7", "2" = "#F4A460")) + labs(x = expression(x[1]), y = expression(x[2])) + coord_fixed() + theme_bw() + theme(panel.grid = element_blank(), legend.position = "bottom") + guides(fill = guide_legend(override.aes = list(shape = 22)), shape = guide_legend(override.aes = list(fill = "grey60"))) ## ----------------------------------------------------------------------------- set.seed(1) fitGraph <- fitSSC(dataMat, groupVec, numClusters = 2, deltaFair = 0.05, tSeq = seq(0, 1, by = 0.02), verbose = FALSE) fitModel <- fitSMM(dataMat, groupVec, numClusters = 2, deltaFair = 0.05, tSeq = seq(0, 1, by = 0.02), verbose = FALSE) summaryTab <- data.frame( Family = c("Centroid (fitSKM)", "Graph (fitSSC)", "Model (fitSMM)"), tOptimal = c(fitCentroid$tOptimal, fitGraph$tOptimal, fitModel$tOptimal), DeltaSoftBaseline = round(c(fitCentroid$violationBaseline, fitGraph$violationBaseline, fitModel$violationBaseline), 3), DeltaSoftFair = round(c(fitCentroid$violation, fitGraph$violation, fitModel$violation), 3) ) summaryTab