lynx.model.HeteroAttnVGAE¶
- class lynx.model.HeteroAttnVGAE(*args, **kwargs)[source]¶
Bases:
BaseModelLearning latent manifold w/ Conditional VGAE on hetero-graph Generative path: DESI (u) -> Latent (z) -> Xenium (x)
- Parameters:
configs (ml_collections.ConfigDict)
device (torch.device)
- __init__(configs, device=torch.device)[source]¶
- Parameters:
configs (ml_collections.ConfigDict)
device (torch.device)
Methods
__init__(configs[, device])checkpoint(curr_loss, min_loss, patience, ...)evaluate(adata_ref, adata_query, graph_data)Full model inference
fit(dataset, train_configs[, DEBUG, log_wandb])Full model training
get_anneal_weight(beta, epoch, warmup_epochs)guide(data)Variational guide
load_state(save_path)lognorm(x)model(data)Generative model
model_train(model, dataset, train_configs[, ...])monitor_metrics(data, device[, key])(Debug-only) Monitor latent factor correlations & reconstruction
plot_latent_corr(pz_corr_scores, qz_corr_scores)plot_loss(train_losses, val_losses)predict(data, device)Get latent (z) & reconstructions from batched data object
set_desc(pbar, epoch, train_loss, val_loss)setup(model, train_configs)Setup optimizer & inference objects
train_step(model, dataloader, svi, device[, key])Single-epoch training step
val_step(model, dataloader, svi, device[, key])Single-epoch validation step
Attributes
init_model_weightsset_seed- evaluate(adata_ref, adata_query, graph_data, n_subgraphs=1, device=torch.device)[source]¶
Full model inference
- Parameters:
adata_ref (scanpy.AnnData)
adata_query (scanpy.AnnData)
graph_data (HeteroDataset)
n_subgraphs (int)
device (torch.device)