thealper2/graphcodebert-code-clone-detection
030
1"""Dataset loading, leakage-safe splitting and GraphCodeBERT feature building.2 3Design notes4------------5The PoolC fold contains 6.7 M *pairs* but only ~45 k *distinct snippets* -- the6pairs are combinations of a small pool of solutions. So the expensive work7(tree-sitter parsing, data-flow extraction, BPE tokenisation) is done **once per8distinct snippet** and cached on disk; a pair is then just two integer indices9plus a label. Nothing proportional to 6.7 M rows is ever tokenised, and the10graph-guided attention masks are materialised lazily in the collator.11"""12 13from __future__ import annotations14 15import hashlib16import json17import logging18import time19from collections import Counter20from dataclasses import dataclass21from pathlib import Path22from typing import Any, Iterator23 24import numpy as np25import pyarrow.parquet as pq26import torch27from torch.utils.data import Dataset28 29from config import (30 CODE1_COLUMN,31 CODE2_COLUMN,32 DATASET_LANGUAGE,33 FORBIDDEN_FEATURE_COLUMNS,34 GROUP1_COLUMN,35 GROUP2_COLUMN,36 LABEL_COLUMN,37 Config,38)39from dfg_parser import DataFlowExtractionError, extract_dataflow40 41logger = logging.getLogger(__name__)42 43_META_COLUMNS = [LABEL_COLUMN, GROUP1_COLUMN, GROUP2_COLUMN]44 45 46# --------------------------------------------------------------------------- #47# 1. Schema verification48# --------------------------------------------------------------------------- #49def verify_dataset_schema(dataset_name: str) -> dict[str, Any]:50 """Check the remote dataset still matches what this pipeline expects.51 52 Raises immediately (rather than silently mis-reading columns) if the schema53 drifts. Returns a report dict for the experiment log.54 """55 from datasets import load_dataset_builder56 57 builder = load_dataset_builder(dataset_name)58 features = builder.info.features59 splits = {k: v.num_examples for k, v in (builder.info.splits or {}).items()}60 61 missing = [62 c63 for c in (CODE1_COLUMN, CODE2_COLUMN, LABEL_COLUMN, GROUP1_COLUMN, GROUP2_COLUMN)64 if c not in features65 ]66 if missing:67 raise ValueError(68 f"{dataset_name} is missing expected columns {missing}. "69 f"Available: {sorted(features)}. Adapt config.py before continuing."70 )71 for col in (CODE1_COLUMN, CODE2_COLUMN):72 if features[col].dtype != "string":73 raise ValueError(f"Column {col!r} must be a string, got {features[col]}.")74 75 report = {76 "dataset_name": dataset_name,77 "features": {k: str(v) for k, v in features.items()},78 "splits": splits,79 "forbidden_feature_columns": list(FORBIDDEN_FEATURE_COLUMNS),80 "language": DATASET_LANGUAGE,81 }82 logger.info("Dataset schema OK: %s", json.dumps(report["splits"]))83 return report84 85 86def _parquet_files(dataset_name: str, split: str) -> list[str]:87 """Resolve the local parquet shards for one split (downloads on first use)."""88 from huggingface_hub import snapshot_download89 90 root = Path(91 snapshot_download(dataset_name, repo_type="dataset", allow_patterns=["data/*", "*.json"])92 )93 files = sorted(root.glob(f"data/{split}-*.parquet"))94 if not files:95 files = sorted(root.glob(f"**/{split}-*.parquet"))96 if not files:97 raise FileNotFoundError(98 f"No parquet shards for split {split!r} under {root}. "99 f"Found: {[p.name for p in root.rglob('*.parquet')]}"100 )101 return [str(p) for p in files]102 103 104# --------------------------------------------------------------------------- #105# 2. Snippet pool + pair index (cached)106# --------------------------------------------------------------------------- #107@dataclass108class SplitIndex:109 """A split reduced to integer indices into the shared snippet pool."""110 111 name: str112 snippet_id1: np.ndarray # int32 [n_pairs]113 snippet_id2: np.ndarray # int32 [n_pairs]114 labels: np.ndarray # int8 [n_pairs]115 group1: np.ndarray # int32 [n_pairs]116 group2: np.ndarray # int32 [n_pairs]117 118 def __len__(self) -> int:119 return int(self.labels.shape[0])120 121 def select(self, rows: np.ndarray, name: str | None = None) -> "SplitIndex":122 return SplitIndex(123 name=name or self.name,124 snippet_id1=self.snippet_id1[rows],125 snippet_id2=self.snippet_id2[rows],126 labels=self.labels[rows],127 group1=self.group1[rows],128 group2=self.group2[rows],129 )130 131 def class_distribution(self) -> dict[str, Any]:132 counts = Counter(self.labels.tolist())133 n = max(len(self), 1)134 return {135 "num_examples": len(self),136 "negatives_label_0": int(counts.get(0, 0)),137 "positives_label_1": int(counts.get(1, 0)),138 "positive_ratio": round(counts.get(1, 0) / n, 6),139 "num_groups": int(np.unique(np.concatenate([self.group1, self.group2])).size),140 }141 142 def snippet_ids(self) -> np.ndarray:143 return np.unique(np.concatenate([self.snippet_id1, self.snippet_id2]))144 145 146def _hash_code(text: str) -> bytes:147 return hashlib.blake2b(text.encode("utf-8", "ignore"), digest_size=16).digest()148 149 150def build_snippet_pool(151 cfg: Config, splits: tuple[str, ...]152) -> tuple[list[str], dict[str, SplitIndex], dict[str, Any]]:153 """Deduplicate every snippet across ``splits`` and index the pairs.154 155 Cached under ``cfg.cache_dir`` -- the scan over the parquet shards is the156 only pass that ever touches all 6.7 M rows, and it runs once.157 """158 cache = Path(cfg.cache_dir) / "pairs" / _fingerprint(cfg.dataset_name, splits)159 if (cache / "meta.json").exists():160 logger.info("Reusing cached snippet pool at %s", cache)161 return _load_pool(cache)162 163 cache.mkdir(parents=True, exist_ok=True)164 t0 = time.time()165 code_to_id: dict[bytes, int] = {}166 snippets: list[str] = []167 indices: dict[str, SplitIndex] = {}168 #: A snippet appearing under two different group ids is a (rare) duplicate169 #: solution; we log it because it is the only within-split ambiguity.170 snippet_groups: dict[int, set[int]] = {}171 172 for split in splits:173 files = _parquet_files(cfg.dataset_name, split)174 id1_chunks, id2_chunks, lab_chunks, g1_chunks, g2_chunks = [], [], [], [], []175 for path in files:176 pf = pq.ParquetFile(path)177 for batch in pf.iter_batches(178 batch_size=50_000,179 columns=[CODE1_COLUMN, CODE2_COLUMN, *_META_COLUMNS],180 ):181 cols = batch.to_pydict()182 for code_col, group_col, out in (183 (CODE1_COLUMN, GROUP1_COLUMN, id1_chunks),184 (CODE2_COLUMN, GROUP2_COLUMN, id2_chunks),185 ):186 ids = np.empty(len(cols[code_col]), dtype=np.int32)187 for i, (text, group) in enumerate(zip(cols[code_col], cols[group_col])):188 key = _hash_code(text)189 sid = code_to_id.get(key)190 if sid is None:191 sid = len(snippets)192 code_to_id[key] = sid193 snippets.append(text)194 ids[i] = sid195 snippet_groups.setdefault(sid, set()).add(int(group))196 out.append(ids)197 lab_chunks.append(np.asarray(cols[LABEL_COLUMN], dtype=np.int8))198 g1_chunks.append(np.asarray(cols[GROUP1_COLUMN], dtype=np.int32))199 g2_chunks.append(np.asarray(cols[GROUP2_COLUMN], dtype=np.int32))200 logger.info(" scanned %s", Path(path).name)201 202 idx = SplitIndex(203 name=split,204 snippet_id1=np.concatenate(id1_chunks),205 snippet_id2=np.concatenate(id2_chunks),206 labels=np.concatenate(lab_chunks),207 group1=np.concatenate(g1_chunks),208 group2=np.concatenate(g2_chunks),209 )210 _assert_label_matches_groups(idx)211 indices[split] = idx212 logger.info("Split %s: %s", split, json.dumps(idx.class_distribution()))213 214 ambiguous = sorted(sid for sid, gs in snippet_groups.items() if len(gs) > 1)215 leakage = _cross_split_leakage(indices)216 stats = {217 "num_unique_snippets": len(snippets),218 "scan_seconds": round(time.time() - t0, 1),219 "snippets_with_multiple_groups": len(ambiguous),220 "cross_split_snippet_overlap": leakage,221 "per_split": {k: v.class_distribution() for k, v in indices.items()},222 }223 _save_pool(cache, snippets, indices, stats)224 logger.info("Snippet pool: %s", json.dumps(stats, indent=2))225 return snippets, indices, stats226 227 228def _assert_label_matches_groups(idx: SplitIndex) -> None:229 """The label is exactly ``group1 == group2``; assert it and shout about it.230 231 This is why ``code1_group``/``code2_group`` are on the forbidden list: they232 are a perfect proxy for the target.233 """234 implied = (idx.group1 == idx.group2).astype(np.int8)235 mismatches = int((implied != idx.labels).sum())236 if mismatches:237 raise ValueError(238 f"Split {idx.name}: {mismatches} rows where `similar` disagrees with "239 "(code1_group == code2_group). The dataset changed; revisit the split logic."240 )241 242 243def _cross_split_leakage(indices: dict[str, SplitIndex]) -> dict[str, int]:244 """Count snippets and groups shared between splits (must be zero)."""245 out: dict[str, int] = {}246 names = list(indices)247 for i, a in enumerate(names):248 for b in names[i + 1 :]:249 sa, sb = set(indices[a].snippet_ids().tolist()), set(indices[b].snippet_ids().tolist())250 ga = set(np.unique(np.concatenate([indices[a].group1, indices[a].group2])).tolist())251 gb = set(np.unique(np.concatenate([indices[b].group1, indices[b].group2])).tolist())252 out[f"{a}|{b}:snippets"] = len(sa & sb)253 out[f"{a}|{b}:groups"] = len(ga & gb)254 return out255 256 257def _fingerprint(*parts: Any) -> str:258 return hashlib.blake2b(repr(parts).encode(), digest_size=8).hexdigest()259 260 261def _save_pool(262 cache: Path, snippets: list[str], indices: dict[str, SplitIndex], stats: dict263) -> None:264 import pyarrow as pa265 266 pq.write_table(pa.table({"code": snippets}), cache / "snippets.parquet")267 for name, idx in indices.items():268 np.savez(269 cache / f"{name}.npz",270 snippet_id1=idx.snippet_id1,271 snippet_id2=idx.snippet_id2,272 labels=idx.labels,273 group1=idx.group1,274 group2=idx.group2,275 )276 (cache / "meta.json").write_text(277 json.dumps({"splits": list(indices), "stats": stats}, indent=2), encoding="utf-8"278 )279 280 281def _load_pool(cache: Path) -> tuple[list[str], dict[str, SplitIndex], dict[str, Any]]:282 meta = json.loads((cache / "meta.json").read_text(encoding="utf-8"))283 snippets = pq.read_table(cache / "snippets.parquet").column("code").to_pylist()284 indices = {}285 for name in meta["splits"]:286 z = np.load(cache / f"{name}.npz")287 indices[name] = SplitIndex(288 name=name,289 snippet_id1=z["snippet_id1"],290 snippet_id2=z["snippet_id2"],291 labels=z["labels"],292 group1=z["group1"],293 group2=z["group2"],294 )295 return snippets, indices, meta["stats"]296 297 298# --------------------------------------------------------------------------- #299# 3. Leakage-safe splitting300# --------------------------------------------------------------------------- #301def split_heldout_by_group(302 heldout: SplitIndex, test_group_fraction: float, seed: int303) -> tuple[SplitIndex, SplitIndex, dict[str, Any]]:304 """Partition the held-out split into validation/test along *group* boundaries.305 306 The dataset ships only ``train`` and ``val``; ``val`` is carved into a307 validation and a test half by assigning whole problem groups to one side.308 Pairs whose two snippets straddle the boundary are dropped -- keeping them309 would put the same group on both sides.310 """311 groups = np.unique(np.concatenate([heldout.group1, heldout.group2]))312 rng = np.random.default_rng(seed)313 shuffled = groups.copy()314 rng.shuffle(shuffled)315 n_test = max(1, int(round(len(shuffled) * test_group_fraction)))316 if n_test >= len(shuffled):317 raise ValueError("test_group_fraction leaves no groups for validation.")318 test_groups = set(shuffled[:n_test].tolist())319 val_groups = set(shuffled[n_test:].tolist())320 321 in_test = np.isin(heldout.group1, list(test_groups)) & np.isin(322 heldout.group2, list(test_groups)323 )324 in_val = np.isin(heldout.group1, list(val_groups)) & np.isin(heldout.group2, list(val_groups))325 dropped = int(len(heldout) - in_test.sum() - in_val.sum())326 327 validation = heldout.select(np.flatnonzero(in_val), name="validation")328 test = heldout.select(np.flatnonzero(in_test), name="test")329 330 overlap = set(validation.snippet_ids().tolist()) & set(test.snippet_ids().tolist())331 if overlap:332 raise AssertionError(f"{len(overlap)} snippets leaked between validation and test.")333 334 report = {335 "heldout_groups": int(len(groups)),336 "validation_groups": len(val_groups),337 "test_groups": len(test_groups),338 "dropped_cross_boundary_pairs": dropped,339 "validation": validation.class_distribution(),340 "test": test.class_distribution(),341 }342 return validation, test, report343 344 345def subsample(346 idx: SplitIndex, max_samples: int, seed: int, balanced: bool = True347) -> tuple[SplitIndex, dict[str, Any]]:348 """Take at most ``max_samples`` rows, optionally keeping the classes balanced.349 350 Subsampling never crosses group boundaries (it only removes rows), so it351 cannot introduce leakage.352 """353 if max_samples < 0 or max_samples >= len(idx):354 return idx, {"subsampled": False, "kept": len(idx)}355 356 rng = np.random.default_rng(seed)357 if balanced:358 per_class = max_samples // 2359 chosen = []360 for label in (0, 1):361 rows = np.flatnonzero(idx.labels == label)362 take = min(per_class, len(rows))363 chosen.append(rng.choice(rows, size=take, replace=False))364 rows = np.sort(np.concatenate(chosen))365 else:366 rows = np.sort(rng.choice(len(idx), size=max_samples, replace=False))367 368 out = idx.select(rows)369 return out, {"subsampled": True, "kept": len(out), "balanced": balanced}370 371 372def decide_class_weights(373 train: SplitIndex, mode: str, threshold: float374) -> tuple[list[float] | None, dict[str, Any]]:375 """Decide whether class-weighted cross entropy is warranted.376 377 Weighting is *not* applied by default: it is enabled only when the measured378 majority-class share exceeds ``threshold``. The rationale is recorded in the379 returned report and written to ``training_config.json``.380 """381 dist = train.class_distribution()382 n0, n1 = dist["negatives_label_0"], dist["positives_label_1"]383 total = max(n0 + n1, 1)384 majority_share = max(n0, n1) / total385 386 if mode == "off":387 apply = False388 reason = "class_weighting=off (forced by config)."389 elif mode == "on":390 apply = True391 reason = "class_weighting=on (forced by config)."392 else:393 apply = majority_share > threshold394 reason = (395 f"Measured majority-class share {majority_share:.4f} "396 f"{'exceeds' if apply else 'is within'} the {threshold} threshold, "397 f"so weighted cross entropy is {'enabled' if apply else 'NOT used'}."398 )399 400 weights = None401 if apply and n0 > 0 and n1 > 0:402 # Inverse-frequency weights normalised to mean 1.403 w = np.array([total / (2 * n0), total / (2 * n1)], dtype=np.float64)404 weights = (w / w.mean()).tolist()405 406 report = {407 "mode": mode,408 "threshold": threshold,409 "majority_class_share": round(majority_share, 6),410 "applied": weights is not None,411 "weights": weights,412 "reason": reason,413 }414 logger.info("Class weighting decision: %s", reason)415 return weights, report416 417 418# --------------------------------------------------------------------------- #419# 4. GraphCodeBERT snippet features420# --------------------------------------------------------------------------- #421@dataclass422class SnippetFeatures:423 """Pre-tokenised snippet pool, laid out as flat numpy arrays.424 425 ``dfg_adj_*`` store the (ragged) node adjacency lists so that no data-flow426 edge is ever clipped away silently.427 """428 429 input_ids: np.ndarray # int32 [N, L]430 position_idx: np.ndarray # int16 [N, L]431 dfg_to_code: np.ndarray # int32 [N, max_nodes, 2] (offsets into the432 #: *untruncated* sub-token stream, so they can exceed the sequence length)433 num_nodes: np.ndarray # int16 [N]434 node_index: np.ndarray # int16 [N] number of real code tokens (incl. <s>/</s>)435 max_length: np.ndarray # int16 [N] code tokens + data-flow nodes436 dfg_adj_values: np.ndarray # int16 [total_edges]437 dfg_adj_offsets: np.ndarray # int64 [N, max_nodes + 1]438 seq_length: int439 stats: dict[str, Any]440 441 def __len__(self) -> int:442 return int(self.input_ids.shape[0])443 444 445def build_snippet_features(446 cfg: Config, snippets: list[str], tokenizer: Any, num_proc: int | None = None447) -> SnippetFeatures:448 """Run data-flow extraction + tokenisation over every distinct snippet.449 450 Uses ``datasets.map`` (batched, multi-process, Arrow-cached) so that a rerun451 with the same config costs nothing.452 """453 from datasets import Dataset as HFDataset454 455 seq_len = cfg.total_sequence_length456 max_nodes = seq_len - 3 # hard upper bound; the real cap is computed per snippet457 cache = Path(cfg.cache_dir) / "features"458 cache.mkdir(parents=True, exist_ok=True)459 key = _fingerprint(460 cfg.dataset_name, cfg.model_name_or_path, cfg.code_length, cfg.data_flow_length, len(snippets)461 )462 npz_path = cache / f"snippets_{key}.npz"463 stats_path = cache / f"snippets_{key}.stats.json"464 465 if npz_path.exists() and stats_path.exists():466 logger.info("Reusing cached snippet features at %s", npz_path)467 z = np.load(npz_path)468 return SnippetFeatures(469 input_ids=z["input_ids"],470 position_idx=z["position_idx"],471 dfg_to_code=z["dfg_to_code"],472 num_nodes=z["num_nodes"],473 node_index=z["node_index"],474 max_length=z["max_length"],475 dfg_adj_values=z["dfg_adj_values"],476 dfg_adj_offsets=z["dfg_adj_offsets"],477 seq_length=seq_len,478 stats=json.loads(stats_path.read_text(encoding="utf-8")),479 )480 481 ds = HFDataset.from_dict({"code": snippets})482 num_proc = num_proc if num_proc and num_proc > 1 else None483 t0 = time.time()484 ds = ds.map(485 _make_feature_fn(tokenizer, cfg.code_length, cfg.data_flow_length),486 batched=True,487 batch_size=256,488 num_proc=num_proc,489 remove_columns=["code"],490 desc="GraphCodeBERT data-flow + tokenisation",491 )492 493 n = len(ds)494 cols = ds.with_format(None)495 status_counter: Counter[str] = Counter()496 n_nodes_all: list[int] = []497 498 input_ids = np.full((n, seq_len), tokenizer.pad_token_id, dtype=np.int32)499 position_idx = np.full((n, seq_len), tokenizer.pad_token_id, dtype=np.int16)500 num_nodes = np.zeros(n, dtype=np.int16)501 node_index = np.zeros(n, dtype=np.int16)502 max_length = np.zeros(n, dtype=np.int16)503 dfg_to_code_list: list[list[list[int]]] = []504 adj_values: list[int] = []505 adj_offsets = np.zeros((n, max_nodes + 1), dtype=np.int64)506 507 observed_max_nodes = 0508 for i, row in enumerate(cols):509 ids = row["input_ids"]510 pos = row["position_idx"]511 input_ids[i, : len(ids)] = ids512 position_idx[i, : len(pos)] = pos513 d2c = row["dfg_to_code"]514 adj = row["dfg_to_dfg"]515 num_nodes[i] = len(d2c)516 observed_max_nodes = max(observed_max_nodes, len(d2c))517 node_index[i] = int(np.sum(np.asarray(pos) > 1))518 max_length[i] = int(np.sum(np.asarray(pos) != tokenizer.pad_token_id))519 dfg_to_code_list.append(d2c)520 base = len(adj_values)521 adj_offsets[i, 0] = base522 for j, nb in enumerate(adj):523 adj_values.extend(nb)524 adj_offsets[i, j + 1] = len(adj_values)525 adj_offsets[i, len(adj) + 1 :] = len(adj_values)526 status_counter.update(row["status_flags"])527 n_nodes_all.append(len(d2c))528 529 observed_max_nodes = max(observed_max_nodes, 1)530 # int32: these are offsets into the untruncated sub-token stream, which for a531 # pathologically long snippet reaches ~1e5 -- well past int16.532 dfg_to_code = np.zeros((n, observed_max_nodes, 2), dtype=np.int32)533 for i, d2c in enumerate(dfg_to_code_list):534 if d2c:535 dfg_to_code[i, : len(d2c)] = np.asarray(d2c, dtype=np.int32)536 adj_offsets = adj_offsets[:, : observed_max_nodes + 1]537 538 nodes_arr = np.asarray(n_nodes_all)539 stats = {540 "num_snippets": n,541 "extraction_seconds": round(time.time() - t0, 1),542 "status_counts": dict(status_counter),543 "snippets_with_empty_dataflow": int((nodes_arr == 0).sum()),544 "dataflow_nodes_mean": round(float(nodes_arr.mean()), 2),545 "dataflow_nodes_p50": int(np.percentile(nodes_arr, 50)),546 "dataflow_nodes_p95": int(np.percentile(nodes_arr, 95)),547 "dataflow_nodes_max": int(nodes_arr.max()),548 "code_tokens_mean": round(float(node_index.mean()), 2),549 "code_tokens_truncated": int((node_index >= cfg.code_length - 1).sum()),550 "total_dataflow_edges": len(adj_values),551 "sequence_length": seq_len,552 }553 logger.info("Snippet features: %s", json.dumps(stats, indent=2))554 if status_counter.get("dfg_failed", 0) or status_counter.get("dfg_recursion_limit", 0):555 logger.warning(556 "Data-flow extraction degraded for %d snippets (kept with an empty graph, not dropped).",557 status_counter.get("dfg_failed", 0) + status_counter.get("dfg_recursion_limit", 0),558 )559 560 features = SnippetFeatures(561 input_ids=input_ids,562 position_idx=position_idx,563 dfg_to_code=dfg_to_code,564 num_nodes=num_nodes,565 node_index=node_index,566 max_length=max_length,567 dfg_adj_values=np.asarray(adj_values, dtype=np.int16),568 dfg_adj_offsets=adj_offsets,569 seq_length=seq_len,570 stats=stats,571 )572 np.savez_compressed(573 npz_path,574 input_ids=features.input_ids,575 position_idx=features.position_idx,576 dfg_to_code=features.dfg_to_code,577 num_nodes=features.num_nodes,578 node_index=features.node_index,579 max_length=features.max_length,580 dfg_adj_values=features.dfg_adj_values,581 dfg_adj_offsets=features.dfg_adj_offsets,582 )583 stats_path.write_text(json.dumps(stats, indent=2), encoding="utf-8")584 return features585 586 587def _make_feature_fn(tokenizer: Any, code_length: int, data_flow_length: int):588 """Build the batched ``datasets.map`` function (picklable via closure)."""589 seq_len = code_length + data_flow_length590 cls_id, sep_id = tokenizer.cls_token_id, tokenizer.sep_token_id591 pad_id, unk_id = tokenizer.pad_token_id, tokenizer.unk_token_id592 593 def fn(batch: dict[str, list]) -> dict[str, list]:594 out_ids, out_pos, out_d2c, out_d2d, out_status = [], [], [], [], []595 for code in batch["code"]:596 flags: list[str] = []597 try:598 code_tokens, dfg, status = extract_dataflow(code, DATASET_LANGUAGE)599 except DataFlowExtractionError as exc:600 # Never drop the example: fall back to a plain tokenisation with601 # no data-flow component, and record the reason.602 logger.warning("Data-flow extraction failed, using empty graph: %s", exc)603 code_tokens, dfg = code.split(), []604 status = {"comment_strip": "n/a", "parse": "failed", "dfg": "failed", "error": str(exc)}605 for stage in ("comment_strip", "parse", "dfg"):606 if status[stage] not in ("ok", "n/a"):607 flags.append(f"{stage}_{status[stage]}")608 if not flags:609 flags.append("ok")610 611 # GraphCodeBERT tokenises each code token separately; the '@ ' prefix612 # trick forces a word-boundary BPE split for non-initial tokens.613 sub_tokens = [614 tokenizer.tokenize("@ " + t)[1:] if i != 0 else tokenizer.tokenize(t)615 for i, t in enumerate(code_tokens)616 ]617 ori2cur = {-1: (0, 0)}618 for i in range(len(sub_tokens)):619 prev_end = ori2cur[i - 1][1]620 ori2cur[i] = (prev_end, prev_end + len(sub_tokens[i]))621 flat = [y for x in sub_tokens for y in x]622 623 # Reserve room for the data-flow nodes, then for <s>/</s>.624 keep = seq_len - 3 - min(len(dfg), data_flow_length)625 flat = flat[:keep][: code_length - 3]626 627 source_tokens = [tokenizer.cls_token] + flat + [tokenizer.sep_token]628 source_ids = tokenizer.convert_tokens_to_ids(source_tokens)629 # Code tokens get positions 2..; data-flow nodes get 0; padding gets630 # pad_token_id (1). This is what the model uses to tell them apart.631 position_idx = [i + pad_id + 1 for i in range(len(source_tokens))]632 633 dfg = dfg[: seq_len - len(source_tokens)]634 source_ids += [unk_id] * len(dfg)635 position_idx += [0] * len(dfg)636 padding = seq_len - len(source_ids)637 source_ids += [pad_id] * padding638 position_idx += [pad_id] * padding639 640 # Re-index edges so they point at node slots, not original token ids.641 reverse = {x[1]: i for i, x in enumerate(dfg)}642 dfg_to_dfg = [[reverse[i] for i in x[-1] if i in reverse] for x in dfg]643 dfg_to_code = [[ori2cur[x[1]][0] + 1, ori2cur[x[1]][1] + 1] for x in dfg]644 645 out_ids.append(source_ids)646 out_pos.append(position_idx)647 out_d2c.append(dfg_to_code)648 out_d2d.append(dfg_to_dfg)649 out_status.append(flags)650 651 return {652 "input_ids": out_ids,653 "position_idx": out_pos,654 "dfg_to_code": out_d2c,655 "dfg_to_dfg": out_d2d,656 "status_flags": out_status,657 }658 659 return fn660 661 662# --------------------------------------------------------------------------- #663# 5. Torch dataset + graph-guided attention-mask collator664# --------------------------------------------------------------------------- #665class ClonePairDataset(Dataset):666 """Pairs as ``(snippet_id1, snippet_id2, label)`` over a shared feature pool."""667 668 def __init__(self, index: SplitIndex, features: SnippetFeatures) -> None:669 self.index = index670 self.features = features671 672 def __len__(self) -> int:673 return len(self.index)674 675 def __getitem__(self, i: int) -> tuple[int, int, int]:676 return (677 int(self.index.snippet_id1[i]),678 int(self.index.snippet_id2[i]),679 int(self.index.labels[i]),680 )681 682 683def build_graph_attention_mask(features: SnippetFeatures, sid: int) -> np.ndarray:684 """Construct GraphCodeBERT's graph-guided masked attention for one snippet.685 686 Four rules, exactly as in the paper:687 688 1. code tokens attend to code tokens;689 2. the special tokens ``<s>``/``</s>`` attend to everything real;690 3. a data-flow node attends to (and is attended by) the code tokens it was691 identified from;692 4. a data-flow node attends to its adjacent nodes in the graph.693 """694 L = features.seq_length695 mask = np.zeros((L, L), dtype=bool)696 node_index = int(features.node_index[sid])697 max_length = int(features.max_length[sid])698 n_nodes = int(features.num_nodes[sid])699 700 # (1) sequence attends to sequence701 mask[:node_index, :node_index] = True702 # (2) special tokens attend to all real positions703 ids = features.input_ids[sid]704 for pos in np.flatnonzero((ids == 0) | (ids == 2)):705 if pos < node_index:706 mask[pos, :max_length] = True707 # (3) nodes <-> the code tokens they come from708 d2c = features.dfg_to_code[sid]709 for j in range(n_nodes):710 a, b = int(d2c[j, 0]), int(d2c[j, 1])711 if a < node_index and b < node_index:712 mask[j + node_index, a:b] = True713 mask[a:b, j + node_index] = True714 # (4) nodes <-> adjacent nodes715 offsets = features.dfg_adj_offsets[sid]716 for j in range(n_nodes):717 nbrs = features.dfg_adj_values[offsets[j] : offsets[j + 1]]718 for a in nbrs:719 if int(a) + node_index < L:720 mask[j + node_index, int(a) + node_index] = True721 return mask722 723 724@dataclass725class CloneCollator:726 """Collate pairs into the tensors ``GraphCodeBERTForCloneDetection`` expects."""727 728 features: SnippetFeatures729 730 def __call__(self, batch: list[tuple[int, int, int]]) -> dict[str, torch.Tensor]:731 ids1 = [b[0] for b in batch]732 ids2 = [b[1] for b in batch]733 labels = [b[2] for b in batch]734 f = self.features735 return {736 "input_ids_1": torch.from_numpy(f.input_ids[ids1].astype(np.int64)),737 "position_idx_1": torch.from_numpy(f.position_idx[ids1].astype(np.int64)),738 "attn_mask_1": torch.from_numpy(739 np.stack([build_graph_attention_mask(f, i) for i in ids1])740 ),741 "input_ids_2": torch.from_numpy(f.input_ids[ids2].astype(np.int64)),742 "position_idx_2": torch.from_numpy(f.position_idx[ids2].astype(np.int64)),743 "attn_mask_2": torch.from_numpy(744 np.stack([build_graph_attention_mask(f, i) for i in ids2])745 ),746 "labels": torch.tensor(labels, dtype=torch.long),747 }748 749 750# --------------------------------------------------------------------------- #751# 6. One-call pipeline used by train.py / evaluate.py752# --------------------------------------------------------------------------- #753@dataclass754class PreparedData:755 train: ClonePairDataset | None756 validation: ClonePairDataset757 test: ClonePairDataset758 features: SnippetFeatures759 collator: CloneCollator760 class_weights: list[float] | None761 report: dict[str, Any]762 763 764def prepare_data(cfg: Config, tokenizer: Any, with_train: bool = True) -> PreparedData:765 """Load, split, verify and featurise the dataset end to end."""766 schema = verify_dataset_schema(cfg.dataset_name)767 snippets, indices, pool_stats = build_snippet_pool(768 cfg, (cfg.train_split, cfg.heldout_split)769 )770 771 train_idx = indices[cfg.train_split]772 validation_idx, test_idx, split_report = split_heldout_by_group(773 indices[cfg.heldout_split], cfg.test_group_fraction, cfg.seed774 )775 776 train_idx, train_sub = subsample(777 train_idx, cfg.max_train_samples, cfg.seed, cfg.balance_subsamples778 )779 validation_idx, val_sub = subsample(780 validation_idx, cfg.max_eval_samples, cfg.seed, cfg.balance_subsamples781 )782 test_idx, test_sub = subsample(783 test_idx, cfg.max_test_samples, cfg.seed, cfg.balance_subsamples784 )785 786 _assert_no_leakage({"train": train_idx, "validation": validation_idx, "test": test_idx})787 788 class_weights, weight_report = decide_class_weights(789 train_idx, cfg.class_weighting, cfg.class_weight_threshold790 )791 792 features = build_snippet_features(cfg, snippets, tokenizer, cfg.preprocessing_num_workers)793 collator = CloneCollator(features)794 795 report = {796 "schema": schema,797 "snippet_pool": pool_stats,798 "heldout_split_strategy": split_report,799 "subsampling": {"train": train_sub, "validation": val_sub, "test": test_sub},800 "class_distribution": {801 "train": train_idx.class_distribution(),802 "validation": validation_idx.class_distribution(),803 "test": test_idx.class_distribution(),804 },805 "class_weighting": weight_report,806 "feature_extraction": features.stats,807 }808 809 return PreparedData(810 train=ClonePairDataset(train_idx, features) if with_train else None,811 validation=ClonePairDataset(validation_idx, features),812 test=ClonePairDataset(test_idx, features),813 features=features,814 collator=collator,815 class_weights=class_weights,816 report=report,817 )818 819 820def _assert_no_leakage(splits: dict[str, SplitIndex]) -> None:821 """Hard gate: no snippet and no group may appear in two splits."""822 names = list(splits)823 for i, a in enumerate(names):824 sa = set(splits[a].snippet_ids().tolist())825 ga = set(np.unique(np.concatenate([splits[a].group1, splits[a].group2])).tolist())826 for b in names[i + 1 :]:827 sb = set(splits[b].snippet_ids().tolist())828 gb = set(np.unique(np.concatenate([splits[b].group1, splits[b].group2])).tolist())829 if sa & sb:830 raise AssertionError(f"LEAK: {len(sa & sb)} snippets shared by {a} and {b}.")831 if ga & gb:832 raise AssertionError(f"LEAK: {len(ga & gb)} groups shared by {a} and {b}.")833 logger.info("Leakage check passed: splits share no snippet and no problem group.")834 835 836def iter_batches(dataset: Dataset, collator: CloneCollator, batch_size: int) -> Iterator[dict]:837 """Small helper for scripts that need batches without a Trainer."""838 for start in range(0, len(dataset), batch_size):839 yield collator([dataset[i] for i in range(start, min(start + batch_size, len(dataset)))])840 