| import torchmetrics |
| import torch |
|
|
| from torch.nn.functional import cross_entropy |
| from ..model_interface import register_model |
| from .base import SaprotBaseModel |
|
|
|
|
| @register_model |
| class SaprotClassificationModel(SaprotBaseModel): |
| def __init__(self, num_labels: int, **kwargs): |
| """ |
| Args: |
| num_labels: number of labels |
| **kwargs: other arguments for SaprotBaseModel |
| """ |
| self.num_labels = num_labels |
| super().__init__(task="classification", **kwargs) |
| |
| def initialize_metrics(self, stage): |
| return {f"{stage}_acc": torchmetrics.Accuracy()} |
|
|
| def forward(self, inputs, coords=None): |
| if coords is not None: |
| inputs = self.add_bias_feature(inputs, coords) |
|
|
| |
| if self.freeze_backbone: |
| repr = torch.stack(self.get_hidden_states(inputs, reduction="mean")) |
| x = self.model.classifier.dropout(repr) |
| x = self.model.classifier.dense(x) |
| x = torch.tanh(x) |
| x = self.model.classifier.dropout(x) |
| logits = self.model.classifier.out_proj(x) |
|
|
| else: |
| logits = self.model(**inputs).logits |
|
|
| return logits |
|
|
| def loss_func(self, stage, logits, labels): |
| label = labels['labels'] |
| loss = cross_entropy(logits, label) |
|
|
| |
| for metric in self.metrics[stage].values(): |
| metric.update(logits.detach(), label) |
|
|
| if stage == "train": |
| log_dict = self.get_log_dict("train") |
| log_dict["train_loss"] = loss |
| self.log_info(log_dict) |
|
|
| |
| self.reset_metrics("train") |
|
|
| return loss |
|
|
| def test_epoch_end(self, outputs): |
| log_dict = self.get_log_dict("test") |
| log_dict["test_loss"] = torch.cat(self.all_gather(outputs), dim=-1).mean() |
|
|
| print(log_dict) |
| self.log_info(log_dict) |
|
|
| self.reset_metrics("test") |
|
|
| def validation_epoch_end(self, outputs): |
| log_dict = self.get_log_dict("valid") |
| log_dict["valid_loss"] = torch.cat(self.all_gather(outputs), dim=-1).mean() |
|
|
| self.log_info(log_dict) |
| self.reset_metrics("valid") |
| self.check_save_condition(log_dict["valid_acc"], mode="max") |