zee / evaluate.py
zeechimp's picture
Update evaluate.py
6f1ff07 verified
Raw History Blame Contribute Delete
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 ".")