File size: 8,339 Bytes
2ef4ea4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
"""GraphCodeBERT clone-detection model.

Reimplements the architecture from Microsoft's ``GraphCodeBERT/clonedetection``
on top of ``transformers`` v5:

* the two snippets are encoded **separately** by one shared GraphCodeBERT
  encoder, each with its own graph-guided masked attention;
* a data-flow node's input embedding is the average of the embeddings of the
  code tokens it was identified from;
* the two ``<s>`` representations are concatenated and fed to a
  ``Linear(2H -> H) -> tanh -> Linear(H -> 2)`` head.

The only real adaptation is the attention mask: ``transformers`` v5 builds masks
through ``masking_utils`` and only forwards a mask untouched when it is already
4-D, so the boolean ``[B, L, L]`` graph mask is expanded to an additive
``[B, 1, L, L]`` mask here.
"""

from __future__ import annotations

import torch
import torch.nn as nn
from transformers import RobertaConfig, RobertaModel, RobertaPreTrainedModel
from transformers.modeling_outputs import SequenceClassifierOutput

__all__ = ["GraphCodeBERTForCloneDetection", "CloneClassificationHead"]


def _autocast_dtype(device: torch.device, fallback: torch.dtype) -> torch.dtype:
    """Dtype the attention scores will actually have, honouring autocast.

    SDPA requires an additive ``attn_mask`` whose dtype matches the query, so a
    hard-coded float32 mask would break under ``fp16=True``.
    """
    try:
        if torch.is_autocast_enabled(device.type):
            return torch.get_autocast_dtype(device.type)
    except TypeError:  # older signature without a device argument
        if device.type == "cuda" and torch.is_autocast_enabled():
            return torch.get_autocast_gpu_dtype()
    return fallback


