rm(list = ls())
library(MASS) # Import function mvrnorm
library(rTensor) # Tensor library
library(pracma) # Import function sqrtm
library(PMA) # PMD method
source("SCCA-code.R") # SCCA method
source("GLAA_SVD.R") # GLAA algorithm
source("models.R") # model settings
source("utility.R") # auxiliary functions

# -------------------- Scenario 1 --------------------#
# It is time-consuming to reproduce the results in Table 1 based on 100 replicates. This code reproduces one replicate with p1 = p2 = 100, and p3 = 1. The ranks for the first two modes are r1 = r2 = 2. The sparsity levels for the first two modes are s1 = s2 = 5. And the sample size n = 500. Please refer to Section 5.1 for more detailed information. 

times <- 1 # The number of replicates

## Set random seed
RNGkind("L'Ecuyer-CMRG")
set.seed(123)

model <- Model1(pp = 100) # Set p1 = p2 = 100.
p <- model$p # The dimension p1 and p2.
r <- model$r # The ranks r1 and r2
s <- model$s # The sparsity level s1 and s2.
n <- model$n # The sample size.
sparse_mode <- model$sparse_mode # Specify the sparse modes.
a.list <- model$a.list # The tuning parameter
Gamma <- model$Gamma # The basis matrices Gamma1, Gamma2
data.gen <- model$data.gen # The function generating the simulated data in Scenario 1.

