DataRaptor/HateGuard
0
1from torch import nn2from transformers import AutoConfig, AutoModel, AutoTokenizer3import torch4 5 6def weight_init_normal(module, model):7 if isinstance(module, nn.Linear):8 module.weight.data.normal_(mean=0.0, std=model.config.initializer_range)9 if module.bias is not None:10 module.bias.data.zero_()11 elif isinstance(module, nn.Embedding):12 module.weight.data.normal_(mean=0.0, std=model.config.initializer_range)13 if module.padding_idx is not None:14 module.weight.data[module.padding_idx].zero_()15 elif isinstance(module, nn.LayerNorm):16 module.bias.data.zero_()17 module.weight.data.fill_(1.0)18 19 20 21class MeanPooling(nn.Module):22 def __init__(self):23 super(MeanPooling, self).__init__()24 25 def forward(self, last_hidden_state, attention_mask):26 input_mask_expanded = attention_mask.unsqueeze(-1).expand(last_hidden_state.size()).float()27 sum_embeddings = torch.sum(last_hidden_state * input_mask_expanded, 1)28 sum_mask = input_mask_expanded.sum(1)29 sum_mask = torch.clamp(sum_mask, min=1e-9)30 mean_embeddings = sum_embeddings / sum_mask31 return mean_embeddings32 33 34class MeanPoolingLayer(nn.Module):35 def __init__(self, 36 hidden_size,37 target_size,38 dropout = 0,39 ):40 super(MeanPoolingLayer, self).__init__()41 self.pool = MeanPooling()42 self.fc = nn.Sequential(43 nn.Dropout(dropout),44 nn.Linear(hidden_size, target_size),45 nn.Sigmoid()46 )47 48 def forward(self, inputs, mask):49 last_hidden_states = inputs[0]50 feature = self.pool(last_hidden_states, mask)51 outputs = self.fc(feature)52 return outputs53 54 55 56class HSLanguageModel(nn.Module):57 def __init__(self,58 backbone = 'microsoft/deberta-v3-small',59 target_size = 1,60 head_dropout = 0,61 reinit_nlayers = 0,62 freeze_nlayers = 0,63 reinit_head = True,64 grad_checkpointing = False,65 ):66 super(HSLanguageModel, self).__init__()67 68 self.config = AutoConfig.from_pretrained(backbone, output_hidden_states=True)69 self.model = AutoModel.from_pretrained(backbone, config=self.config)70 self.head = MeanPoolingLayer(71 self.config.hidden_size,72 target_size,73 head_dropout74 )75 self.tokenizer = AutoTokenizer.from_pretrained(backbone);76 77 78 if grad_checkpointing == True:79 print('Gradient ckpt enabled')80 self.model.gradient_checkpointing_enable()81 82 if reinit_nlayers > 0:83 # Reinit last n encoder layers84 # [TODO] Check if it is autoencoding model: Bert, Roberta, DistilBert, Albert, XLMRoberta, BertModel85 for layer in self.model.encoder.layer[-reinit_nlayers:]: 86 self._init_weights(layer)87 88 if freeze_nlayers > 0:89 self.model.embeddings.requires_grad_(False)90 self.model.encoder.layer[:freeze_nlayers].requires_grad_(False)91 92 if reinit_head:93 # Reinit layers in head94 self._init_weights(self.head)95 96 97 def _init_weights(self, layer):98 for module in layer.modules():99 init_fn = weight_init_normal100 init_fn(module, self)101 102 103 def forward(self, inputs):104 outputs = self.model(**inputs)105 outputs = self.head(outputs, inputs['attention_mask'])106 return outputs107 108 109if __name__ == '__main__':110 111 model = HSLanguageModel()112 113 114 115 116 