def _to_additive_mask(bool_mask: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
    """``[B, L, L]`` boolean -> ``[B, 1, L, L]`` additive mask (0 / -inf)."""
    additive = torch.zeros(bool_mask.shape, dtype=dtype, device=bool_mask.device)
    additive.masked_fill_(~bool_mask, torch.finfo(dtype).min)
    return additive.unsqueeze(1)


class CloneClassificationHead(nn.Module):
    """Pairwise head over the two ``<s>`` vectors (GraphCodeBERT's own head)."""

    def __init__(self, config: RobertaConfig) -> None:
        super().__init__()
        self.dense = nn.Linear(config.hidden_size * 2, config.hidden_size)
        self.dropout = nn.Dropout(config.hidden_dropout_prob)
        self.out_proj = nn.Linear(config.hidden_size, 2)

    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
        """``hidden_states``: ``[B*2, L, H]`` -> logits ``[B, 2]``."""
        x = hidden_states[:, 0, :]  # <s> of each snippet
        x = x.reshape(-1, x.size(-1) * 2)  # pair the two snippets back up
        x = self.dropout(x)
        x = torch.tanh(self.dense(x))
        x = self.dropout(x)
        return self.out_proj(x)


class GraphCodeBERTForCloneDetection(RobertaPreTrainedModel):
    """Binary clone classifier: ``0 = not clone``, ``1 = clone``."""

    config_class = RobertaConfig
    base_model_prefix = "roberta"
    supports_gradient_checkpointing = True

    def __init__(self, config: RobertaConfig, class_weights: list[float] | None = None) -> None:
        super().__init__(config)
        config.num_labels = 2
        self.roberta = RobertaModel(config, add_pooling_layer=False)
        self.classifier = CloneClassificationHead(config)
        self.register_buffer(
            "class_weights",
            torch.tensor(class_weights, dtype=torch.float32) if class_weights else None,
            persistent=False,
        )
        self.post_init()

    # ------------------------------------------------------------------ #
    def _embed_with_dataflow(
        self, input_ids: torch.Tensor, position_idx: torch.Tensor, attn_mask: torch.Tensor
    ) -> torch.Tensor:
        """Word embeddings where each data-flow node averages its code tokens.

        ``position_idx`` encodes the role of every slot: ``0`` = data-flow node,
        ``1`` (= ``<pad>``) = padding, ``>= 2`` = real code token.
        """
        nodes_mask = position_idx.eq(0)
        token_mask = position_idx.ge(2)

        embeddings = self.roberta.embeddings.word_embeddings(input_ids)
        # For every node row, the code-token columns it may look at.
        nodes_to_token = nodes_mask[:, :, None] & token_mask[:, None, :] & attn_mask
        nodes_to_token = nodes_to_token.to(embeddings.dtype)
        nodes_to_token = nodes_to_token / (nodes_to_token.sum(-1) + 1e-10)[:, :, None]
        averaged = torch.einsum("abc,acd->abd", nodes_to_token, embeddings)
        return embeddings * (~nodes_mask)[:, :, None] + averaged * nodes_mask[:, :, None]

    def _encode(
        self, input_ids: torch.Tensor, position_idx: torch.Tensor, attn_mask: torch.Tensor
    ) -> torch.Tensor:
        embeddings = self._embed_with_dataflow(input_ids, position_idx, attn_mask)
        dtype = _autocast_dtype(input_ids.device, embeddings.dtype)
        outputs = self.roberta(
            inputs_embeds=embeddings,
            attention_mask=_to_additive_mask(attn_mask, dtype),
            position_ids=position_idx,
            token_type_ids=torch.zeros_like(position_idx),
        )
        return outputs.last_hidden_state

    # ------------------------------------------------------------------ #
    def forward(
        self,
        input_ids_1: torch.Tensor,
        position_idx_1: torch.Tensor,
        attn_mask_1: torch.Tensor,
        input_ids_2: torch.Tensor,
        position_idx_2: torch.Tensor,
        attn_mask_2: torch.Tensor,
        labels: torch.Tensor | None = None,
    ) -> SequenceClassifierOutput:
        """Encode both snippets with the shared encoder and classify the pair.

        Args:
            input_ids_*: ``[B, L]`` token ids; data-flow slots hold ``<unk>``.
            position_idx_*: ``[B, L]`` role/position ids (see ``_embed_with_dataflow``).
            attn_mask_*: ``[B, L, L]`` boolean graph-guided attention mask.
            labels: ``[B]`` with values in ``{0, 1}``.
        """
        batch_size, seq_len = input_ids_1.shape
        # Stack both snippets into one encoder call: [B, L] x2 -> [B*2, L].
        input_ids = torch.cat((input_ids_1[:, None], input_ids_2[:, None]), 1).view(-1, seq_len)
        position_idx = torch.cat((position_idx_1[:, None], position_idx_2[:, None]), 1).view(
            -1, seq_len
        )
        attn_mask = torch.cat((attn_mask_1[:, None], attn_mask_2[:, None]), 1).view(
            -1, seq_len, seq_len
        )

        hidden = self._encode(input_ids, position_idx, attn_mask)
        logits = self.classifier(hidden)

        loss = None
        if labels is not None:
            weight = None
            if self.class_weights is not None:
                weight = self.class_weights.to(device=logits.device, dtype=logits.dtype)
            loss = nn.functional.cross_entropy(logits, labels.view(-1), weight=weight)

        return SequenceClassifierOutput(loss=loss, logits=logits)


def load_model(
    model_name_or_path: str,
    attn_implementation: str = "sdpa",
    class_weights: list[float] | None = None,
    gradient_checkpointing: bool = False,
) -> GraphCodeBERTForCloneDetection:
    """Load GraphCodeBERT weights into the pairwise clone-detection head."""
    model = GraphCodeBERTForCloneDetection.from_pretrained(
        model_name_or_path,
        attn_implementation=attn_implementation,
    )
    # Set after loading: `from_pretrained` should not have to carry runtime-only
    # arguments, and the weights are a training artefact, not part of the config.
    model.class_weights = (
        torch.tensor(class_weights, dtype=torch.float32) if class_weights else None
    )
    if model.config.model_type != "roberta":
        raise ValueError(
            f"Expected a RoBERTa-architecture checkpoint (GraphCodeBERT), "
            f"got model_type={model.config.model_type!r}."
        )
    if gradient_checkpointing:
        model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
    return model


def count_parameters(model: nn.Module) -> dict[str, int]:
    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
    total = sum(p.numel() for p in model.parameters())
    return {"trainable_parameters": trainable, "total_parameters": total}