output <- sapply(seq_len(times), function(i){
  cat("Time", i, '\n')
  # ------------------------ Data generation ------------------------ #
  data <- data.gen(n) # Generate three training data sets
  x <- data$x
  y <- data$y
  z <- data$z
  # standardization
  x <- scale(x)
  y <- scale(y)
  z <- scale(z)

  data.test <- data.gen(n) # Generate three test data sets
  x.test <- data.test$x
  y.test <- data.test$y
  z.test <- data.test$z
  # standardization
  x.test <- scale(x.test)
  y.test <- scale(y.test)
  z.test <- scale(z.test)
  # ----------------------------------------------------------------- #

  # Compute the sample estimator Delta tilde for the test data set
  Delta.test <- tensor_mean(x.test, y.test, z.test)
  Delta.test <- as.tensor(Delta.test)

  # ----------------------- GLAA ------------------------- #
  fit.results <- lapply(a.list, function(a){ # Implement Algorithm 1 for each tuning parameter a.
    fit <- STATSVD.one(x, y, z, r, s, sparse_mode = sparse_mode, tmax = 50, a = a) # Sparse tensor decomposition algorithm (Algorithm 1)
    dist.true <- sapply(1:length(p), function(i){
      subspace(fit[[i]], Gamma[[i]])
    })
    true.dist <- mean(dist.true)  # The average subspace distance for the first two modes.
    s.list <- lapply(1:length(p), function(i){
      which(apply(fit[[i]], 1, function(x){any(x != 0)})) # The estimated active set for the first two modes.
    })
    
    # ----- validation ------ #
    proj <- list(fit[[1]] %*% t(fit[[1]]), fit[[2]] %*% t(fit[[2]]))
    Delta.test.2 <- ttl(Delta.test, proj, ms = c(1,2)) # The projection of Delta tilde from the test data set onto the two subspaces spanned by the estimated basis matrices.
    error <- rTensor::fnorm(Delta.test - Delta.test.2) # The error defined in (5)
    # ----------------------- #
    list(true.dist = true.dist, s.list = s.list, error = error)
  })
  
  true.dist <- do.call(c, lapply(fit.results, "[[", 1)) 
  s.list <- lapply(fit.results, "[[", 2) 
  error <- do.call(c, lapply(fit.results, "[[", 3)) 
  
  ind <- which.min(error) # Select the optimal tuning parameter.
  dist <- true.dist[ind] # The subspace estimation error corresponding to the optimal tuning parameter.
  s.list.final <- s.list[[ind]] # The selected active sets corresponding to the optimal tuning parameter.
  cat(paste0("The ", ind, "-th parameter is the optimal one. \n"))
  
  TFPR.list <- sapply(1:length(p), function(i){
    if((p[i] - s[i]) == 0){c(1,0)}
    else{
      TPR <- sum(s.list.final[[i]] %in% 1:s[i])/s[i] # The True Positive Rate on each mode
      FPR <- sum(s.list.final[[i]] %in% (s[i]+1):p[i])/(p[i] - s[i]) # The False Positive Rate on each mode
      c(TPR, FPR)
    }
  })
  cat("Set X: ", paste(s.list.final[[1]], collapse = " "), " | TPR(X):", TFPR.list[1,1], " | FPR(X):", TFPR.list[2,1], "\n",
      "Set Y: ", paste(s.list.final[[2]], collapse = " "), " | TPR(Y):", TFPR.list[1,2], " | FPR(Y):", TFPR.list[2,2], "\n",
      "Mean dist: ", dist, "\n\n", sep = "")
  GLAA.result <- c(c(TFPR.list), dist)
  # ---------------------------------- #
  
  # ----------------  Univariate LA -------------------- #
  ula <- tensor_mean(x, y ,z)
  ula <- as.tensor(ula) # Compute the sample estimator Delta tilde

  var.ula.list <- ula.variable(ula, p, s) # Variable selection for ULA: for each mode-k matricization, select rows with the first s_k largest l_2 norm.
  TFPR.ula <- sapply(1:length(p), function(i){ # The TPR and FPR for ULA
    ind <- var.ula.list[[i]]
    if((p[i] == s[i])){c(1,0)}
    else{
      TPR <- sum(ind %in% 1:s[i])/s[i]
      FPR <- sum(ind %in% (s[i]+1):p[i])/(p[i] - s[i])
      c(TPR, FPR)
    }
  })
  Gamma.ula <- lapply(1:length(p), function(k){ # Estimate the basis matrices for ULA: the first r_k left singular vectors of each mode-k matricization.
    Gamma <- svd(k_unfold(ula, k)@data)$u[,1:r[k], drop = FALSE]
    Gamma
  })
  dist.ula.list <- sapply(1:length(p), function(i){ 
    subspace(Gamma.ula[[i]], Gamma[[i]])
  })
  dist.ula <- mean(dist.ula.list) # The average subspace distance for ULA.
  ula.result <- c(c(TFPR.ula), dist.ula)
  # ----------------------------------------------- #


  # ---------------------- PMD ------------------------ #
  perm.out <- CCA.permute(x, y, typex="standard",typez="standard", trace = FALSE) # Automatically select tuning parameters for PMD method.
  pmd <- CCA(x, y, typex="standard", typez="standard", K=r[1],
                        penaltyx=perm.out$bestpenaltyx, penaltyz=perm.out$bestpenaltyz,
                        v=perm.out$v.init, trace = FALSE) # Implement PMD method using the selected tuning parameters and initialized v vectors.
  Gamma.pmd <- list(pmd$u, pmd$v) # The canonical directions from PMD
  TFPR.pmd <- sapply(1:2, function(i){ # The TPR and FPR for PMD
    item <- Gamma.pmd[[i]]
    ind <- which(apply(item, 1, function(x){any(x != 0)}))
    TPR <- sum(ind %in% 1:s[i])/s[i]
    FPR <- sum(ind %in% (s[i]+1):p[i])/(p[i] - s[i])
    c(TPR, FPR)
  })
  dist.pmd.list <- sapply(1:2, function(i){
    subspace(Gamma.pmd[[i]], Gamma[[i]])
  })
  dist.pmd <- mean(dist.pmd.list) # The average subspace distance for PMD.
  pmd.result <- c(c(TFPR.pmd), dist.pmd)
  # ----------------------------------------------------- #


  # ---------- SCCA ---------- #
  sigma.X.hat <- cov(x)
  sigma.Y.hat <- cov(y)
  sigma.YX.hat <- cov(y, x)
  lambda.scca <- (1:10)/100 # Tuning parameters for SCCA.

  obj.init<-init0(sigma.YX.hat=sigma.YX.hat, sigma.X.hat=sigma.X.hat, sigma.Y.hat=sigma.Y.hat, init.method='svd', npairs=r[1], n=n) # The initialization.
  obj.cv<-cv.SCCA.equal(x = x, y = y, alpha.init = obj.init$alpha.init, beta.init = obj.init$beta.init, lambda = lambda.scca) # Select the optimal tuning parameter of SCCA.
  bestlambda <- obj.cv$bestlambda # The optimal tuning parameter
  obj <- SCCA(x = x, y = y, lambda.alpha = rep(bestlambda, r[1]), lambda.beta = rep(bestlambda, r[1]), npairs = r[1], init.method = 'svd') # Implement SCCA.

  Gamma.scca <- list(obj$beta, obj$alpha) # The canonical directions from SCCA
  TFPR.scca <- sapply(1:2, function(i){ # The TPR and FPR for SCCA
    item <- Gamma.scca[[i]]
    ind <- which(apply(item, 1, function(x){any(x != 0)}))
    TPR <- sum(ind %in% 1:s[i])/s[i]
    FPR <- sum(ind %in% (s[i]+1):p[i])/(p[i] - s[i])
    c(TPR, FPR)
  })
  dist.scca.list <- sapply(1:2, function(i){
    subspace(Gamma.scca[[i]], Gamma[[i]])
  })
  dist.scca <- mean(dist.scca.list) # The average subspace distance for SCCA.
  scca.result <- c(c(TFPR.scca), dist.scca)
  # -------------------------------- #
  c(GLAA.result, ula.result, pmd.result, scca.result)
})

output <- as.data.frame(t(output))  # Record the output
colnames(output) <- c("GLAA_T1", "GLAA_F1", "GLAA_T2", "GLAA_F2", "GLAA_D", "ULA_T1", "ULA_F1", "ULA_T2", "ULA_F2", "ULA_D", "PMD_T1", "PMD_F1", "PMD_T2", "PMD_F2", "PMD_D", "SCCA_T1", "SCCA_F1", "SCCA_T2", "SCCA_F2", "SCCA_D")
print(output)
# --------------------------------------------------------#

# For reproducibility checking: the output based on one replicate is
# --------------------------------------------------------#
# > options(digits = 4)
# > output
# GLAA_T1 GLAA_F1 GLAA_T2 GLAA_F2  GLAA_D ULA_T1 ULA_F1 ULA_T2 ULA_F2  ULA_D PMD_T1 PMD_F1 PMD_T2 PMD_F2  PMD_D SCCA_T1 SCCA_F1 SCCA_T2 SCCA_F2
#       1       0       1       0 0.09477      1      0      1      0 0.7965      1 0.1368    0.2 0.2211 0.9595     0.4  0.2421       0  0.2632
# SCCA_D
# 0.9929
# --------------------------------------------------------#

