{"metadata":{"kernelspec":{"name":"ir","display_name":"R","language":"R"},"language_info":{"name":"R","codemirror_mode":"r","pygments_lexer":"r","mimetype":"text/x-r-source","file_extension":".r","version":"4.0.5"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59094,"databundleVersionId":7010844,"sourceType":"competition"},{"sourceId":7712331,"sourceType":"datasetVersion","datasetId":4441094}],"dockerImageVersionId":30618,"isInternetEnabled":true,"language":"r","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## <p style=\"border: 3px solid #3B2F2F; border-radius: 10px; padding: 15px; background-color: #ffc7ba; text-align: center; font-family: 'Arial', Times, serif; font-weight: bold; letter-spacing: 1px; color: #3B2F2F; font-size: 24px; margin-bottom: 10px;\"> NN on melted data and embedding-based feature representations using torch</p>\n\nThis notebook is based on the 30th place solution by FRENIO REDEKER ([write up](https://www.kaggle.com/competitions/open-problems-single-cell-perturbations/discussion/461649)), espesially on 3 of his notebooks:\n- [Export TabMod NN Embeddings](https://www.kaggle.com/code/frenio/30-op2scp-export-tabmod-nn-embeddings#Train-without-PCA) - melting the data to get 3 features - cell type, compound and gene and use NN with embedding layer to generate embeddings for each\n- [Look at TabMod NN Embeddings](https://www.kaggle.com/code/frenio/30-op2scp-look-at-tabmod-nn-embeddings) - plotting 2 first principal components for each set of embeddings\n- [Tabular Model NN with PCA10 Denoising](https://www.kaggle.com/code/frenio/30-op2scp-tabular-model-nn-with-pca10-denoising) - submitting predictions generated with NN\n\nHere I use only 1000 random genes to run the kaggle notebook in reasonable time with CPU).\n\nIn the original version by FRENIO REDEKER:   \n**Public LB score 0.565  \nPrivat LB score 0.784**","metadata":{}},{"cell_type":"markdown","source":"****","metadata":{}},{"cell_type":"code","source":"library(arrow)\nlibrary(data.table)\nlibrary(qs)\nlibrary(tictoc)\nlibrary(ggplot2)\nlibrary(patchwork)\nlibrary(ggrepel)\nlibrary(pheatmap)\nlibrary(torch)\nlibrary(luz)","metadata":{"_uuid":"051d70d956493feee0c6d64651c6a088724dca2a","_execution_state":"idle","execution":{"iopub.status.busy":"2024-04-04T03:22:32.960011Z","iopub.execute_input":"2024-04-04T03:22:32.962089Z","iopub.status.idle":"2024-04-04T03:22:33.004533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Helper functions**","metadata":{}},{"cell_type":"code","source":"shape <- function(dt) {\n    cat(\"\\nShape:\", nrow(dt), \"rows x\", ncol(dt), \"columns\")\n}\nfig <- function(width, heigth){\n  options(repr.plot.width = width, repr.plot.height = heigth)\n}","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:22:33.153096Z","iopub.execute_input":"2024-04-04T03:22:33.155115Z","iopub.status.idle":"2024-04-04T03:22:33.172353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Read Data","metadata":{}},{"cell_type":"markdown","source":"**Train data**","metadata":{}},{"cell_type":"code","source":"de_train <- read_parquet('/kaggle/input/open-problems-single-cell-perturbations/de_train.parquet')\nde_train <- setDT(de_train)\n\nsample_genes <- sample(ncol(de_train)-5, 1000)\nde_train <- de_train[, c(1:5, sample_genes), with = FALSE]\n\nhead(de_train)\nshape(de_train)","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:22:33.367390Z","iopub.execute_input":"2024-04-04T03:22:33.369640Z","iopub.status.idle":"2024-04-04T03:22:35.153851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Test data**","metadata":{}},{"cell_type":"code","source":"id_map <- fread(\"/kaggle/input/open-problems-single-cell-perturbations/id_map.csv\")\ngene_cols <- names(de_train)[unlist(lapply(de_train, is.numeric))]\nsuppressWarnings(id_map[, (gene_cols) := 0])\n\nhead(id_map)\nshape(id_map)","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:23:01.292765Z","iopub.execute_input":"2024-04-04T03:23:01.296546Z","iopub.status.idle":"2024-04-04T03:23:02.604525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Functions","metadata":{}},{"cell_type":"markdown","source":"**prepare data** - performs the following operations with a data table:\n- splitting to train and valid subsets (80/20) by adding a column 'idx' with either \"train\" or \"valid\" values \n- Reshaping the result (returns the data in long format)\n- if denoise = TRUE, before reshaping also performs:  \n1) Principal Component Analysis (PCA) on the target columns (genes) after scaling and centering of the data,   \n2) dimensionality reduction (retains only the first n_comp principal components from the PCA results),  \n3) inverse transformation (inversely transforms the reduced principal components back to the original space using the PCA rotation matrix, scale, and center),  \n4) bind the result for the train subset with the untouched valid subset. \n","metadata":{}},{"cell_type":"code","source":"prepare_data <- function(dt, denoise = FALSE, n_comp = 614, seed = 1) {\n  \n  gene_cols <- names(dt)[unlist(lapply(dt, is.numeric))]\n  id_cols <- c(\"cell_type\", \"sm_name\")\n  \n  melt_dt <- melt(\n    dt, value.name = \"value\",\n    variable.name = \"gene\",\n    variable.factor = FALSE,\n    measure.vars = gene_cols,\n    id.vars = id_cols\n  )\n  \n  set.seed(seed)\n  train_idx <- sample(1:nrow(melt_dt), size = floor(0.8 * nrow(melt_dt)))\n  valid_idx <- setdiff(1:nrow(melt_dt), train_idx)\n  melt_dt[train_idx, idx := \"train\"][valid_idx, idx := \"valid\"]\n  \n  if (denoise) {        \n    \n    pca_data <- prcomp(dt[, ..gene_cols], scale. = TRUE, center = TRUE)\n    reduced_data <- pca_data$x[, 1:n_comp]\n    \n    reduced_data_inv_trans <- t(\n      t(reduced_data %*% t(pca_data$rotation[, 1:n_comp])) *\n        pca_data$scale + pca_data$center\n    )\n    \n    denoised_melt_dt <- melt(\n      cbind(dt[, ..id_cols], reduced_data_inv_trans),\n      value.name = \"value\",\n      variable.name = \"gene\",\n      variable.factor = FALSE,\n      measure.vars = gene_cols,\n      id.vars = id_cols\n    )\n    \n    denoised_melt_dt[train_idx, idx := \"train\"][valid_idx, idx := \"valid\"]\n    denoised_melt_dt[valid_idx, value := melt_dt[valid_idx, value]]\n    return(denoised_melt_dt)\n  }\n  return(melt_dt)\n}","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:23:08.551359Z","iopub.execute_input":"2024-04-04T03:23:08.553306Z","iopub.status.idle":"2024-04-04T03:23:08.568625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**get_train_history** -  extracts training and validation metrics (loss and RMSE) for each epoch from a model fitted during training and returns a a data.table object in long format with columns for loss and RMSE.","metadata":{}},{"cell_type":"code","source":"get_train_history <- function(fitted) {\n    \n    train <- rbindlist(lapply(1:length(fitted$records$metrics$train), function(lst) {\n        data.table(\n            epoch = lst,\n            set = \"train\",\n            loss = fitted$records$metrics$train[[lst]]$loss,\n            rmse = fitted$records$metrics$train[[lst]]$rmse\n        )\n    }))\n    valid <- rbindlist(lapply(1:length(fitted$records$metrics$valid), function(lst) {\n            data.table(\n                epoch = lst,\n                set = \"valid\",\n                loss = fitted$records$metrics$valid[[lst]]$loss,\n                rmse = fitted$records$metrics$valid[[lst]]$rmse\n            )\n        }))\n    train_history <- rbind(train, valid)\n  \n  return(train_history)   \n}","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:23:12.058230Z","iopub.execute_input":"2024-04-04T03:23:12.060169Z","iopub.status.idle":"2024-04-04T03:23:12.075415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare data","metadata":{}},{"cell_type":"markdown","source":"**Train**","metadata":{}},{"cell_type":"code","source":"train_melted <- prepare_data(de_train, denoise = TRUE, n_comp = 10)\nhead(train_melted)\nshape(train_melted)","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:23:14.633595Z","iopub.execute_input":"2024-04-04T03:23:14.635473Z","iopub.status.idle":"2024-04-04T03:23:15.341318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Test**","metadata":{}},{"cell_type":"code","source":"test_melted <- melt(id_map, value.name = \"value\",\n                    variable.name = \"gene\",\n                    variable.factor = FALSE,\n                    id.vars = c(\"id\", \"cell_type\", \"sm_name\"),\n                    measure.vars = gene_cols)\nhead(test_melted)\nshape(test_melted)","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:23:18.054218Z","iopub.execute_input":"2024-04-04T03:23:18.056094Z","iopub.status.idle":"2024-04-04T03:23:18.126535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Make dictionary for feature categories**","metadata":{}},{"cell_type":"code","source":"cat_features_cols <- c(\"cell_type\", \"sm_name\", \"gene\")\ndict <- lapply(cat_features_cols, function(f) {\n    \n    f_dict <- unique(train_melted[, ..f])\n    f_dict[, (paste0(f, \"_index\")) := as.numeric(as.factor(get(f)))]\n    f_dict <- f_dict[order(get(paste0(f, \"_index\")))]\n})\nnames(dict) <- cat_features_cols\n\nlapply(dict, function(dt) {\n    \n    add_row <- dt[1]\n    add_row[] <- \"...\"\n    rbindlist(list(head(dt, 2), add_row, tail(dt, 2)))\n})","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:23:34.034603Z","iopub.execute_input":"2024-04-04T03:23:34.036619Z","iopub.status.idle":"2024-04-04T03:23:34.158056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Transform categirical feature columns in train and test data to numeric indices according to the dictionary**","metadata":{}},{"cell_type":"code","source":"for (col in cat_features_cols) {\n    \n    levels_order <- dict[[col]][, get(col)]\n    \n    train_melted[, (col) := as.numeric(\n        factor(get(col), levels = levels_order)\n    )]\n    test_melted[, (col) := as.numeric(\n        factor(get(col), levels = levels_order)\n    )]\n}","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:23:47.449114Z","iopub.execute_input":"2024-04-04T03:23:47.450933Z","iopub.status.idle":"2024-04-04T03:23:47.542096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train NN","metadata":{}},{"cell_type":"markdown","source":"**Make a dataset with categorical columns transformed into a numeric matrix suitable for training models**","metadata":{}},{"cell_type":"code","source":"make_dataset <- dataset(\n        initialize = function(dt) {\n            self$x_cat <- self$get_categorical(dt)\n            self$y <- dt[[\"value\"]]\n        },\n        .getitem = function(i) {\n            x_cat <- self$x_cat[i, ]\n            y <- self$y[i]\n            list(x = x_cat, y = y)\n        },\n        .length = function() {\n            length(self$y)\n        },\n        get_categorical = function(dt, cols = cat_features_cols) {\n            as.matrix(\n                dt[, lapply(.SD, as.integer),\n                   .SDcols = cat_features_cols]\n            )\n        }\n    )","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:23:57.077094Z","iopub.execute_input":"2024-04-04T03:23:57.079152Z","iopub.status.idle":"2024-04-04T03:23:57.100684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds <- make_dataset(train_melted)\nds[1]","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:24:14.304895Z","iopub.execute_input":"2024-04-04T03:24:14.306823Z","iopub.status.idle":"2024-04-04T03:24:14.359734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Split data into training and validation subsets**","metadata":{}},{"cell_type":"code","source":"train_indices <- which(train_melted[, idx == \"train\"])\nvalid_indices <- which(train_melted[, idx == \"valid\"])\n\ntrain_ds <- dataset_subset(ds, train_indices)\ntrain_dl <-  dataloader(train_ds, batch_size = 4096, shuffle = TRUE)\n\nvalid_ds <- dataset_subset(ds, valid_indices)\nvalid_dl <- dataloader(valid_ds, batch_size = 4096, shuffle = FALSE)","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:24:33.690937Z","iopub.execute_input":"2024-04-04T03:24:33.692827Z","iopub.status.idle":"2024-04-04T03:24:33.816970Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Set embeddings' dimension**","metadata":{}},{"cell_type":"code","source":"num_embeddings <- t(train_melted[, lapply(.SD, uniqueN),\n                                 .SDcols = cat_features_cols]\n)\n\nembedding_sizes <- data.table(\n  col = rownames(num_embeddings),\n  num = num_embeddings[, 1],\n  size = c(5, 26, 600)\n)\nembedding_sizes","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:24:49.619289Z","iopub.execute_input":"2024-04-04T03:24:49.621742Z","iopub.status.idle":"2024-04-04T03:24:49.829383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Embedding module**","metadata":{}},{"cell_type":"code","source":"embedding_module <- nn_module(\n  initialize = function(embedding_sizes) {\n    self$embeddings <- nn_module_list(\n        lapply(1:nrow(embedding_sizes), function(i) {\n            nn_embedding(num_embeddings = embedding_sizes[i, num],\n                         embedding_dim = embedding_sizes[i, size])\n        })\n    )\n  },\n  forward = function(x) {\n    embedded <- vector(mode = \"list\",\n                       length = length(self$embeddings)\n    )\n    for (i in 1:length(self$embeddings)) {\n        \n        embedded[[i]] <- self$embeddings[[i]](x[, i])\n    }\n    torch_cat(embedded, dim = 2)\n  }\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:24:57.456148Z","iopub.execute_input":"2024-04-04T03:24:57.457983Z","iopub.status.idle":"2024-04-04T03:24:57.478278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**NN model**","metadata":{}},{"cell_type":"code","source":"net <- nn_module(\n  initialize = function(embedding_sizes,\n                        fc1_dim,\n                        fc2_dim,\n                        fc3_dim) {\n    self$embedder <- embedding_module(embedding_sizes)\n    self$fc1 <- nn_linear(sum(embedding_sizes$size), fc1_dim)\n    self$fc2 <- nn_linear(fc1_dim, fc2_dim)\n    self$fc3 <- nn_linear(fc2_dim, fc3_dim)\n    self$output <- nn_linear(fc3_dim, 1)\n  },\n  forward = function(x) {\n    embedded <- self$embedder(x)\n    score <- embedded %>%\n      self$fc1() %>%\n      nnf_relu() %>%\n      self$fc2() %>%\n      nnf_relu() %>%\n      self$fc3() %>%\n      nnf_relu() %>%\n      self$output()\n    score[, 1]\n  }\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:27:00.699371Z","iopub.execute_input":"2024-04-04T03:27:00.701723Z","iopub.status.idle":"2024-04-04T03:27:00.721129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Fit model**","metadata":{}},{"cell_type":"code","source":"fc1_dim <- 1000\nfc2_dim <- 500\nfc3_dim <- 250\nnum_epochs <- 20\n\ntic()\nfitted <- net %>%\n  setup(\n      optimizer = optim_adam,\n      loss = nn_mse_loss(),\n      metrics = list(luz_metric_mae(), luz_metric_rmse())\n  ) %>%\n  set_hparams(\n      embedding_sizes = embedding_sizes,\n      fc1_dim = fc1_dim,\n      fc2_dim = fc2_dim,\n      fc3_dim = fc3_dim\n  ) %>%\n  fit(train_dl,\n      epochs = 20,\n      valid_data = valid_dl,\n      accelerator = accelerator(cpu = if (cuda_is_available()) FALSE else TRUE),\n      callbacks = list(\n          luz_callback_lr_scheduler(\n              lr_one_cycle,\n              max_lr = 3e-4,\n              epochs = 20,\n              steps_per_epoch = length(train_dl),\n              call_on = \"on_batch_end\"\n          )\n      ),\n      verbose = TRUE\n     )\ntoc()","metadata":{"execution":{"iopub.status.busy":"2024-04-04T03:27:09.144810Z","iopub.execute_input":"2024-04-04T03:27:09.147045Z","iopub.status.idle":"2024-04-04T04:01:53.741626Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Plot train history","metadata":{}},{"cell_type":"code","source":"train_hist <- get_train_history(fitted) \ndcast(train_hist, epoch ~ set, value.var = \"loss\")","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:02:51.914875Z","iopub.execute_input":"2024-04-04T04:02:51.918865Z","iopub.status.idle":"2024-04-04T04:02:51.994745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig(25, 10)\n\np1 <- ggplot(train_hist, aes(x = epoch, y = loss, group = set)) +\n        geom_line(aes(color = set)) +\n        geom_point(aes(color = set)) +\n        theme_bw(base_size = 22) +\n        theme(legend.title=element_blank())\n\np2 <- ggplot(train_hist, aes(x = epoch, y = rmse, group = set)) +\n        geom_line(aes(color = set)) +\n        geom_point(aes(color = set)) +\n        theme_bw(base_size = 22) +\n        theme(legend.title=element_blank())\n\np1 + p2 + plot_layout(guides = \"collect\") +\nplot_annotation(title = \"Train history\") &\ntheme(text = element_text(size = 22)) ","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:03:06.051584Z","iopub.execute_input":"2024-04-04T04:03:06.053566Z","iopub.status.idle":"2024-04-04T04:03:06.877979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Evaluate the model (on entire train data)**","metadata":{}},{"cell_type":"code","source":"test_dl <- dataloader(ds, batch_size = 4096)\npreds <- predict(fitted, test_dl)\n\npreds <- as.matrix(preds$to(device = \"cpu\"))      \ntrue_targets <- as.matrix(train_melted[, .(value)])\n\nmae <- mean((true_targets - preds)^2)\ncat(\"RMSE:\", sqrt(mae), \"\\n\", \"MAE:\", mae)","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:03:14.770006Z","iopub.execute_input":"2024-04-04T04:03:14.771790Z","iopub.status.idle":"2024-04-04T04:04:12.775938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Get embeddings","metadata":{}},{"cell_type":"code","source":"embedding_weights <- vector(mode = \"list\")\n\nfor (i in 1:length(fitted$model$embedder$embeddings)) {\n  embedding_weights[[i]] <-\n    fitted$model$embedder$embeddings[[i]]$\n    parameters$weight$to(device = \"cpu\")\n}\nnames(embedding_weights) <- cat_features_cols","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:04:23.186386Z","iopub.execute_input":"2024-04-04T04:04:23.188448Z","iopub.status.idle":"2024-04-04T04:04:23.219513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"embeds_list <- lapply(1:length(embedding_weights), function(i) {\n    \n    emb_dt <- data.table(as.matrix(embedding_weights[[i]]))\n    setnames(emb_dt, names(emb_dt),\n             paste0(cat_features_cols[i], \"_emb\", 1:ncol(emb_dt))\n            )\n    index_name <- paste0(cat_features_cols[i], \"_index\")\n    emb_dt[, (index_name) := 1:nrow(emb_dt)]\n    \n    merge(dict[[cat_features_cols[i]]],\n          emb_dt, by = index_name, all = TRUE)\n})\nnames(embeds_list) <- cat_features_cols","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:12:13.987291Z","iopub.execute_input":"2024-04-04T04:12:13.989269Z","iopub.status.idle":"2024-04-04T04:12:14.083134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lapply(embeds_list, head, 3)","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:12:16.470789Z","iopub.execute_input":"2024-04-04T04:12:16.472624Z","iopub.status.idle":"2024-04-04T04:12:16.906282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"qsave(embeds_list, \"embeds_list.qs\")","metadata":{"execution":{"iopub.status.busy":"2024-04-01T05:28:49.320825Z","iopub.execute_input":"2024-04-01T05:28:49.324055Z","iopub.status.idle":"2024-04-01T05:28:49.372857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Explore embedding representations","metadata":{}},{"cell_type":"code","source":"plot_emb_pca <- function(emb_dt, labels = TRUE) {\n    \n    pca <- prcomp(emb_dt[, 3:ncol(emb_dt)], center = TRUE, scale = TRUE)    \n    plt_dt <- cbind(emb_dt, pca$x[, 1:2])\n    \n    plt <- ggplot(plt_dt, aes(x = PC1, y = PC2)) +\n    geom_point(size = 3) +\n    theme_bw(base_size = 22) +\n    theme(panel.grid = element_blank()) + #,aspect.ratio = 1\n    ggtitle(names(emb_dt)[2])\n    \n    if (labels) {\n        plt <- plt +\n            geom_label_repel(aes(label = unlist(emb_dt[, 2])), size = 6)\n    }\n    return(plt)\n}","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:04:35.031736Z","iopub.execute_input":"2024-04-04T04:04:35.033750Z","iopub.status.idle":"2024-04-04T04:04:35.049818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig(10, 10)\nplot_emb_pca(embeds_list[[\"cell_type\"]])","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:04:37.841063Z","iopub.execute_input":"2024-04-04T04:04:37.843857Z","iopub.status.idle":"2024-04-04T04:04:38.418620Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig(25, 18)\nplot_emb_pca(embeds_list[[\"sm_name\"]])","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:04:43.405739Z","iopub.execute_input":"2024-04-04T04:04:43.407721Z","iopub.status.idle":"2024-04-04T04:04:48.221474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig(25, 18)\nplot_emb_pca(embeds_list[[\"gene\"]], labels = FALSE)","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:04:50.888983Z","iopub.execute_input":"2024-04-04T04:04:50.890831Z","iopub.status.idle":"2024-04-04T04:04:52.377033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**PCA of embedding weights, biplot visualizing factor loadings for cell type embeddings**","metadata":{}},{"cell_type":"code","source":"fig(25, 10)\nfeature_col <- \"cell_type\"\nemb_dt <- embeds_list[[feature_col]]\npca <- prcomp(emb_dt[, 3:ncol(emb_dt)], center = TRUE, scale = TRUE)\nbiplot(pca)","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:05:13.472195Z","iopub.execute_input":"2024-04-04T04:05:13.474262Z","iopub.status.idle":"2024-04-04T04:05:13.752890Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Cosine similarity**","metadata":{}},{"cell_type":"code","source":"calculate_cosine_similarity <- function(embeds_list, feature_col) {\n  embeds <- as.matrix(embeds_list[[feature_col]][, -c(1:2)])\n    \n  norms <- sqrt(rowSums(embeds^2))  \n  dot_products <- embeds %*% t(embeds)  \n  similarity_matrix <- dot_products / (norms %*% t(norms))  \n  diag(similarity_matrix) <- 0\n  \n  rownames(similarity_matrix) <- \n    colnames(similarity_matrix) <- \n    unlist(embeds_list[[feature_col]][, 2])\n  \n  return(similarity_matrix)\n}","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:05:19.559029Z","iopub.execute_input":"2024-04-04T04:05:19.561000Z","iopub.status.idle":"2024-04-04T04:05:19.578316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig(12, 6)\nfeature_col <- \"cell_type\"\nsim_res <- calculate_cosine_similarity(embeds_list, feature_col)\n\nbreaksList = seq(min(sim_res), max(sim_res), by = 0.1)\nmyColors <- c(colorRampPalette(c(\"darkblue\", \"white\",\"darkred\"))(length(breaksList)))\n\nheat <- pheatmap(sim_res, fontsize = 16, fontsize_row = 16, fontsize_col = 16,\n                 color = myColors,\n                 breaks = breaksList,\n                 angle_col = 90,\n                 main = paste(\"Cosine similarity between\", feature_col, \"embeddings\"))","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:05:28.903512Z","iopub.execute_input":"2024-04-04T04:05:28.905493Z","iopub.status.idle":"2024-04-04T04:05:29.145662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig(30, 30)\nfeature_col <- \"sm_name\"\nsim_res <- calculate_cosine_similarity(embeds_list, feature_col)\n\nbreaksList = seq(min(sim_res), max(sim_res), by = 0.1)\nmyColors <- c(colorRampPalette(c(\"darkblue\", \"white\",\"darkred\"))(length(breaksList)))\n\nheat <- pheatmap(sim_res, fontsize = 12, fontsize_row = 12, fontsize_col = 12,\n                 color = myColors,\n                 breaks = breaksList,\n                 angle_col = 90,\n                 main = paste(\"Cosine similarity between\", feature_col, \"embeddings\"))","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:05:47.327363Z","iopub.execute_input":"2024-04-04T04:05:47.329963Z","iopub.status.idle":"2024-04-04T04:05:48.499362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict","metadata":{}},{"cell_type":"code","source":"test_ds <- make_dataset(test_melted[, !\"id\"])\ntest_dl <- dataloader(test_ds, batch_size = 4096)\npreds <- predict(fitted, test_dl)\npreds <- as.matrix(preds$to(device = \"cpu\"))      \n\nsubmit <- test_melted[, .(id, gene_index = gene, pred = preds[, 1])]\nsubmit <- submit[dict[[\"gene\"]], on = \"gene_index\"]\nsubmit <- dcast(submit, id ~ gene, value.var = \"pred\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit[1:10, 1:10]","metadata":{"execution":{"iopub.status.busy":"2024-04-04T04:08:53.569717Z","iopub.execute_input":"2024-04-04T04:08:53.571814Z","iopub.status.idle":"2024-04-04T04:08:53.616832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fwrite(submit, \"submission.csv\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}