Skip to content

rslearn.train.optimizer

optimizer

Optimizers for rslearn.

OptimizerFactory

A factory class that initializes the optimizer given the LightningModule.

Source code in rslearn/train/optimizer.py
class OptimizerFactory:
    """A factory class that initializes the optimizer given the LightningModule."""

    def build(self, lm: L.LightningModule) -> Optimizer:
        """Build the optimizer configured by this factory class."""
        raise NotImplementedError

build

build(lm: LightningModule) -> Optimizer

Build the optimizer configured by this factory class.

Source code in rslearn/train/optimizer.py
def build(self, lm: L.LightningModule) -> Optimizer:
    """Build the optimizer configured by this factory class."""
    raise NotImplementedError

AdamW dataclass

Bases: OptimizerFactory

Factory for AdamW optimzier.

Source code in rslearn/train/optimizer.py
@dataclass
class AdamW(OptimizerFactory):
    """Factory for AdamW optimzier."""

    lr: float = 0.001
    betas: tuple[float, float] = (0.9, 0.999)
    eps: float | None = None
    weight_decay: float | None = None

    def build(self, lm: L.LightningModule) -> Optimizer:
        """Build the AdamW optimizer."""
        params = [p for p in lm.parameters() if p.requires_grad]
        kwargs = {k: v for k, v in asdict(self).items() if v is not None}
        return torch.optim.AdamW(params, **kwargs)

build

build(lm: LightningModule) -> Optimizer

Build the AdamW optimizer.

Source code in rslearn/train/optimizer.py
def build(self, lm: L.LightningModule) -> Optimizer:
    """Build the AdamW optimizer."""
    params = [p for p in lm.parameters() if p.requires_grad]
    kwargs = {k: v for k, v in asdict(self).items() if v is not None}
    return torch.optim.AdamW(params, **kwargs)