codesage/codesage-large
25698
1#!/usr/bin/env python2# coding=utf-83# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.4 5import math6import torch7import torch.utils.checkpoint8from torch import nn9from torch.nn import CrossEntropyLoss, MSELoss, BCEWithLogitsLoss10from transformers.activations import ACT2FN11from transformers.modeling_utils import Conv1D, PreTrainedModel12from transformers.utils import logging13from .config_codesage import CodeSageConfig14from transformers.modeling_outputs import (15 BaseModelOutputWithPooling,16 MaskedLMOutput,17 SequenceClassifierOutput18)19 20logger = logging.get_logger(__name__)21 22CODESAGE_PRETRAINED_MODEL_ARCHIVE_LIST = [23 "codesage/codesage-small",24 "codesage/codesage-base",25 "codesage/codesage-large",26 # See all CodeSage models at https://huggingface.co/models?filter=codesage27]28 29 30class CodeSageAttention(nn.Module):31 def __init__(self, config):32 super().__init__()33 34 self.hidden_size = config.hidden_size35 self.num_heads = config.num_attention_heads36 self.head_dim = config.hidden_size // self.num_heads37 if self.head_dim * self.num_heads != config.hidden_size:38 raise ValueError(39 f"`hidden_size` must be divisible by num_heads "40 f"(got `hidden_size`: {config.hidden_size} and `num_heads`: {self.num_heads})."41 )42 43 self.c_attn = Conv1D(3 * self.hidden_size, self.hidden_size)44 self.c_proj = Conv1D(self.hidden_size, self.hidden_size)45 46 self.attention_dropout = nn.Dropout(config.attention_dropout_prob)47 self.residual_dropout = nn.Dropout(config.residual_dropout_prob)48 49 def attn(self, query, key, value, attention_mask=None, head_mask=None):50 attn_weights = torch.matmul(query, key.transpose(-1, -2))51 attn_weights = attn_weights / math.sqrt(self.head_dim)52 if attention_mask is not None:53 attn_weights = attn_weights + attention_mask54 55 attn_weights = nn.Softmax(dim=-1)(attn_weights)56 attn_weights = self.attention_dropout(attn_weights)57 if head_mask is not None:58 attn_weights = attn_weights * head_mask59 60 attn_output = torch.matmul(attn_weights, value)61 return attn_output, attn_weights62 63 def split_heads(self, tensor, num_heads, attn_head_size):64 """65 Splits hidden_size dim into attn_head_size and num_heads66 """67 new_shape = tensor.size()[:-1] + (num_heads, attn_head_size)68 tensor = tensor.view(*new_shape)69 return tensor.permute(0, 2, 1, 3) # (batch, head, seq_length, head_features)70 71 def merge_heads(self, tensor, num_heads, attn_head_size):72 """73 Merges attn_head_size dim and num_attn_heads dim into hidden_size74 """75 tensor = tensor.permute(0, 2, 1, 3).contiguous()76 new_shape = tensor.size()[:-2] + (num_heads * attn_head_size,)77 return tensor.view(new_shape)78 79 def forward(80 self,81 hidden_states,82 attention_mask=None,83 head_mask=None,84 output_attentions=False,85 ):86 query, key, value = self.c_attn(hidden_states).split(self.hidden_size, dim=2)87 query = self.split_heads(query, self.num_heads, self.head_dim)88 key = self.split_heads(key, self.num_heads, self.head_dim)89 value = self.split_heads(value, self.num_heads, self.head_dim)90 91 attn_output, attn_weights = self.attn(query, key, value, attention_mask, head_mask)92 93 attn_output = self.merge_heads(attn_output, self.num_heads, self.head_dim)94 attn_output = self.c_proj(attn_output)95 attn_output = self.residual_dropout(attn_output)96 97 outputs = (attn_output, attn_weights) if output_attentions else (attn_output,)98 return outputs # a, present, (attentions)99 100 101class CodeSageMLP(nn.Module):102 def __init__(self, intermediate_size, config):103 super().__init__()104 105 self.c_fc = Conv1D(intermediate_size, config.hidden_size)106 self.act = ACT2FN[config.activation_function]107 self.c_proj = Conv1D(config.hidden_size, intermediate_size)108 self.dropout = nn.Dropout(config.residual_dropout_prob)109 110 def forward(self, hidden_states):111 hidden_states = self.c_fc(hidden_states)112 hidden_states = self.act(hidden_states)113 hidden_states = self.c_proj(hidden_states)114 hidden_states = self.dropout(hidden_states)115 return hidden_states116 117 118class CodeSageBlock(nn.Module):119 def __init__(self, config):120 super().__init__()121 hidden_size = config.hidden_size122 inner_dim = config.intermediate_size if config.intermediate_size is not None else 4 * hidden_size123 self.ln_1 = nn.LayerNorm(hidden_size, eps=config.layer_norm_epsilon)124 self.attn = CodeSageAttention(config)125 self.ln_2 = nn.LayerNorm(hidden_size, eps=config.layer_norm_epsilon)126 self.mlp = CodeSageMLP(inner_dim, config)127 128 def forward(129 self,130 hidden_states,131 attention_mask=None,132 head_mask=None,133 output_attentions=False,134 ):135 residual = hidden_states136 hidden_states = self.ln_1(hidden_states)137 attn_outputs = self.attn(138 hidden_states,139 attention_mask=attention_mask,140 head_mask=head_mask,141 output_attentions=output_attentions142 )143 attn_output = attn_outputs[0] # output_attn: a, present, (attentions)144 outputs = attn_outputs[1:]145 hidden_states = attn_output + residual146 147 residual = hidden_states148 hidden_states = self.ln_2(hidden_states)149 feed_forward_hidden_states = self.mlp(hidden_states)150 hidden_states = residual + feed_forward_hidden_states151 152 outputs = (hidden_states,) + outputs[1:]153 return outputs # hidden_states, present, (attentions)154 155 156class CodeSagePreTrainedModel(PreTrainedModel):157 config_class = CodeSageConfig158 base_model_prefix = "transformer"159 160 def _init_weights(self, module):161 """Initialize the weights."""162 if isinstance(module, (nn.Linear, Conv1D)):163 module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)164 if module.bias is not None:165 module.bias.data.zero_()166 elif isinstance(module, nn.Embedding):167 module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)168 if module.padding_idx is not None:169 module.weight.data[module.padding_idx].zero_()170 elif isinstance(module, nn.LayerNorm):171 module.bias.data.zero_()172 module.weight.data.fill_(1.0)173 174 175class CodeSageModel(CodeSagePreTrainedModel):176 def __init__(self, config):177 super().__init__(config)178 179 self.wte = nn.Embedding(config.vocab_size, config.hidden_size)180 self.wpe = nn.Embedding(config.max_position_embeddings, config.hidden_size)181 182 self.drop = nn.Dropout(config.embedding_dropout_prob)183 self.h = nn.ModuleList([CodeSageBlock(config) for _ in range(config.num_hidden_layers)])184 self.ln_f = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_epsilon)185 186 self.init_weights()187 188 def get_input_embeddings(self):189 return self.wte190 191 def set_input_embeddings(self, new_embeddings: torch.Tensor):192 self.wte = new_embeddings193 194 def forward(195 self,196 input_ids=None,197 attention_mask=None,198 position_ids=None,199 head_mask=None,200 inputs_embeds=None,201 output_attentions=None,202 output_hidden_states=None,203 return_dict=None204 ):205 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions206 output_hidden_states = (207 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states208 )209 return_dict = return_dict if return_dict is not None else self.config.use_return_dict210 211 if input_ids is not None and inputs_embeds is not None:212 raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")213 if input_ids is not None:214 input_shape = input_ids.size()215 elif inputs_embeds is not None:216 input_shape = inputs_embeds.size()[:-1]217 else:218 raise ValueError("You have to specify either input_ids or inputs_embeds")219 220 device = input_ids.device if input_ids is not None else inputs_embeds.device221 if position_ids is None:222 position_ids = torch.arange(input_shape[-1], dtype=torch.long, device=device)223 position_ids = position_ids.unsqueeze(0).view(-1, input_shape[-1])224 else:225 position_ids = position_ids.view(-1, input_shape[-1])226 227 extended_attention_mask = None228 if attention_mask is not None:229 assert attention_mask.dim() == 2230 extended_attention_mask = attention_mask[:, None, None, :]231 extended_attention_mask = extended_attention_mask.to(dtype=self.dtype) # fp16 compatibility232 extended_attention_mask = (1.0 - extended_attention_mask) * -10000.0233 234 head_mask = self.get_head_mask(head_mask, self.config.num_hidden_layers)235 if inputs_embeds is None:236 inputs_embeds = self.wte(input_ids)237 238 position_embeds = self.wpe(position_ids)239 hidden_states = inputs_embeds + position_embeds240 241 hidden_states = self.drop(hidden_states)242 output_shape = input_shape + (hidden_states.size(-1),)243 244 all_self_attentions = () if output_attentions else None245 all_hidden_states = () if output_hidden_states else None246 for i, block in enumerate(self.h):247 if output_hidden_states:248 all_hidden_states = all_hidden_states + (hidden_states,)249 250 outputs = block(251 hidden_states,252 attention_mask=extended_attention_mask,253 head_mask=head_mask[i],254 output_attentions=output_attentions,255 )256 257 hidden_states = outputs[0]258 if output_attentions:259 all_self_attentions = all_self_attentions + (outputs[1],)260 261 hidden_states = self.ln_f(hidden_states)262 hidden_states = hidden_states.view(*output_shape)263 if output_hidden_states:264 all_hidden_states = all_hidden_states + (hidden_states,)265 266 pooled_output = None # max-pooled output267 if attention_mask is not None:268 pooled_output = (hidden_states * attention_mask[:, :, None]).sum(1) / attention_mask.sum(1)[:, None]269 270 if not return_dict:271 return tuple(272 v273 for v in [hidden_states, pooled_output, all_hidden_states, all_self_attentions]274 if v is not None275 )276 277 return BaseModelOutputWithPooling(278 last_hidden_state=hidden_states,279 pooler_output=pooled_output,280 hidden_states=all_hidden_states,281 attentions=all_self_attentions282 )283 284 285class CodeSageForMaskedLM(CodeSagePreTrainedModel):286 _tied_weights_keys = ["lm_head.weight"]287 288 def __init__(self, config):289 super().__init__(config)290 self.transformer = CodeSageModel(config)291 self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)292 293 self.init_weights()294 295 def get_output_embeddings(self):296 return self.lm_head297 298 def set_output_embeddings(self, new_embeddings):299 self.lm_head = new_embeddings300 301 def forward(302 self,303 input_ids=None,304 attention_mask=None,305 position_ids=None,306 head_mask=None,307 inputs_embeds=None,308 labels=None,309 output_attentions=None,310 output_hidden_states=None,311 return_dict=None312 ):313 return_dict = return_dict if return_dict is not None else self.config.use_return_dict314 315 transformer_outputs = self.transformer(316 input_ids,317 attention_mask=attention_mask,318 position_ids=position_ids,319 head_mask=head_mask,320 inputs_embeds=inputs_embeds,321 output_attentions=output_attentions,322 output_hidden_states=output_hidden_states,323 return_dict=return_dict324 )325 hidden_states = transformer_outputs[0]326 lm_logits = self.lm_head(hidden_states)327 328 masked_lm_loss = None329 if labels is not None:330 loss_fct = CrossEntropyLoss()331 masked_lm_loss = loss_fct(lm_logits.view(-1, lm_logits.size(-1)), labels.view(-1))332 333 if not return_dict:334 output = (lm_logits,) + transformer_outputs[1:]335 return ((masked_lm_loss,) + output) if masked_lm_loss is not None else output336 337 return MaskedLMOutput(338 loss=masked_lm_loss,339 logits=lm_logits,340 hidden_states=transformer_outputs.hidden_states,341 attentions=transformer_outputs.attentions,342 )343 344 345class CodeSageForSequenceClassification(CodeSagePreTrainedModel):346 347 def __init__(self, config):348 super().__init__(config)349 self.num_labels = config.num_labels350 self.config = config351 352 self.transformer = CodeSageModel(config)353 classifier_dropout = (354 config.classifier_dropout 355 if hasattr(config, 'classifier_dropout') and config.classifier_dropout is not None 356 else config.residual_dropout_prob357 )358 self.dropout = nn.Dropout(classifier_dropout)359 self.classifier = nn.Linear(config.hidden_size, config.num_labels)360 361 # Initialize weights and apply final processing362 self.post_init()363 364 def forward(365 self,366 input_ids=None,367 attention_mask=None,368 position_ids=None,369 head_mask=None,370 inputs_embeds=None,371 labels=None,372 output_attentions=None,373 output_hidden_states=None,374 return_dict=None,375 ):376 return_dict = return_dict if return_dict is not None else self.config.use_return_dict377 assert attention_mask is not None, "attention_mask is needed to perform max-pooling"378 379 outputs = self.transformer(380 input_ids,381 attention_mask=attention_mask,382 position_ids=position_ids,383 head_mask=head_mask,384 inputs_embeds=inputs_embeds,385 output_attentions=output_attentions,386 output_hidden_states=output_hidden_states,387 return_dict=return_dict,388 )389 390 pooled_output = outputs[1]391 pooled_output = self.dropout(pooled_output)392 logits = self.classifier(pooled_output)393 394 loss = None395 if labels is not None:396 if self.config.problem_type is None:397 if self.num_labels == 1:398 self.config.problem_type = "regression"399 elif self.num_labels > 1 and (labels.dtype == torch.long or labels.dtype == torch.int):400 self.config.problem_type = "single_label_classification"401 else:402 self.config.problem_type = "multi_label_classification"403 404 if self.config.problem_type == "regression":405 loss_fct = MSELoss()406 if self.num_labels == 1:407 loss = loss_fct(logits.squeeze(), labels.squeeze())408 else:409 loss = loss_fct(logits, labels)410 elif self.config.problem_type == "single_label_classification":411 loss_fct = CrossEntropyLoss()412 loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1))413 elif self.config.problem_type == "multi_label_classification":414 loss_fct = BCEWithLogitsLoss()415 loss = loss_fct(logits, labels)416 417 if not return_dict:418 output = (logits,) + outputs[2:]419 return ((loss,) + output) if loss is not None else output420 421 return SequenceClassifierOutput(422 loss=loss,423 logits=logits,424 hidden_states=outputs.hidden_states,425 attentions=outputs.attentions,426 )427 