modlyn.models.SimpleLogReg
¶
- class modlyn.models.SimpleLogReg(adata, label_column, learning_rate=0.01, weight_decay=0.01)¶
Bases:
LightningModuleA simple LightningModule for classification tasks using a linear layer.
- Parameters:
adata (
AnnData) – AnAnnDatato infer dimensions from.label_column (
str) – Name of the column inobsthat 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) –AnnDataobject containing the training data.adata_val (
AnnData|None) –AnnDataobject 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)¶