ellenhp/query2osm
0
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 