thealper2/graphcodebert-code-clone-detection
030
1"""GraphCodeBERT clone-detection model.2 3Reimplements the architecture from Microsoft's ``GraphCodeBERT/clonedetection``4on top of ``transformers`` v5:5 6* the two snippets are encoded **separately** by one shared GraphCodeBERT7 encoder, each with its own graph-guided masked attention;8* a data-flow node's input embedding is the average of the embeddings of the9 code tokens it was identified from;10* the two ``<s>`` representations are concatenated and fed to a11 ``Linear(2H -> H) -> tanh -> Linear(H -> 2)`` head.12 13The only real adaptation is the attention mask: ``transformers`` v5 builds masks14through ``masking_utils`` and only forwards a mask untouched when it is already154-D, so the boolean ``[B, L, L]`` graph mask is expanded to an additive16``[B, 1, L, L]`` mask here.17"""18 19from __future__ import annotations20 21import torch22import torch.nn as nn23from transformers import RobertaConfig, RobertaModel, RobertaPreTrainedModel24from transformers.modeling_outputs import SequenceClassifierOutput25 26__all__ = ["GraphCodeBERTForCloneDetection", "CloneClassificationHead"]27 28 29def _autocast_dtype(device: torch.device, fallback: torch.dtype) -> torch.dtype:30 """Dtype the attention scores will actually have, honouring autocast.31 32 SDPA requires an additive ``attn_mask`` whose dtype matches the query, so a33 hard-coded float32 mask would break under ``fp16=True``.34 """35 try:36 if torch.is_autocast_enabled(device.type):37 return torch.get_autocast_dtype(device.type)38 except TypeError: # older signature without a device argument39 if device.type == "cuda" and torch.is_autocast_enabled():40 return torch.get_autocast_gpu_dtype()41 return fallback42 43 44def _to_additive_mask(bool_mask: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:45 """``[B, L, L]`` boolean -> ``[B, 1, L, L]`` additive mask (0 / -inf)."""46 additive = torch.zeros(bool_mask.shape, dtype=dtype, device=bool_mask.device)47 additive.masked_fill_(~bool_mask, torch.finfo(dtype).min)48 return additive.unsqueeze(1)49 50 51class CloneClassificationHead(nn.Module):52 """Pairwise head over the two ``<s>`` vectors (GraphCodeBERT's own head)."""53 54 def __init__(self, config: RobertaConfig) -> None:55 super().__init__()56 self.dense = nn.Linear(config.hidden_size * 2, config.hidden_size)57 self.dropout = nn.Dropout(config.hidden_dropout_prob)58 self.out_proj = nn.Linear(config.hidden_size, 2)59 60 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:61 """``hidden_states``: ``[B*2, L, H]`` -> logits ``[B, 2]``."""62 x = hidden_states[:, 0, :] # <s> of each snippet63 x = x.reshape(-1, x.size(-1) * 2) # pair the two snippets back up64 x = self.dropout(x)65 x = torch.tanh(self.dense(x))66 x = self.dropout(x)67 return self.out_proj(x)68 69 70class GraphCodeBERTForCloneDetection(RobertaPreTrainedModel):71 """Binary clone classifier: ``0 = not clone``, ``1 = clone``."""72 73 config_class = RobertaConfig74 base_model_prefix = "roberta"75 supports_gradient_checkpointing = True76 77 def __init__(self, config: RobertaConfig, class_weights: list[float] | None = None) -> None:78 super().__init__(config)79 config.num_labels = 280 self.roberta = RobertaModel(config, add_pooling_layer=False)81 self.classifier = CloneClassificationHead(config)82 self.register_buffer(83 "class_weights",84 torch.tensor(class_weights, dtype=torch.float32) if class_weights else None,85 persistent=False,86 )87 self.post_init()88 89 # ------------------------------------------------------------------ #90 def _embed_with_dataflow(91 self, input_ids: torch.Tensor, position_idx: torch.Tensor, attn_mask: torch.Tensor92 ) -> torch.Tensor:93 """Word embeddings where each data-flow node averages its code tokens.94 95 ``position_idx`` encodes the role of every slot: ``0`` = data-flow node,96 ``1`` (= ``<pad>``) = padding, ``>= 2`` = real code token.97 """98 nodes_mask = position_idx.eq(0)99 token_mask = position_idx.ge(2)100 101 embeddings = self.roberta.embeddings.word_embeddings(input_ids)102 # For every node row, the code-token columns it may look at.103 nodes_to_token = nodes_mask[:, :, None] & token_mask[:, None, :] & attn_mask104 nodes_to_token = nodes_to_token.to(embeddings.dtype)105 nodes_to_token = nodes_to_token / (nodes_to_token.sum(-1) + 1e-10)[:, :, None]106 averaged = torch.einsum("abc,acd->abd", nodes_to_token, embeddings)107 return embeddings * (~nodes_mask)[:, :, None] + averaged * nodes_mask[:, :, None]108 109 def _encode(110 self, input_ids: torch.Tensor, position_idx: torch.Tensor, attn_mask: torch.Tensor111 ) -> torch.Tensor:112 embeddings = self._embed_with_dataflow(input_ids, position_idx, attn_mask)113 dtype = _autocast_dtype(input_ids.device, embeddings.dtype)114 outputs = self.roberta(115 inputs_embeds=embeddings,116 attention_mask=_to_additive_mask(attn_mask, dtype),117 position_ids=position_idx,118 token_type_ids=torch.zeros_like(position_idx),119 )120 return outputs.last_hidden_state121 122 # ------------------------------------------------------------------ #123 def forward(124 self,125 input_ids_1: torch.Tensor,126 position_idx_1: torch.Tensor,127 attn_mask_1: torch.Tensor,128 input_ids_2: torch.Tensor,129 position_idx_2: torch.Tensor,130 attn_mask_2: torch.Tensor,131 labels: torch.Tensor | None = None,132 ) -> SequenceClassifierOutput:133 """Encode both snippets with the shared encoder and classify the pair.134 135 Args:136 input_ids_*: ``[B, L]`` token ids; data-flow slots hold ``<unk>``.137 position_idx_*: ``[B, L]`` role/position ids (see ``_embed_with_dataflow``).138 attn_mask_*: ``[B, L, L]`` boolean graph-guided attention mask.139 labels: ``[B]`` with values in ``{0, 1}``.140 """141 batch_size, seq_len = input_ids_1.shape142 # Stack both snippets into one encoder call: [B, L] x2 -> [B*2, L].143 input_ids = torch.cat((input_ids_1[:, None], input_ids_2[:, None]), 1).view(-1, seq_len)144 position_idx = torch.cat((position_idx_1[:, None], position_idx_2[:, None]), 1).view(145 -1, seq_len146 )147 attn_mask = torch.cat((attn_mask_1[:, None], attn_mask_2[:, None]), 1).view(148 -1, seq_len, seq_len149 )150 151 hidden = self._encode(input_ids, position_idx, attn_mask)152 logits = self.classifier(hidden)153 154 loss = None155 if labels is not None:156 weight = None157 if self.class_weights is not None:158 weight = self.class_weights.to(device=logits.device, dtype=logits.dtype)159 loss = nn.functional.cross_entropy(logits, labels.view(-1), weight=weight)160 161 return SequenceClassifierOutput(loss=loss, logits=logits)162 163 164def load_model(165 model_name_or_path: str,166 attn_implementation: str = "sdpa",167 class_weights: list[float] | None = None,168 gradient_checkpointing: bool = False,169) -> GraphCodeBERTForCloneDetection:170 """Load GraphCodeBERT weights into the pairwise clone-detection head."""171 model = GraphCodeBERTForCloneDetection.from_pretrained(172 model_name_or_path,173 attn_implementation=attn_implementation,174 )175 # Set after loading: `from_pretrained` should not have to carry runtime-only176 # arguments, and the weights are a training artefact, not part of the config.177 model.class_weights = (178 torch.tensor(class_weights, dtype=torch.float32) if class_weights else None179 )180 if model.config.model_type != "roberta":181 raise ValueError(182 f"Expected a RoBERTa-architecture checkpoint (GraphCodeBERT), "183 f"got model_type={model.config.model_type!r}."184 )185 if gradient_checkpointing:186 model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})187 return model188 189 190def count_parameters(model: nn.Module) -> dict[str, int]:191 trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)192 total = sum(p.numel() for p in model.parameters())193 return {"trainable_parameters": trainable, "total_parameters": total}194 