Team Ai
Modelpublic

thealper2/graphcodebert-code-clone-detection

sourceHugging Facemitupdated 17d agoView on Hugging Face
0likes30downloads
modeling.py194 linesDownload Raw Back to code
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