CamQuest/codesearch
0
1"""Validate UniXcoderEncoder against the official microsoft/CodeBERT recipe.2 3Tests, in order of strength:4 1. GOLDEN: element-wise match vs the official unixcoder.py (fetched from GitHub)5 on single unpadded sequences — where our 1-D attention mask and upstream's6 2-D mask are provably identical, so raw pooled vectors must match to ~1e-5.7 2. ID framing: [<s>=0, <encoder-only>=6, </s>=2] + bpe + [</s>=2].8 3. Padding-invariance: batched (padded) == each item alone, to ~1e-5 (mask test).9 4. Unit-norm + dim==768 after normalize.10 5. Toy retrieval: 5 code snippets vs their NL descriptions -> MRR == 1.0, and the11 correct-pair cosine is discriminative (vs the degenerate ~0.20 of ST's default).12 13Exit non-zero on any failure.14"""15from __future__ import annotations16 17import sys18import types19import urllib.request20 21import numpy as np22import torch23 24sys.path.insert(0, "src")25from codesearch.embedding import UniXcoderEncoder # noqa: E40226 27MODEL = "microsoft/unixcoder-base"28_OFFICIAL_URL = (29 "https://raw.githubusercontent.com/microsoft/CodeBERT/master/UniXcoder/unixcoder.py"30)31 32CODE = [33 "def add(a, b):\n return a + b",34 "def is_even(n):\n return n % 2 == 0",35 "def reverse(s):\n return s[::-1]",36 "def to_upper(s):\n return s.upper()",37 "def read_file(path):\n with open(path) as f:\n return f.read()",38]39NL = [40 "add two numbers",41 "check whether a number is even",42 "reverse a string",43 "convert a string to uppercase",44 "read the contents of a file",45]46 47_FAILS: list[str] = []48 49 50def check(cond: bool, msg: str) -> None:51 print((" PASS " if cond else " FAIL ") + msg)52 if not cond:53 _FAILS.append(msg)54 55 56def load_official():57 """Fetch official unixcoder.py, import it as a module, return the UniXcoder class."""58 src = urllib.request.urlopen(_OFFICIAL_URL, timeout=30).read().decode()59 mod = types.ModuleType("official_unixcoder")60 exec(compile(src, "official_unixcoder.py", "exec"), mod.__dict__)61 return mod.UniXcoder62 63 64def main() -> int:65 enc = UniXcoderEncoder(MODEL)66 67 # --- 2. ID framing -------------------------------------------------------68 ids = enc._build_ids(CODE[0])69 check(ids[:3] == [0, 6, 2] and ids[-1] == 2,70 f"ID framing [<s>,<enc-only>,</s>]...[</s>] got {ids[:3]}...{ids[-1]}")71 72 # --- 1. GOLDEN cross-check vs the official recipe ------------------------73 # The as-shipped official forward() cannot run under transformers 5.x: its 2-D74 # (mask_i*mask_j) attention mask crashes on batched input and goes silently75 # causal on single input. So we validate against the official recipe run the76 # way tf5 requires: (a) its tokenizer framing must match ours id-for-id, and77 # (b) the official model config (is_decoder=True) fed a full-bidirectional 4-D78 # mask + mean-pool must match our vectors element-wise. is_decoder is a79 # mask-only flag (zero weight effect), so this is a true faithfulness check.80 try:81 from transformers import AutoConfig, AutoModel82 83 Official = load_official()84 off = Official(MODEL) # only used for its .tokenize()85 86 # (a) framing: official ids == ours87 max_id_mismatch = 088 for text in CODE + NL:89 off_ids = off.tokenize([text], max_length=512, mode="<encoder-only>")[0]90 max_id_mismatch = max(max_id_mismatch, int(off_ids != enc._build_ids(text)))91 check(max_id_mismatch == 0, "golden framing: official.tokenize ids == _build_ids")92 93 # (b) element-wise: official config (is_decoder=True) + full bidirectional94 cfg = AutoConfig.from_pretrained(MODEL)95 cfg.is_decoder = True96 ref_model = AutoModel.from_pretrained(MODEL, config=cfg)97 ref_model.eval()98 max_ediff = 0.099 for text in CODE + NL:100 ids = torch.tensor([enc._build_ids(text)]) # single, unpadded101 full4d = torch.zeros(1, 1, ids.shape[1], ids.shape[1]) # all-zeros => full bidir102 with torch.no_grad():103 ref = ref_model(ids, attention_mask=full4d)[0][0].mean(0)104 ref_vec = ref.numpy().astype(np.float32)105 mine_raw = enc.encode(text, normalize_embeddings=False)106 max_ediff = max(max_ediff, float(np.abs(mine_raw - ref_vec).max()))107 check(max_ediff < 1e-4,108 f"golden element-wise match vs official bidirectional recipe: max|Δ|={max_ediff:.2e} (<1e-4)")109 except Exception as e: # network / upstream unavailable110 print(f" SKIP golden cross-check unavailable: {type(e).__name__}: {e}")111 112 # --- 3. Padding-invariance (batched vs single) ---------------------------113 batched = enc.encode(CODE, normalize_embeddings=True)114 singles = np.vstack([enc.encode(c, normalize_embeddings=True) for c in CODE])115 pad_diff = float(np.abs(batched - singles).max())116 check(pad_diff < 1e-5, f"padding-invariance: max|Δ|={pad_diff:.2e} (<1e-5)")117 118 # --- 4. Unit norm + dim --------------------------------------------------119 norms = np.linalg.norm(batched, axis=1)120 check(batched.shape[1] == 768, f"dim==768 got {batched.shape[1]}")121 check(np.allclose(norms, 1.0, atol=1e-4), f"unit-norm rows norms in [{norms.min():.4f},{norms.max():.4f}]")122 123 # --- 5. Toy retrieval sanity (MRR) --------------------------------------124 code_emb = enc.encode(CODE, normalize_embeddings=True)125 nl_emb = enc.encode(NL, normalize_embeddings=True)126 sims = nl_emb @ code_emb.T # (query, doc) cosine127 ranks = []128 for i in range(len(NL)):129 order = np.argsort(-sims[i])130 ranks.append(int(np.where(order == i)[0][0]) + 1)131 mrr = float(np.mean([1.0 / r for r in ranks]))132 diag = float(np.mean(np.diag(sims)))133 check(mrr == 1.0, f"toy retrieval MRR==1.0 ranks={ranks} mrr={mrr:.3f}")134 check(diag > 0.30, f"correct-pair cosine discriminative: mean diag={diag:.3f} (>>0.20 degenerate)")135 136 print("\n" + ("ALL PASS" if not _FAILS else f"{len(_FAILS)} FAILURE(S): " + "; ".join(_FAILS)))137 return 1 if _FAILS else 0138 139 140if __name__ == "__main__":141 sys.exit(main())142 