scdori.get_tf_expression#
- scdori.get_tf_expression(tf_expression_mode, model, device, train_loader, rna_anndata, atac_anndata, num_cells, tf_indices, encoding_batch_onehot, config_file)#
Compute TF expression per topic.
If
tf_expression_modeis “True”, this function computes the mean TF expression for the top-k cells in each topic. Otherwise, it uses a normalized topic-TF decoder matrix from the model.- Parameters:
tf_expression_mode (str) – Mode for TF expression. “True” calculates per-topic TF expression from top-k cells, “latent” uses the topic-TF decoder matrix.
model (torch.nn.Module) – The scDoRI model containing encoder and decoder modules.
device (torch.device) – The device (CPU or CUDA) used for PyTorch tensors.
train_loader (DataLoader) – DataLoader for training data.
rna_anndata (anndata.AnnData) – RNA single-cell data in AnnData format.
atac_anndata (anndata.AnnData) – ATAC single-cell data in AnnData format.
num_cells (np.ndarray) – number of cells constituting each input metacell, set to 1 for single cell data.
tf_indices (list of int) – Indices of TF features in the RNA data.
encoding_batch_onehot (np.ndarray) – One-hot encoding for batch information.
config_file (python file) – Configuration object with model training.
- Returns:
torch.Tensor A (num_topics x num_tfs) tensor of TF expression values for each topic.