Download train_final.py from zeechimp/zee: direct link, hf CLI and curl.
- Browser
- Download file 9.25 kB
-
https://huggingface.co/zeechimp/zee/resolve/main/train_final.py
- Command line
-
hf download hf://zeechimp/zee/train_final.py
-
curl -L -o train_final.py https://huggingface.co/zeechimp/zee/resolve/main/train_final.py
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) |