lynx.model.HeteroAttnVGAE

class lynx.model.HeteroAttnVGAE(*args, **kwargs)[source]

Bases: BaseModel

Learning 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_weights

set_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)

fit(dataset, train_configs, DEBUG=False, log_wandb=False)[source]

Full model training

guide(data)[source]

Variational guide

model(data)[source]

Generative model

predict(data, device)[source]

Get latent (z) & reconstructions from batched data object