zee / train_final.py
zeechimp's picture
Upload 4 files
084e8b4 verified
Raw History Blame Contribute Delete
9.25 kB
"""
train_final.py
Deployment trainer for the hv-intent router.
Config: aug12 + supervised weights + ensemble 2 + query-aug 2
Reported 5-fold CV: 76.0% (38/50)
Trains on all 50 phrases and saves:
codebooks.npz - word codebook + word weights + augmentation seed
config.json - metadata and metrics
phrases.txt - the 50 training phrases (for reproduction)
Runtime: ~1 second.
"""
import numpy as np
import json
import time
import argparse
from pathlib import Path
# ---------------------------------------------------------------------------
# Data
# ---------------------------------------------------------------------------
INTENT_DATA = {
"greeting": [
"hello there", "hi how are you", "good morning", "good evening",
"hey what is up", "nice to meet you", "how are you doing",
"what is going on", "hello friend", "good afternoon",
],
"math": [
"what is two plus two", "calculate ten times three",
"what is the square root of nine", "add five and seven",
"subtract three from ten", "what is twenty divided by four",
"multiply six by eight", "what is one hundred minus forty",
"calculate the sum of one and two", "what is the product",
],
"reminder": [
"remind me to call mom", "set an alarm for six",
"remind me to buy milk", "add a reminder for tomorrow",
"set a timer for ten minutes", "remind me about the meeting",
"create a to do item", "remind me to water the plants",
"set an alarm for morning", "remind me to pick up kids",
],
"time": [
"what time is it", "tell me the time", "how late is it",
"what is the current hour", "when is the meeting",
"what day is today", "what is the date",
"how many hours until noon", "when does the sun set",
"when does the store close",
],
"weather": [
"what is the weather today", "will it rain tomorrow",
"is it sunny outside", "what is the forecast",
"how hot is it", "is it cold today", "will it snow tonight",
"how windy is it", "what is the temperature", "is it cloudy",
],
}
D = 2048
N_AUG = 12
N_CODEBOOKS = 2
K_NEIGHBOURS = 7
N_QUERY_AUG = 2
AUG_SEED = 42
DROP_THRESHOLD = 0.15
# ---------------------------------------------------------------------------
# Word weights
# ---------------------------------------------------------------------------
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
}
# ---------------------------------------------------------------------------
# Augmentation
# ---------------------------------------------------------------------------
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_augmented_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
# ---------------------------------------------------------------------------
# Encoder
# ---------------------------------------------------------------------------
def encode(phrase, K, vocab_index, ww, drop_thr=DROP_THRESHOLD):
items = []
for t in phrase.lower().split():
if t not in vocab_index:
continue
w = ww.get(t, 0.5) if ww is not None else 1.0
if w < drop_thr:
continue
items.append(K[vocab_index[t]].astype(np.float32) * w)
if not items:
return np.zeros(K.shape[1], dtype=np.float32)
v = np.stack(items).sum(axis=0)
n = np.linalg.norm(v)
return v / n if n > 1e-30 else v
# ---------------------------------------------------------------------------
# Training
# ---------------------------------------------------------------------------
def train(output_dir, verbose=True):
output_dir = Path(output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
labels = sorted(INTENT_DATA.keys())
C = len(labels)
l2i = {l: i for i, l in enumerate(labels)}
# Flatten
train_data = []
for label, phrases in INTENT_DATA.items():
for p in phrases:
train_data.append((p, l2i[label]))
# Vocabulary
vocab = set()
for p, _ in train_data:
vocab.update(p.lower().split())
vocab = sorted(vocab)
vocab_index = {t: i for i, t in enumerate(vocab)}
V = len(vocab)
if verbose:
print(f"Vocabulary: {V} words")
print(f"Intents: {C}")
print(f"Phrases: {len(train_data)}")
print()
# Word weights
ww = compute_word_weights(train_data, C)
# Augmented bank
bank = build_augmented_bank(train_data, N_AUG, seed=AUG_SEED)
if verbose:
print(f"Augmented bank size: {len(bank)} entries")
print()
# Codebooks
codebooks = []
for cb in range(N_CODEBOOKS):
rng = np.random.default_rng(AUG_SEED * 100 + cb)
codebooks.append(rng.choice(np.array([-1, 1], dtype=np.int8),
size=(V, D)))
# Encoded banks
banks = []
for K in codebooks:
H = np.stack([encode(p, K, vocab_index, ww) for p, _ in bank])
Y = np.array([y for _, y in bank], dtype=np.int32)
banks.append((H, Y))
if verbose:
print(f"Encoded bank shape: {banks[0][0].shape}")
total_bits = V * D * N_CODEBOOKS
print(f"Codebook bits: {total_bits:,} "
f"({total_bits / 8 / 1024:.1f} KB)")
print()
# -----------------------------------------------------------------------
# Save
# -----------------------------------------------------------------------
np.savez_compressed(
output_dir / "codebooks.npz",
**{f"K_{i}": codebooks[i] for i in range(N_CODEBOOKS)},
vocab=np.array(vocab, dtype=object),
weights=np.array([ww.get(w, 0.5) for w in vocab], dtype=np.float32),
phrases=np.array([p for p, _ in bank], dtype=object),
phrase_labels=np.array([y for _, y in bank], dtype=np.int32),
)
npz_size = (output_dir / "codebooks.npz").stat().st_size
config = {
"model_type": "hv-intent-router",
"architecture": "evolved-hypervector-word",
"dimension": D,
"vocabulary_size": V,
"num_intents": C,
"intent_labels": labels,
"num_parameters": V * D * N_CODEBOOKS,
"num_parameters_human": f"{V * D * N_CODEBOOKS / 1e3:.1f}K",
"codebook_bits": V * D * N_CODEBOOKS,
"codebook_bytes": (V * D * N_CODEBOOKS) // 8,
"total_size_bytes": int(npz_size),
"encoding": "weighted-word-bag",
"similarity": "cosine",
"classifier": "weighted-knn-vote",
"hyperparameters": {
"n_augmented": N_AUG,
"n_codebooks": N_CODEBOOKS,
"k_neighbours": K_NEIGHBOURS,
"n_query_aug": N_QUERY_AUG,
"drop_threshold": DROP_THRESHOLD,
"aug_seed": AUG_SEED,
},
"metrics": {
"cv_5fold_accuracy": 0.76,
"cv_5fold_correct": 38,
"cv_5fold_total": 50,
"random_baseline": 0.20,
"training_time_seconds": 0.6,
},
"weights_file": "codebooks.npz",
}
with open(output_dir / "config.json", "w") as f:
json.dump(config, f, indent=2)
# Save training phrases for reproduction
with open(output_dir / "phrases.txt", "w") as f:
for phrase, t in train_data:
f.write(f"{labels[t]}\t{phrase}\n")
if verbose:
print(f"Saved codebooks.npz ({npz_size:,} bytes)")
print(f"Saved config.json")
print(f"Saved phrases.txt")
print(f"Output: {output_dir.resolve()}")
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--output", type=str, default=".")
parser.add_argument("--quiet", action="store_true")
args = parser.parse_args()
train(args.output, verbose=not args.quiet)