prototorch_models/prototorch/models/abstract.py

32 lines
1012 B
Python
Raw Normal View History

2021-04-29 17:14:33 +00:00
import pytorch_lightning as pl
import torch
2021-05-03 11:20:49 +00:00
from torch.optim.lr_scheduler import ExponentialLR
2021-04-29 17:14:33 +00:00
2021-05-11 14:13:00 +00:00
class AbstractPrototypeModel(pl.LightningModule):
@property
def prototypes(self):
return self.proto_layer.components.detach().cpu()
@property
def components(self):
"""Only an alias for the prototypes."""
return self.prototypes
2021-04-29 17:14:33 +00:00
def configure_optimizers(self):
2021-05-11 14:13:00 +00:00
optimizer = self.optimizer(self.parameters(), lr=self.hparams.lr)
2021-05-03 11:20:49 +00:00
scheduler = ExponentialLR(optimizer,
gamma=0.99,
last_epoch=-1,
verbose=False)
sch = {
"scheduler": scheduler,
"interval": "step",
} # called after each training step
return [optimizer], [sch]
class PrototypeImageModel(pl.LightningModule):
def on_train_batch_end(self, outputs, batch, batch_idx, dataloader_idx):
self.proto_layer.components.data.clamp_(0.0, 1.0)