Team Ai
Apppublic

CamQuest/codesearch

sourceHugging Facemitupdated 3mo agoView on Hugging Face
0likes
test_unixcoder_encoder.py142 linesDownload Raw Back to scripts
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