CoolFace
Apppublic

ellenhp/query2osm

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
hydra.py113 linesDownload Raw Back to root
1from transformers import BertConfig, BertModel2import torch.nn as nn3import torch4from typing import Optional, Union, Tuple, List5from transformers.modeling_outputs import SequenceClassifierOutput6from torch.nn import CrossEntropyLoss7 8 9class HydraConfig(BertConfig):10    model_type = "hydra"11    label_groups = None12 13    def __init__(self, **kwargs):14        super().__init__(**kwargs)15 16    def num_labels(self):17        return sum([len(group) for group in self.label_groups])18 19    def distilbert_config(self):20        return BertConfig(**self.__dict__)21 22 23class HydraSequenceClassifierOutput(SequenceClassifierOutput):24    classifications: List[dict]25 26    def __init__(self, classifications=None, **kwargs):27        super().__init__(**kwargs)28        self.classifications = classifications29 30 31class Hydra(BertModel):32    config_class = HydraConfig33 34    def __init__(self, config: HydraConfig):35        super().__init__(config)36        self.config = config37        self.pre_classifier = nn.Linear(config.hidden_size, config.hidden_size)38        self.classifier = nn.Linear(config.hidden_size, sum(39            [len(group) for group in config.label_groups]))40        self.dropout = nn.Dropout(config.hidden_dropout_prob)41 42        self.embeddings.requires_grad_(False)43 44        self.post_init()45 46    def forward(47        self,48        input_ids: Optional[torch.Tensor] = None,49        attention_mask: Optional[torch.Tensor] = None,50        head_mask: Optional[torch.Tensor] = None,51        inputs_embeds: Optional[torch.Tensor] = None,52        labels: Optional[torch.LongTensor] = None,53        output_attentions: Optional[bool] = None,54        output_hidden_states: Optional[bool] = None,55        return_dict: Optional[bool] = None,56    ) -> Union[SequenceClassifierOutput, Tuple[torch.Tensor, ...]]:57        return_dict = return_dict if return_dict is not None else self.config.use_return_dict58 59        distilbert_output = super().forward(60            input_ids=input_ids,61            attention_mask=attention_mask,62            head_mask=head_mask,63            inputs_embeds=inputs_embeds,64            output_attentions=output_attentions,65            output_hidden_states=output_hidden_states,66            return_dict=return_dict67        )68        hidden_state = distilbert_output[0]  # (bs, seq_len, dim)69        pooled_output = hidden_state[:, 0]  # (bs, dim)70        pooled_output = self.pre_classifier(pooled_output)  # (bs, dim)71        pooled_output = nn.ReLU()(pooled_output)  # (bs, dim)72        pooled_output = self.dropout(pooled_output)  # (bs, dim)73        logits = self.classifier(pooled_output)  # (bs, num_labels)74 75        loss = None76        if labels is not None:77 78            loss_fct = CrossEntropyLoss()79            loss = loss_fct(logits, labels)80 81            if not return_dict:82                output = (logits,) + distilbert_output[1:]83                return ((loss,) + output) if loss is not None else output84 85        classifications = []86        if logits.shape[0] == 1:87            offset = 088            for group in self.config.label_groups:89                inverted = {group[pair]: pair for pair in group}90                softmax = nn.Softmax(dim=1)91                output = softmax(logits[:, offset:offset + len(group)])92                classification = []93                for i, val in enumerate(output[0]):94                    classification.append((inverted[i], val.item()))95                classification.sort(key=lambda x: x[1], reverse=True)96                classifications.append(classification)97                offset += len(group)98 99        return HydraSequenceClassifierOutput(100            loss=loss,101            logits=logits,102            hidden_states=distilbert_output.hidden_states,103            attentions=distilbert_output.attentions,104            classifications=classifications105        )106 107    def to(self, device):108        super().to(device)109        self.pre_classifier.to(device)110        self.classifier.to(device)111        self.dropout.to(device)112        return self113