Team Ai
Modelpublic

thealper2/graphcodebert-code-clone-detection

sourceHugging Facemitupdated 17d agoView on Hugging Face
0likes30downloads
config.py276 linesDownload Raw Back to code
1"""Central configuration for GraphCodeBERT binary code-clone detection.2 3Every tunable lives in :class:`Config`.  ``Config.from_cli`` turns the dataclass4fields into ``argparse`` flags automatically, so ``train.py`` / ``evaluate.py``5never drift apart from this file.6"""7 8from __future__ import annotations9 10import argparse11import dataclasses12import json13from dataclasses import dataclass, field, fields14from pathlib import Path15from typing import Any16 17# --------------------------------------------------------------------------- #18# Facts discovered by inspecting the dataset (see README "Dataset structure").19# They are kept here as documented defaults, *not* as blind assumptions:20# preprocess.py re-verifies every one of them at runtime and raises if the21# remote dataset ever changes.22# --------------------------------------------------------------------------- #23CODE1_COLUMN = "code1"24CODE2_COLUMN = "code2"25LABEL_COLUMN = "similar"26GROUP1_COLUMN = "code1_group"27GROUP2_COLUMN = "code2_group"28#: Metadata columns that MUST NOT reach the model. ``code1_group``/``code2_group``29#: determine the label exactly (``similar == (code1_group == code2_group)``), so30#: feeding them in any form would be a 100 % label leak.31FORBIDDEN_FEATURE_COLUMNS = (32    GROUP1_COLUMN,33    GROUP2_COLUMN,34    "pair_id",35    "question_pair_id",36)37#: The source language of the snippets, needed to pick the tree-sitter grammar.38DATASET_LANGUAGE = "python"39 40 41@dataclass42class Config:43    """All knobs for preprocessing, training and evaluation."""44 45    # ---------------- dataset ---------------- #46    dataset_name: str = "PoolC/1-fold-clone-detection-600k-5fold"47    #: HF split that becomes the training set (group-disjoint from ``val``).48    train_split: str = "train"49    #: HF split that is partitioned by *group* into validation and test.50    heldout_split: str = "val"51    #: Fraction of the held-out split's groups reserved for the test set.52    test_group_fraction: float = 0.553 54    #: Cap on the number of pairs per split. ``-1`` = use everything.55    #: The full training split has 5.39 M pairs; see README for why the default56    #: is a subsample and how to raise it.57    max_train_samples: int = 50_00058    max_eval_samples: int = 20_00059    max_test_samples: int = 20_00060    #: Keep the 50/50 label balance exactly when subsampling.61    balance_subsamples: bool = True62 63    # ---------------- model ---------------- #64    model_name_or_path: str = "microsoft/graphcodebert-base"65    #: GraphCodeBERT clone-detection defaults from the original paper/repo.66    code_length: int = 51267    data_flow_length: int = 12868    attn_implementation: str = "sdpa"69 70    # ---------------- training ---------------- #71    learning_rate: float = 2e-572    num_train_epochs: float = 3.073    per_device_train_batch_size: int = 474    per_device_eval_batch_size: int = 475    gradient_accumulation_steps: int = 476    weight_decay: float = 0.0177    warmup_ratio: float = 0.178    max_grad_norm: float = 1.079    fp16: bool = True80    bf16: bool = False81    gradient_checkpointing: bool = False82    optim: str = "adamw_torch"83    lr_scheduler_type: str = "linear"84 85    #: Class-weighted cross entropy. ``"auto"`` enables it only when the86    #: measured training distribution is more skewed than87    #: ``class_weight_threshold``; ``"off"`` never, ``"on"`` always.88    class_weighting: str = "auto"89    class_weight_threshold: float = 0.690 91    # ---------------- evaluation / checkpointing ---------------- #92    eval_strategy: str = "steps"93    eval_steps: int = 100094    save_strategy: str = "steps"95    save_steps: int = 100096    save_total_limit: int = 297    logging_steps: int = 10098    metric_for_best_model: str = "f1"99    greater_is_better: bool = True100    load_best_model_at_end: bool = True101 102    # ---------------- runtime ---------------- #103    seed: int = 42104    full_determinism: bool = False105    dataloader_num_workers: int = 4106    #: Processes used for tree-sitter data-flow extraction.107    preprocessing_num_workers: int = 8108    output_dir: str = "./outputs"109    model_dir: str = "./models/graphcodebert-clone-detection"110    logging_dir: str = "./logs"111    cache_dir: str = "./outputs/feature_cache"112    report_to: str = "none"113    run_sanity_check: bool = True114    sanity_check_samples: int = 64115 116    # ---------------- derived ---------------- #117    @property118    def total_sequence_length(self) -> int:119        """Length of one encoded snippet: code tokens + data-flow nodes."""120        return self.code_length + self.data_flow_length121 122    @property123    def effective_batch_size(self) -> int:124        return self.per_device_train_batch_size * self.gradient_accumulation_steps125 126    # ---------------- (de)serialisation ---------------- #127    def to_dict(self) -> dict[str, Any]:128        d = dataclasses.asdict(self)129        d["total_sequence_length"] = self.total_sequence_length130        d["effective_batch_size"] = self.effective_batch_size131        return d132 133    def save(self, path: str | Path) -> None:134        path = Path(path)135        path.parent.mkdir(parents=True, exist_ok=True)136        path.write_text(json.dumps(self.to_dict(), indent=2), encoding="utf-8")137 138    @classmethod139    def from_json(cls, path: str | Path) -> "Config":140        raw = json.loads(Path(path).read_text(encoding="utf-8"))141        known = {f.name for f in fields(cls)}142        return cls(**{k: v for k, v in raw.items() if k in known})143 144    # ---------------- CLI ---------------- #145    @classmethod146    def build_parser(cls, description: str = "") -> argparse.ArgumentParser:147        parser = argparse.ArgumentParser(148            description=description,149            formatter_class=argparse.ArgumentDefaultsHelpFormatter,150        )151        parser.add_argument(152            "--config_json",153            type=str,154            default=None,155            help="Load defaults from a saved training_config.json, then apply CLI overrides.",156        )157        for f in fields(cls):158            flag = f"--{f.name}"159            if f.type is bool or f.type == "bool":160                # Accept --flag / --flag true / --flag false.161                parser.add_argument(162                    flag,163                    type=_str2bool,164                    nargs="?",165                    const=True,166                    default=None,167                    help=f"(bool) default: {f.default}",168                )169            else:170                parser.add_argument(flag, type=type(f.default), default=None)171        return parser172 173    @classmethod174    def from_cli(cls, argv: list[str] | None = None, description: str = "") -> "Config":175        parser = cls.build_parser(description)176        args, unknown = parser.parse_known_args(argv)177        if unknown:178            raise SystemExit(f"Unrecognised arguments: {unknown}")179        cfg = cls.from_json(args.config_json) if args.config_json else cls()180        for f in fields(cls):181            value = getattr(args, f.name, None)182            if value is not None:183                setattr(cfg, f.name, value)184        cfg.validate()185        return cfg186 187    def validate(self) -> None:188        """Fail fast on impossible combinations instead of dying mid-training."""189        if self.fp16 and self.bf16:190            raise ValueError("Enable at most one of fp16 / bf16.")191        if not 0.0 < self.test_group_fraction < 1.0:192            raise ValueError("test_group_fraction must lie strictly between 0 and 1.")193        if self.code_length <= 3:194            raise ValueError("code_length must leave room for <s> and </s>.")195        if self.data_flow_length < 0:196            raise ValueError("data_flow_length must be >= 0.")197        # GraphCodeBERT position ids run up to code_length + 1; RoBERTa's198        # embedding table holds 514 slots (512 + <pad> + offset).199        if self.code_length > 512:200            raise ValueError(201                "code_length > 512 exceeds GraphCodeBERT's position embeddings (514 slots)."202            )203        if self.class_weighting not in {"auto", "on", "off"}:204            raise ValueError("class_weighting must be one of: auto, on, off.")205        if self.metric_for_best_model not in {206            "f1",207            "accuracy",208            "precision",209            "recall",210            "loss",211        }:212            raise ValueError(f"Unsupported metric_for_best_model: {self.metric_for_best_model}")213        if self.load_best_model_at_end and self.eval_strategy != self.save_strategy:214            raise ValueError("load_best_model_at_end requires eval_strategy == save_strategy.")215        if (216            self.load_best_model_at_end217            and self.eval_strategy == "steps"218            and self.save_steps % self.eval_steps != 0219        ):220            raise ValueError("save_steps must be a multiple of eval_steps.")221 222 223#: Placeholder namespace in the Makefile default; replaced by the logged-in user.224PLACEHOLDER_HUB_NAMESPACE = "your-username"225DEFAULT_HUB_MODEL_NAME = "graphcodebert-clone-detection"226 227 228def resolve_hub_repo_id(repo_id: str | None, token: str | None = None) -> str:229    """Expand a bare model name into ``<namespace>/<name>`` for the Hub.230 231    Accepts ``None``, a bare name, or a full ``user/name``. The namespace is232    taken from the caller's Hugging Face credential (``huggingface-cli login``,233    ``HF_TOKEN``, or an explicit ``token``), so the token never has to be typed234    on the command line.235 236    Raises:237        RuntimeError: if no namespace is given and no credential is available.238    """239    from huggingface_hub import HfApi, get_token240 241    name = (repo_id or DEFAULT_HUB_MODEL_NAME).strip().strip("/")242    if "/" in name:243        namespace, _, model_name = name.partition("/")244        if namespace != PLACEHOLDER_HUB_NAMESPACE:245            return f"{namespace}/{model_name}"246        name = model_name or DEFAULT_HUB_MODEL_NAME247 248    effective = token or get_token()249    if not effective:250        raise RuntimeError(251            "No Hugging Face credential found, so the namespace for "252            f"{name!r} cannot be resolved. Run `huggingface-cli login`, export "253            "HF_TOKEN, or pass the full repo id as HF_REPO=user/name."254        )255    try:256        who = HfApi().whoami(token=effective)257    except Exception as exc:258        raise RuntimeError(259            f"Could not identify the logged-in Hugging Face account: {exc}. "260            "Pass the full repo id as HF_REPO=user/name."261        ) from exc262    namespace = who.get("name")263    if not namespace:264        raise RuntimeError("Hugging Face account has no username; pass HF_REPO=user/name.")265    return f"{namespace}/{name}"266 267 268def _str2bool(value: str | bool) -> bool:269    if isinstance(value, bool):270        return value271    if value.lower() in {"true", "t", "yes", "y", "1"}:272        return True273    if value.lower() in {"false", "f", "no", "n", "0"}:274        return False275    raise argparse.ArgumentTypeError(f"Expected a boolean, got {value!r}")276