File size: 10,900 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
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
"""Central configuration for GraphCodeBERT binary code-clone detection.

Every tunable lives in :class:`Config`.  ``Config.from_cli`` turns the dataclass
fields into ``argparse`` flags automatically, so ``train.py`` / ``evaluate.py``
never drift apart from this file.
"""

from __future__ import annotations

import argparse
import dataclasses
import json
from dataclasses import dataclass, field, fields
from pathlib import Path
from typing import Any

# --------------------------------------------------------------------------- #
# Facts discovered by inspecting the dataset (see README "Dataset structure").
# They are kept here as documented defaults, *not* as blind assumptions:
# preprocess.py re-verifies every one of them at runtime and raises if the
# remote dataset ever changes.
# --------------------------------------------------------------------------- #
CODE1_COLUMN = "code1"
CODE2_COLUMN = "code2"
LABEL_COLUMN = "similar"
GROUP1_COLUMN = "code1_group"
GROUP2_COLUMN = "code2_group"
#: Metadata columns that MUST NOT reach the model. ``code1_group``/``code2_group``
#: determine the label exactly (``similar == (code1_group == code2_group)``), so
#: feeding them in any form would be a 100 % label leak.
FORBIDDEN_FEATURE_COLUMNS = (
    GROUP1_COLUMN,
    GROUP2_COLUMN,
    "pair_id",
    "question_pair_id",
)
#: The source language of the snippets, needed to pick the tree-sitter grammar.
DATASET_LANGUAGE = "python"


@dataclass
class Config:
    """All knobs for preprocessing, training and evaluation."""

    # ---------------- dataset ---------------- #
    dataset_name: str = "PoolC/1-fold-clone-detection-600k-5fold"
    #: HF split that becomes the training set (group-disjoint from ``val``).
    train_split: str = "train"
    #: HF split that is partitioned by *group* into validation and test.
    heldout_split: str = "val"
    #: Fraction of the held-out split's groups reserved for the test set.
    test_group_fraction: float = 0.5

    #: Cap on the number of pairs per split. ``-1`` = use everything.
    #: The full training split has 5.39 M pairs; see README for why the default
    #: is a subsample and how to raise it.
    max_train_samples: int = 50_000
    max_eval_samples: int = 20_000
    max_test_samples: int = 20_000
    #: Keep the 50/50 label balance exactly when subsampling.
    balance_subsamples: bool = True

    # ---------------- model ---------------- #
    model_name_or_path: str = "microsoft/graphcodebert-base"
    #: GraphCodeBERT clone-detection defaults from the original paper/repo.
    code_length: int = 512
    data_flow_length: int = 128
    attn_implementation: str = "sdpa"

    # ---------------- training ---------------- #
    learning_rate: float = 2e-5
    num_train_epochs: float = 3.0
    per_device_train_batch_size: int = 4
    per_device_eval_batch_size: int = 4
    gradient_accumulation_steps: int = 4
    weight_decay: float = 0.01
    warmup_ratio: float = 0.1
    max_grad_norm: float = 1.0
    fp16: bool = True
    bf16: bool = False
    gradient_checkpointing: bool = False
    optim: str = "adamw_torch"
    lr_scheduler_type: str = "linear"

    #: Class-weighted cross entropy. ``"auto"`` enables it only when the
    #: measured training distribution is more skewed than
    #: ``class_weight_threshold``; ``"off"`` never, ``"on"`` always.
    class_weighting: str = "auto"
    class_weight_threshold: float = 0.6

    # ---------------- evaluation / checkpointing ---------------- #
    eval_strategy: str = "steps"
    eval_steps: int = 1000
    save_strategy: str = "steps"
    save_steps: int = 1000
    save_total_limit: int = 2
    logging_steps: int = 100
    metric_for_best_model: str = "f1"
    greater_is_better: bool = True
    load_best_model_at_end: bool = True

    # ---------------- runtime ---------------- #
    seed: int = 42
    full_determinism: bool = False
    dataloader_num_workers: int = 4
    #: Processes used for tree-sitter data-flow extraction.
    preprocessing_num_workers: int = 8
    output_dir: str = "./outputs"
    model_dir: str = "./models/graphcodebert-clone-detection"
    logging_dir: str = "./logs"
    cache_dir: str = "./outputs/feature_cache"
    report_to: str = "none"
    run_sanity_check: bool = True
    sanity_check_samples: int = 64

    # ---------------- derived ---------------- #
    @property
    def total_sequence_length(self) -> int:
        """Length of one encoded snippet: code tokens + data-flow nodes."""
        return self.code_length + self.data_flow_length

    @property
    def effective_batch_size(self) -> int:
        return self.per_device_train_batch_size * self.gradient_accumulation_steps

    # ---------------- (de)serialisation ---------------- #
    def to_dict(self) -> dict[str, Any]:
        d = dataclasses.asdict(self)
        d["total_sequence_length"] = self.total_sequence_length
        d["effective_batch_size"] = self.effective_batch_size
        return d

    def save(self, path: str | Path) -> None:
        path = Path(path)
        path.parent.mkdir(parents=True, exist_ok=True)
        path.write_text(json.dumps(self.to_dict(), indent=2), encoding="utf-8")

    @classmethod
    def from_json(cls, path: str | Path) -> "Config":
        raw = json.loads(Path(path).read_text(encoding="utf-8"))
        known = {f.name for f in fields(cls)}
        return cls(**{k: v for k, v in raw.items() if k in known})

    # ---------------- CLI ---------------- #
    @classmethod
    def build_parser(cls, description: str = "") -> argparse.ArgumentParser:
        parser = argparse.ArgumentParser(
            description=description,
            formatter_class=argparse.ArgumentDefaultsHelpFormatter,
        )
        parser.add_argument(
            "--config_json",
            type=str,
            default=None,
            help="Load defaults from a saved training_config.json, then apply CLI overrides.",
        )
        for f in fields(cls):
            flag = f"--{f.name}"
            if f.type is bool or f.type == "bool":
                # Accept --flag / --flag true / --flag false.
                parser.add_argument(
                    flag,
                    type=_str2bool,
                    nargs="?",
                    const=True,
                    default=None,
                    help=f"(bool) default: {f.default}",
                )
            else:
                parser.add_argument(flag, type=type(f.default), default=None)
        return parser

    @classmethod
    def from_cli(cls, argv: list[str] | None = None, description: str = "") -> "Config":
        parser = cls.build_parser(description)
        args, unknown = parser.parse_known_args(argv)
        if unknown:
            raise SystemExit(f"Unrecognised arguments: {unknown}")
        cfg = cls.from_json(args.config_json) if args.config_json else cls()
        for f in fields(cls):
            value = getattr(args, f.name, None)
            if value is not None:
                setattr(cfg, f.name, value)
        cfg.validate()
        return cfg

    def validate(self) -> None:
        """Fail fast on impossible combinations instead of dying mid-training."""
        if self.fp16 and self.bf16:
            raise ValueError("Enable at most one of fp16 / bf16.")
        if not 0.0 < self.test_group_fraction < 1.0:
            raise ValueError("test_group_fraction must lie strictly between 0 and 1.")
        if self.code_length <= 3:
            raise ValueError("code_length must leave room for <s> and </s>.")
        if self.data_flow_length < 0:
            raise ValueError("data_flow_length must be >= 0.")
        # GraphCodeBERT position ids run up to code_length + 1; RoBERTa's
        # embedding table holds 514 slots (512 + <pad> + offset).
        if self.code_length > 512:
            raise ValueError(
                "code_length > 512 exceeds GraphCodeBERT's position embeddings (514 slots)."
            )
        if self.class_weighting not in {"auto", "on", "off"}:
            raise ValueError("class_weighting must be one of: auto, on, off.")
        if self.metric_for_best_model not in {
            "f1",
            "accuracy",
            "precision",
            "recall",
            "loss",
        }:
            raise ValueError(f"Unsupported metric_for_best_model: {self.metric_for_best_model}")
        if self.load_best_model_at_end and self.eval_strategy != self.save_strategy:
            raise ValueError("load_best_model_at_end requires eval_strategy == save_strategy.")
        if (
            self.load_best_model_at_end
            and self.eval_strategy == "steps"
            and self.save_steps % self.eval_steps != 0
        ):
            raise ValueError("save_steps must be a multiple of eval_steps.")


