## modlyn.models.SimpleLogReg

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)