thealper2/graphcodebert-code-clone-detection
030
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 