#: Placeholder namespace in the Makefile default; replaced by the logged-in user.
PLACEHOLDER_HUB_NAMESPACE = "your-username"
DEFAULT_HUB_MODEL_NAME = "graphcodebert-clone-detection"


def resolve_hub_repo_id(repo_id: str | None, token: str | None = None) -> str:
    """Expand a bare model name into ``<namespace>/<name>`` for the Hub.

    Accepts ``None``, a bare name, or a full ``user/name``. The namespace is
    taken from the caller's Hugging Face credential (``huggingface-cli login``,
    ``HF_TOKEN``, or an explicit ``token``), so the token never has to be typed
    on the command line.

    Raises:
        RuntimeError: if no namespace is given and no credential is available.
    """
    from huggingface_hub import HfApi, get_token

    name = (repo_id or DEFAULT_HUB_MODEL_NAME).strip().strip("/")
    if "/" in name:
        namespace, _, model_name = name.partition("/")
        if namespace != PLACEHOLDER_HUB_NAMESPACE:
            return f"{namespace}/{model_name}"
        name = model_name or DEFAULT_HUB_MODEL_NAME

    effective = token or get_token()
    if not effective:
        raise RuntimeError(
            "No Hugging Face credential found, so the namespace for "
            f"{name!r} cannot be resolved. Run `huggingface-cli login`, export "
            "HF_TOKEN, or pass the full repo id as HF_REPO=user/name."
        )
    try:
        who = HfApi().whoami(token=effective)
    except Exception as exc:
        raise RuntimeError(
            f"Could not identify the logged-in Hugging Face account: {exc}. "
            "Pass the full repo id as HF_REPO=user/name."
        ) from exc
    namespace = who.get("name")
    if not namespace:
        raise RuntimeError("Hugging Face account has no username; pass HF_REPO=user/name.")
    return f"{namespace}/{name}"


def _str2bool(value: str | bool) -> bool:
    if isinstance(value, bool):
        return value
    if value.lower() in {"true", "t", "yes", "y", "1"}:
        return True
    if value.lower() in {"false", "f", "no", "n", "0"}:
        return False
    raise argparse.ArgumentTypeError(f"Expected a boolean, got {value!r}")