## ----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

