modlyn.models.SimpleLogReg .md

class modlyn.models.SimpleLogReg(adata, label_column, learning_rate=0.01, weight_decay=0.01)

Bases: LightningModule

A simple LightningModule for classification tasks using a linear layer.

Parameters:
  • adata (AnnData) – An AnnData to infer dimensions from.

  • label_column (str) – Name of the column in obs that contains the target values.

  • learning_rate (float, default: 0.01) – Learning rate for the optimizer.

  • weight_decay (float, default: 0.01) – Weight decay for the optimizer.

forward(inputs)
training_step(batch, batch_idx)
on_train_epoch_end()
Return type:

None

validation_step(batch, batch_idx)
on_validation_epoch_end()
Return type:

None

configure_optimizers()
fit(adata_train, adata_val, train_dataloader_kwargs=None, val_dataloader_kwargs=None, max_epochs=4, log_every_n_steps=1, num_sanity_val_steps=0, max_steps=3000)

Fit the model using a SimpleLogRegDataModule.

Parameters:
  • adata_train (AnnData | None) – AnnData object containing the training data.

  • adata_val (AnnData | None) – AnnData object containing the validation data.

  • train_dataloader_kwargs (default: None) – Additional keyword arguments passed to the torch DataLoader for the training dataset.

  • val_dataloader_kwargs (default: None) – Additional keyword arguments passed to the torch DataLoader for the validation dataset.

  • max_epochs (int, default: 4) – Maximum number of epochs to train.

  • log_every_n_steps (int, default: 1) – Log training metrics every n steps.

  • num_sanity_val_steps (int, default: 0) – Number of sanity validation steps to run before training.

  • max_steps (int, default: 3000) – Maximum number of training steps.

get_weights()

Get the weights of the linear layer as a DataFrame.

Return type:

DataFrame

plot_losses(figsize=(15, 6))

Plot training and validation losses over training steps.

plot_classification_report(adata)