Team Ai
Modelpublic

thealper2/graphcodebert-code-clone-detection

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