Team Ai
Apppublic

CamQuest/codesearch

sourceHugging Facemitupdated 3mo agoView on Hugging Face
0likes
probe_code_encoders.py88 linesDownload Raw Back to scripts
1"""M5 sub-experiment B: probe fair code bi-encoders that load under current transformers.2 3For each candidate: try to load via SentenceTransformer (trust_remote_code where4needed), encode a short code snippet + a NL query, report embedding dim and a5sanity cosine. Failures are caught and reported per-model so one bad model does6not abort the sweep.7"""8import sys9import traceback10 11SNIPPET = "def add(a, b):\n    return a + b"12QUERY = "add two numbers"13 14# (model_id, needs_trust_remote_code, note)15CANDIDATES = [16    ("microsoft/unixcoder-base", False, "RoBERTa arch; may need mean-pooling wrapper"),17    ("microsoft/codebert-base", False, "RoBERTa arch; may need mean-pooling wrapper"),18    ("jinaai/jina-embeddings-v2-base-code", True, "known-blocked baseline (find_pruneable_heads_and_indices)"),19    ("nomic-ai/CodeRankEmbed", True, "newer code bi-encoder, standard-ish arch"),20    ("Alibaba-NLP/gte-modernbert-base", True, "modern general encoder, control"),21]22 23 24def probe_sentence_transformer(model_id, trust):25    import numpy as np26    from sentence_transformers import SentenceTransformer27 28    st = SentenceTransformer(model_id, trust_remote_code=trust)29    emb = st.encode([SNIPPET, QUERY], normalize_embeddings=True)30    dim = emb.shape[1]31    cos = float(np.dot(emb[0], emb[1]))32    return dim, cos33 34 35def probe_mean_pool(model_id):36    """Fallback: raw HF model + mean pooling (for encoders with no ST config)."""37    import numpy as np38    import torch39    from transformers import AutoModel, AutoTokenizer40 41    tok = AutoTokenizer.from_pretrained(model_id)42    model = AutoModel.from_pretrained(model_id)43    model.eval()44 45    def encode(text):46        batch = tok(text, return_tensors="pt", truncation=True, max_length=256)47        with torch.no_grad():48            out = model(**batch).last_hidden_state49        mask = batch["attention_mask"].unsqueeze(-1).float()50        vec = (out * mask).sum(1) / mask.sum(1).clamp(min=1e-9)51        vec = torch.nn.functional.normalize(vec, dim=-1)52        return vec[0].numpy()53 54    v_code = encode(SNIPPET)55    v_query = encode(QUERY)56    return v_code.shape[0], float(np.dot(v_code, v_query))57 58 59def main():60    results = []61    for model_id, trust, note in CANDIDATES:62        print(f"\n{'='*70}\n{model_id}  ({note})\n{'='*70}", flush=True)63        row = {"model": model_id, "note": note}64        try:65            dim, cos = probe_sentence_transformer(model_id, trust)66            row.update(status="OK (ST)", dim=dim, cos=round(cos, 4))67            print(f"  -> OK via SentenceTransformer: dim={dim} cos={cos:.4f}", flush=True)68        except Exception as e:69            print(f"  ST load failed: {type(e).__name__}: {e}", flush=True)70            print("  trying raw mean-pooling fallback...", flush=True)71            try:72                dim, cos = probe_mean_pool(model_id)73                row.update(status="OK (mean-pool)", dim=dim, cos=round(cos, 4))74                print(f"  -> OK via mean-pool: dim={dim} cos={cos:.4f}", flush=True)75            except Exception as e2:76                row.update(status=f"FAIL: {type(e2).__name__}", dim=None, cos=None)77                print(f"  -> FAIL: {type(e2).__name__}: {e2}", flush=True)78                traceback.print_exc()79        results.append(row)80 81    print(f"\n\n{'#'*70}\nSUMMARY\n{'#'*70}")82    for r in results:83        print(f"  {r['status']:>18}  dim={r['dim']}  cos={r['cos']}  {r['model']}")84 85 86if __name__ == "__main__":87    sys.exit(main())88