Antibody_deep_learning / scripts /03_gan_inference.R
anzhi2710gmailcom's picture
Upload folder using huggingface_hub
fe8e241 verified
Raw
History Blame Contribute Delete
3.2 kB
args <- commandArgs(trailingOnly = FALSE)
script_arg <- grep("^--file=", args, value = TRUE)
if (length(script_arg) > 0) {
script_path <- normalizePath(sub("^--file=", "", script_arg[1]))
setwd(normalizePath(file.path(dirname(script_path), "..")))
}
library(reticulate)
use_python(Sys.getenv("RETICULATE_PYTHON"), required = TRUE)
tf <- import("tensorflow", convert = FALSE)
helper <- import_from_path("tf_savedmodel_helper", path = file.path(getwd(), "scripts"), convert = TRUE)
aa_order <- c("A", "R", "N", "D", "C", "Q", "E", "G", "H", "I",
"L", "K", "M", "F", "P", "S", "T", "W", "Y", "V", "X", "-")
decode_seq_cdr3 <- function(seq_list) {
aa_list <- list()
for (i in seq_along(seq_list)) {
peptide1 <- as.data.frame(seq_list[[i]])
colnames(peptide1) <- aa_order
peptide1$X <- NULL
aa <- colnames(peptide1)[apply(peptide1, 1, which.max)]
gaps <- which(aa %in% "-")
if (length(gaps) > 0) {
gap_first <- gaps[gaps > 1][1]
gap_last_candidates <- gaps[gaps < 32]
gap_last <- gap_last_candidates[length(gap_last_candidates)]
if (!is.na(gap_first) && !is.na(gap_last) && gap_first <= gap_last) {
aa[gap_first:gap_last] <- "-"
}
}
aa_string <- paste(aa, collapse = "")
aa_string <- gsub("-", "", aa_string)
aa_list[[i]] <- aa_string
}
data.frame(aa = unlist(aa_list), stringsAsFactors = FALSE)
}
get_input_name <- function(model_dir) {
model <- tf$saved_model$load(model_dir)
serving <- model$signatures$get("serving_default")
sig_text <- py_str(serving$structured_input_signature)
out_text <- py_str(serving$structured_outputs)
cat("\nModel:", model_dir, "\n")
cat("Input signature:", sig_text, "\n")
cat("Output signature:", out_text, "\n")
input_name <- sub(".*'([^']+)': TensorSpec.*", "\\1", sig_text)
cat("Input name:", input_name, "\n")
input_name
}
generate_model <- function(model_id, n_seq = 100, batch_size = 20, latent_dim = 100) {
model_dir <- paste0("weight/GAN/GAN_model_", model_id)
input_name <- get_input_name(model_dir)
seq_list <- list()
k <- 1
set.seed(1000 + model_id)
while (length(seq_list) < n_seq) {
noise <- matrix(rnorm(batch_size * latent_dim), nrow = batch_size, ncol = latent_dim)
fake <- helper$predict_saved_model(model_dir, input_name, noise)
for (i in seq_len(dim(fake)[1])) {
f <- fake[i,,,]
dim(f) <- c(32, 22, 1)
seq_list[[k]] <- f
k <- k + 1
if (length(seq_list) >= n_seq) break
}
}
decode_seq_cdr3(seq_list)
}
seq_all <- readRDS("model/GAN/seq_all.RDS")
gen_seq <- list()
for (i in 1:15) {
cat("\nGenerating model", i, names(seq_all)[i], "\n")
temp <- generate_model(i, n_seq = 100)
gen_seq[[i]] <- temp$aa
cat("unique first 10:\n")
print(head(unique(gen_seq[[i]]), 10))
}
names(gen_seq) <- names(seq_all)
saveRDS(gen_seq, "model/GAN/gen_seq_tf218.RDS")
out <- data.frame()
for (i in seq_along(gen_seq)) {
out <- rbind(out, data.frame(model_id = i, group = names(gen_seq)[i], aa = gen_seq[[i]], stringsAsFactors = FALSE))
}
write.table(out, "model/GAN/gen_seq_tf218.tsv", sep = "\t", quote = FALSE, row.names = FALSE)
cat("\nGAN generation OK\n")