#!/usr/bin/env Rscript
suppressPackageStartupMessages({
  library(RUVSeq)
  library(edgeR)
  library(ggplot2)
  library(ggrepel)
})

# ============================================================
# 0. CONFIG
# ============================================================
outdir  <- "/BRC/yan/heart/snACAT_CM/DAR_results/"
infile  <- "/BRC/yan/heart/snACAT_CM/counts/count_matrix_snACAT_CM.txt"
k_range <- 1:6
k_final <- 6     # update after checking RLE/PCA plots

# ============================================================
# 1. LOAD DATA
# ============================================================
cat("[1] Loading count matrix...\n")
counts <- read.table(infile, header = TRUE, sep = "\t")
rownames(counts) <- paste0(counts$chr, ":", counts$start, "-", counts$end)
counts <- counts[, -(1:3)]
cat("    Peaks:", nrow(counts), " Samples:", ncol(counts), "\n")

# ============================================================
# 2. METADATA
# ============================================================
cat("[2] Building metadata...\n")
meta <- data.frame(
  sample = colnames(counts),
  sex    = ifelse(grepl("female", colnames(counts)), "female", "male"),
  row.names = colnames(counts)
)
x <- factor(meta$sex, levels = c("male", "female"))
cat("    female:", sum(meta$sex == "female"),
    "  male:", sum(meta$sex == "male"), "\n")

# ============================================================
# 3. FILTER
# ============================================================
cat("[3] Filtering low-count peaks...\n")
y    <- DGEList(counts = as.matrix(counts), group = x)
keep <- filterByExpr(y, group = x)
y    <- y[keep, , keep.lib.sizes = FALSE]
cat("    Peaks after filtering:", nrow(y), "\n")

# ============================================================
# 4. RUVr: compute residuals
# ============================================================
cat("[4] Computing deviance residuals for RUVr...\n")
set    <- newSeqExpressionSet(as.matrix(y$counts),
           phenoData = data.frame(x, row.names = colnames(y)))
design <- model.matrix(~ x, data = pData(set))

y_tmp  <- calcNormFactors(y, method = "RLE")
y_tmp  <- estimateDisp(y_tmp, design)
fit    <- glmFit(y_tmp, design)
res    <- residuals(fit, type = "deviance")
seqUQ  <- betweenLaneNormalization(set, which = "upper")

# ============================================================
# 5. RUVr k=1~6: diagnostic plots + save corrected matrix
# ============================================================
cat("[5] Running RUVr k=1~6...\n")
dir.create(outdir, showWarnings = FALSE, recursive = TRUE)

colors <- ifelse(meta$sex == "female", "#E64B35", "#4DBBD5")

for (k in k_range) {
  cat("    k =", k, "\n")
  set2  <- RUVr(seqUQ, rownames(set), k = k, res)
  ndddd <- normCounts(set2)

  suppressMessages({
    pdf(paste0(outdir, "snACAT_CM_RUVr_k", k, "_RLE.pdf"))
    plotRLE(set2, outline = FALSE, ylim = c(-2, 2),
            col = colors, main = paste0("snATAC CM RUVr k=", k))
    dev.off()
  })

  suppressMessages({
    pdf(paste0(outdir, "snACAT_CM_RUVr_k", k, "_PCA.pdf"))
    plotPCA(set2, col = colors,
            main = paste0("snATAC CM RUVr k=", k), cex = 0.6)
    dev.off()
  })

  log_mat <- log1p(t(ndddd))
  pca     <- prcomp(log_mat, scale. = FALSE)
  pct     <- round(summary(pca)$importance[2, 1:2] * 100, 1)

  df_pca <- data.frame(
    PC1    = pca$x[, 1],
    PC2    = pca$x[, 2],
    sex    = meta[rownames(pca$x), "sex"],
    sample = meta[rownames(pca$x), "sample"]
  )
  p <- ggplot(df_pca, aes(PC1, PC2, color = sex, label = sample)) +
    geom_point(size = 3) +
    geom_text_repel(size = 2.5, max.overlaps = 20) +
    scale_color_manual(values = c(female = "#E64B35", male = "#4DBBD5")) +
    labs(title = paste0("snATAC CM RUVr k=", k),
         x = paste0("PC1 (", pct[1], "%)"),
         y = paste0("PC2 (", pct[2], "%)")) +
    theme_bw() +
    theme(plot.title = element_text(hjust = 0.5, face = "bold"))
  ggsave(paste0(outdir, "snACAT_CM_RUVr_k", k, "_PCA_gg.pdf"),
         p, width = 8, height = 6)

  write.table(ndddd,
              paste0(outdir, "snACAT_CM_RUVr_k", k, "_normCounts.txt"),
              sep = "\t", quote = FALSE, col.names = NA)
}

# ============================================================
# 6. DAR with chosen k
# ============================================================
cat("[6] Running DAR with k =", k_final, "...\n")

set2    <- RUVr(seqUQ, rownames(set), k = k_final, res)
W_cols  <- grep("^W_", colnames(pData(set2)), value = TRUE)
df2     <- data.frame(x = x, pData(set2)[, W_cols, drop = FALSE],
                      row.names = colnames(y$counts))
design2 <- model.matrix(~ ., data = df2)

y2 <- DGEList(counts = y$counts, group = x)
y2 <- calcNormFactors(y2, method = "RLE")
y2 <- estimateDisp(y2, design2)
fit2 <- glmFit(y2, design2)

# coef = 2 corresponds to sex effect (female vs male)
lrt <- glmLRT(fit2, coef = 2)
dar <- topTags(lrt, n = Inf, adjust.method = "BH")$table

dar$peak      <- rownames(dar)
dar$direction <- "non-DAR"
dar$direction[dar$FDR < 0.01 & dar$logFC >  log2(1.5)] <- "MORE"
dar$direction[dar$FDR < 0.01 & dar$logFC < -log2(1.5)] <- "LESS"

cat("    MORE DAR:", sum(dar$direction == "MORE"), "\n")
cat("    LESS DAR:", sum(dar$direction == "LESS"), "\n")

# ============================================================
# 7. SAVE DAR results
# ============================================================
cat("[7] Saving results...\n")
write.table(dar,
  paste0(outdir, "snACAT_CM_DAR_female_vs_male_RUVr_k", k_final, ".txt"),
  sep = "\t", quote = FALSE, row.names = FALSE)

write.table(dar[dar$direction != "non-DAR", ],
  paste0(outdir, "snACAT_CM_DAR_female_vs_male_RUVr_k", k_final, "_filter.txt"),
  sep = "\t", quote = FALSE, row.names = FALSE)

cat("[Done]\n")
