Download model/saprot/saprot_classification_model.py from OneScience-Group/SaProt: direct link, hf CLI and curl.
- Browser
- Download file 2.35 kB
-
https://huggingface.co/OneScience-Group/SaProt/resolve/main/model/saprot/saprot_classification_model.py
- Command line
-
hf download hf://OneScience-Group/SaProt/model/saprot/saprot_classification_model.py
-
curl -L -o saprot_classification_model.py https://huggingface.co/OneScience-Group/SaProt/resolve/main/model/saprot/saprot_classification_model.py
2.35 kB
| import torchmetrics | |
| import torch | |
| from torch.nn.functional import cross_entropy | |
| from ..model_interface import register_model | |
| from .base import SaprotBaseModel | |
| 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 backbone is frozen, the embedding will be the average of all residues | |
| 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) | |
| # Update metrics | |
| 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) | |
| # Reset train metrics | |
| 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") |