Download evaluate.py from zeechimp/zee: direct link, hf CLI and curl.
- Browser
- Download file 4.17 kB
-
https://huggingface.co/zeechimp/zee/resolve/main/evaluate.py
- Command line
-
hf download hf://zeechimp/zee/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/zeechimp/zee/resolve/main/evaluate.py
4.17 kB
| """ | |
| evaluate.py | |
| Reproduce the 5-fold cross-validation accuracy. | |
| Reconstructs the fold splits from phrases.txt, retrains on each | |
| 4/5 subset, and reports the held-out accuracy. Result should be | |
| around 76%. | |
| """ | |
| import numpy as np | |
| from hv_intent import load_v2, predict, _encode_phrase | |
| import json | |
| from pathlib import Path | |
| N_FOLDS = 5 | |
| D = 2048 | |
| N_AUG = 12 | |
| N_CODEBOOKS = 2 | |
| AUG_SEED = 42 | |
| DROP_THRESHOLD = 0.15 | |
| K_NEIGHBOURS = 7 | |
| def compute_word_weights(train_data, n_classes): | |
| counts = {} | |
| for phrase, t in train_data: | |
| for w in phrase.lower().split(): | |
| counts.setdefault(w, np.zeros(n_classes))[t] += 1 | |
| floor = 1.0 / n_classes | |
| return { | |
| w: float((c.max() / c.sum() - floor) / (1.0 - floor)) | |
| for w, c in counts.items() if c.sum() >= 1 | |
| } | |
| def augment_drop(phrase, rng): | |
| words = phrase.split() | |
| n = len(words) | |
| if n <= 2: | |
| return phrase | |
| op = int(rng.integers(0, 3)) | |
| if op == 0: | |
| i = int(rng.integers(0, n)) | |
| return " ".join(w for j, w in enumerate(words) if j != i) | |
| elif op == 1: | |
| if n <= 3: | |
| return phrase | |
| idx = rng.choice(n, size=2, replace=False) | |
| return " ".join(w for j, w in enumerate(words) if j not in idx) | |
| else: | |
| return " ".join(words[:-1] if rng.random() < 0.5 else words[1:]) | |
| def build_bank(train_data, n_aug, seed): | |
| rng = np.random.default_rng(seed) | |
| bank = list(train_data) | |
| for phrase, label in train_data: | |
| for _ in range(n_aug): | |
| bank.append((augment_drop(phrase, rng), label)) | |
| return bank | |
| def main(path="."): | |
| path = Path(path) | |
| with open(path / "config.json") as f: | |
| config = json.load(f) | |
| intent_labels = config["intent_labels"] | |
| l2i = {l: i for i, l in enumerate(intent_labels)} | |
| # Load phrases | |
| data = np.load(path / "codebooks.npz", allow_pickle=True) | |
| phrases = list(data["phrases"]) | |
| labels = data["phrase_labels"] | |
| # Group by intent | |
| by_intent = {l: [] for l in intent_labels} | |
| for phrase, t in zip(phrases[:50], labels[:50]): # first 50 are originals | |
| by_intent[intent_labels[t]].append(phrase) | |
| # Make folds | |
| folds = [[] for _ in range(N_FOLDS)] | |
| for label, phrases_list in by_intent.items(): | |
| for i, p in enumerate(phrases_list): | |
| folds[i % N_FOLDS].append((p, l2i[label])) | |
| correct = 0 | |
| total = 0 | |
| for fold_i in range(N_FOLDS): | |
| test_data = folds[fold_i] | |
| train_data = [item for j, f in enumerate(folds) if j != fold_i | |
| for item in f] | |
| ww = compute_word_weights(train_data, len(intent_labels)) | |
| bank = build_bank(train_data, N_AUG, AUG_SEED + fold_i) | |
| # Vocab | |
| vocab = sorted({w for p, _ in train_data for w in p.lower().split()}) | |
| vocab_index = {w: i for i, w in enumerate(vocab)} | |
| codebooks = [] | |
| for cb in range(N_CODEBOOKS): | |
| rng = np.random.default_rng(AUG_SEED * 100 + fold_i * 10 + cb) | |
| codebooks.append(rng.choice(np.array([-1, 1], dtype=np.int8), | |
| size=(len(vocab), D))) | |
| banks = [] | |
| for K in codebooks: | |
| H = np.stack([_encode_phrase(p, K, vocab_index, ww) | |
| for p, _ in bank]) | |
| Y = np.array([y for _, y in bank], dtype=np.int32) | |
| banks.append((H, Y)) | |
| for phrase, t in test_data: | |
| vote = np.zeros(len(intent_labels), dtype=np.float32) | |
| for K, (H, Y) in zip(codebooks, banks): | |
| q = _encode_phrase(phrase, K, vocab_index, ww) | |
| s = H @ q | |
| k_eff = min(K_NEIGHBOURS, s.shape[0]) | |
| top = np.argpartition(-s, k_eff - 1)[:k_eff] | |
| for j in top: | |
| vote[Y[j]] += max(s[j], 0.0) | |
| if int(vote.argmax()) == t: | |
| correct += 1 | |
| total += 1 | |
| print(f"5-fold CV accuracy: {correct}/{total} = {correct/total:.1%}") | |
| print(f"Config: D={D}, aug={N_AUG}, codebooks={N_CODEBOOKS}, " | |
| f"k={K_NEIGHBOURS}") | |
| if __name__ == "__main__": | |
| import sys | |
| main(sys.argv[1] if len(sys.argv) > 1 else ".") |