jeffy-classify 0.1.0a8__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (44) hide show
  1. jeffy/__init__.py +7 -0
  2. jeffy/benchmark_page.py +200 -0
  3. jeffy/build_pack.py +355 -0
  4. jeffy/catalog.py +148 -0
  5. jeffy/compare_jeff.py +454 -0
  6. jeffy/engine.py +194 -0
  7. jeffy/evaluate.py +401 -0
  8. jeffy/examples/__init__.py +15 -0
  9. jeffy/examples/reviews.csv +25 -0
  10. jeffy/model_pack.py +145 -0
  11. jeffy/pack/ag_news/manifest.json +31 -0
  12. jeffy/pack/ag_news/model.npz +0 -0
  13. jeffy/pack/banking77/manifest.json +177 -0
  14. jeffy/pack/banking77/model.npz +0 -0
  15. jeffy/pack/clinc_oos/manifest.json +325 -0
  16. jeffy/pack/clinc_oos/model.npz +0 -0
  17. jeffy/pack/dbpedia/manifest.json +51 -0
  18. jeffy/pack/dbpedia/model.npz +0 -0
  19. jeffy/pack/emotion/manifest.json +35 -0
  20. jeffy/pack/emotion/model.npz +0 -0
  21. jeffy/pack/imdb/manifest.json +27 -0
  22. jeffy/pack/imdb/model.npz +0 -0
  23. jeffy/pack/massive_intent/manifest.json +143 -0
  24. jeffy/pack/massive_intent/model.npz +0 -0
  25. jeffy/pack/pack_manifest.json +152 -0
  26. jeffy/pack/sms_spam/manifest.json +27 -0
  27. jeffy/pack/sms_spam/model.npz +0 -0
  28. jeffy/pack/snli/manifest.json +29 -0
  29. jeffy/pack/snli/model.npz +0 -0
  30. jeffy/pack/sst2/manifest.json +27 -0
  31. jeffy/pack/sst2/model.npz +0 -0
  32. jeffy/pack/tweet_eval_emotion/manifest.json +31 -0
  33. jeffy/pack/tweet_eval_emotion/model.npz +0 -0
  34. jeffy/pack/tweet_eval_offensive/manifest.json +27 -0
  35. jeffy/pack/tweet_eval_offensive/model.npz +0 -0
  36. jeffy/pack/tweet_eval_sentiment/manifest.json +29 -0
  37. jeffy/pack/tweet_eval_sentiment/model.npz +0 -0
  38. jeffy/server.py +532 -0
  39. jeffy/train.py +281 -0
  40. jeffy_classify-0.1.0a8.dist-info/METADATA +325 -0
  41. jeffy_classify-0.1.0a8.dist-info/RECORD +44 -0
  42. jeffy_classify-0.1.0a8.dist-info/WHEEL +4 -0
  43. jeffy_classify-0.1.0a8.dist-info/entry_points.txt +6 -0
  44. jeffy_classify-0.1.0a8.dist-info/licenses/LICENSE +21 -0
jeffy/__init__.py ADDED
@@ -0,0 +1,7 @@
1
+ """Jeffy: a pretrained local decision engine.
2
+
3
+ Reusable embeddings + small classifiers, shipped with pretrained
4
+ capabilities and reproducible benchmarks.
5
+ """
6
+
7
+ __version__ = "0.1.0a7"
@@ -0,0 +1,200 @@
1
+ """Generate benchmark HTML page from evaluation results.
2
+
3
+ Usage:
4
+ python -m jeffy.benchmark_page
5
+ """
6
+
7
+ import json
8
+ from pathlib import Path
9
+
10
+
11
+ def generate_html(results_path: str = "data/eval_results/benchmark.json",
12
+ jeff_path: str = "data/eval_results/jeff_comparison.json",
13
+ output_path: str = "data/eval_results/benchmark.html"):
14
+ with open(results_path) as f:
15
+ data = json.load(f)
16
+
17
+ # Load Jeff results if available
18
+ jeff_data = {}
19
+ jeff_note = "not evaluated"
20
+ jeff_path = Path(jeff_path)
21
+ if jeff_path.exists():
22
+ with open(jeff_path) as f:
23
+ jeff_raw = json.load(f)
24
+ for r in jeff_raw.get("results", []):
25
+ jeff_data[r["task_id"]] = r
26
+ jeff_note = jeff_raw.get("note", "")
27
+
28
+ results = sorted(data["results"], key=lambda r: -r["accuracy"])
29
+ hw = data.get("hardware", {})
30
+
31
+ rows = ""
32
+ for r in results:
33
+ bl_maj = r["baselines"]["majority"]["accuracy"] if r.get("baselines") else None
34
+ bl_tfidf = r["baselines"]["tfidf_lr"]["accuracy"] if r.get("baselines") else None
35
+ lat = r["latency"]["total_p50_ms"] if r.get("latency") else None
36
+ emb_lat = r["latency"]["embed_p50_ms"] if r.get("latency") else None
37
+ clf_lat = r["latency"]["clf_p50_ms"] if r.get("latency") else None
38
+ notes = "; ".join(r.get("notes", []))
39
+ ci = f'[{r["accuracy_ci_lo"]:.3f}, {r["accuracy_ci_hi"]:.3f}]'
40
+
41
+ # Clinc detail
42
+ extra = ""
43
+ if r.get("clinc_detail"):
44
+ cd = r["clinc_detail"]
45
+ extra = f' (in-scope {cd["in_scope_acc"]:.1%}, OOS {cd["oos_acc"]:.1%})'
46
+
47
+ artifact_kb = {"banking77": 334, "clinc_oos": 632, "massive_intent": 271,
48
+ "ag_news": 41, "dbpedia": 81, "emotion": 49, "imdb": 29,
49
+ "sms_spam": 29, "snli": 37, "sst2": 29,
50
+ "tweet_eval_emotion": 41, "tweet_eval_offensive": 29,
51
+ "tweet_eval_sentiment": 37}.get(r["task_id"], "?")
52
+
53
+ rows += f"""<tr>
54
+ <td><strong>{r['task_id']}</strong><br><small>{r['name']}</small></td>
55
+ <td>{r['n_classes']}</td>
56
+ <td>{r['test_examples']}</td>
57
+ <td>{r['eval_split']}</td>
58
+ <td><strong>{r['accuracy']:.1%}</strong><br><small>{ci}</small></td>
59
+ <td>{r['macro_f1']:.1%}</td>
60
+ <td>{f'{bl_maj:.1%}' if bl_maj is not None else '—'}</td>
61
+ <td>{f'{bl_tfidf:.1%}' if bl_tfidf is not None else '—'}</td>
62
+ <td>{f'{jeff_data[r["task_id"]]["jeff"]["accuracy_all"]:.1%}' if r['task_id'] in jeff_data and 'jeff' in jeff_data[r['task_id']] else (f'{jeff_data[r["task_id"]].get("jeff_accuracy",0):.1%}' if r['task_id'] in jeff_data and 'jeff_accuracy' in jeff_data[r['task_id']] else 'not evaluated')}</td>
63
+ <td>{lat:.0f}ms<br><small>emb {emb_lat:.0f} + clf {clf_lat:.2f}</small></td>
64
+ <td>{artifact_kb} KB</td>
65
+ </tr>\n"""
66
+
67
+ avg_acc = sum(r["accuracy"] for r in results) / len(results)
68
+ avg_f1 = sum(r["macro_f1"] for r in results) / len(results)
69
+
70
+ html = f"""<!DOCTYPE html>
71
+ <html lang="en">
72
+ <head>
73
+ <meta charset="utf-8">
74
+ <meta name="viewport" content="width=device-width,initial-scale=1">
75
+ <title>Jeffy Benchmark</title>
76
+ <style>
77
+ :root{{--mono:ui-monospace,SFMono-Regular,Consolas,monospace;--accent:#2563eb}}
78
+ *{{box-sizing:border-box}}
79
+ body{{margin:0;font:14px/1.6 var(--mono);color:#111;background:#fff;padding:32px}}
80
+ .container{{max-width:1400px;margin:auto}}
81
+ h1{{font-size:28px;font-weight:700;margin:0 0 4px}}
82
+ .sub{{color:#555;font-size:13px;margin-bottom:24px}}
83
+ h2{{font-size:18px;margin:32px 0 12px;border-bottom:1px solid #e5e5e5;padding-bottom:8px}}
84
+ table{{border-collapse:collapse;width:100%;font-size:12px}}
85
+ th,td{{padding:8px 10px;text-align:left;border-bottom:1px solid #eee}}
86
+ th{{background:#f8f8f8;font-weight:600;position:sticky;top:0}}
87
+ tr:hover{{background:#f0f7ff}}
88
+ small{{color:#777}}
89
+ strong{{font-weight:600}}
90
+ .meta{{font-size:12px;color:#555;margin:16px 0;line-height:1.8}}
91
+ .meta dt{{font-weight:600;display:inline}}
92
+ .meta dd{{display:inline;margin:0 16px 0 0}}
93
+ .note{{background:#fffbeb;border:1px solid #fde68a;border-radius:6px;padding:12px;font-size:12px;margin:16px 0}}
94
+ .arch{{background:#f0f7ff;border:1px solid #bfdbfe;border-radius:6px;padding:16px;font-size:13px;margin:16px 0}}
95
+ code{{background:#f5f5f5;padding:2px 5px;border-radius:3px;font-size:12px}}
96
+ a{{color:var(--accent)}}
97
+ </style>
98
+ </head>
99
+ <body>
100
+ <div class="container">
101
+ <h1>Jeffy Benchmark</h1>
102
+ <p class="sub">Frozen pretrained embeddings + task-specific logistic heads. Evaluated {data['evaluation_date']}.</p>
103
+
104
+ <div class="arch">
105
+ <strong>Architecture:</strong> BAAI/bge-large-en-v1.5 (1024d frozen encoder, ~1.2 GB) → StandardScaler → LogisticRegression (C=0.01, newton-cg)<br>
106
+ <strong>Each head:</strong> 1 scaler + 1 linear classifier. No fine-tuning, no MLP, no generation.<br>
107
+ <strong>Training regime:</strong> Each head trained on up to 10,000 examples from the dataset's training split (seed=42 for subsampling).<br>
108
+ <strong>Comparison context:</strong> Jeffy heads are task-trained; Jeff and similar zero-shot models are not. This is a deployment comparison, not an equal-supervision experiment.<br>
109
+ <strong>Baselines:</strong> TF-IDF+LR tuned via 3-fold CV over C in {{0.01, 0.1, 1, 10}}, analyzer in {{word, char_wb}}, ngram_range in {{(1,1), (1,2)}}, max_features in {{10000, 30000}}. Vectorizer fit inside each fold.
110
+ </div>
111
+
112
+ <h2>Results</h2>
113
+ <table>
114
+ <thead>
115
+ <tr>
116
+ <th>Task</th><th>Classes</th><th>Test N</th><th>Split</th>
117
+ <th>Accuracy</th><th>Macro F1</th>
118
+ <th>Majority</th><th>TF-IDF+LR</th><th>Jeff</th>
119
+ <th>Latency (p50)</th><th>Head Size</th>
120
+ </tr>
121
+ </thead>
122
+ <tbody>
123
+ {rows}
124
+ <tr style="font-weight:600;border-top:2px solid #333">
125
+ <td>Average (13 tasks)</td><td></td><td></td><td></td>
126
+ <td>{avg_acc:.1%}</td><td>{avg_f1:.1%}</td>
127
+ <td></td><td></td><td></td><td></td><td></td>
128
+ </tr>
129
+ </tbody>
130
+ </table>
131
+
132
+ <div class="note">
133
+ <strong>Jeff comparison:</strong> Not yet evaluated locally. Jeff is a zero-shot decision model (no per-task training).
134
+ A fair comparison requires running Jeff on identical evaluation examples with the same label meanings.
135
+ Jeffy's advantage comes from task-specific training; Jeff's advantage is generalization without training data.
136
+ </div>
137
+
138
+ <h2>Task-Specific Notes</h2>
139
+ <dl class="meta">
140
+ <dt>SST-2:</dt><dd>Evaluated on validation split; official test labels are not public.</dd>
141
+ <dt>SMS Spam:</dt><dd>Random train/test split (test_size=0.2, seed=42); no standard benchmark split.</dd>
142
+ <dt>SNLI:</dt><dd>Input encoded as "premise [SEP] hypothesis". Label -1 (unlabeled) filtered.</dd>
143
+ <dt>CLINC-OOS:</dt><dd>151 classes including out-of-scope. In-scope accuracy {[r for r in results if r['task_id']=='clinc_oos'][0].get('clinc_detail',{}).get('in_scope_acc','?'):.1%}, OOS detection {[r for r in results if r['task_id']=='clinc_oos'][0].get('clinc_detail',{}).get('oos_acc','?'):.1%}.</dd>
144
+ <dt>MASSIVE:</dt><dd>English subset only (config='en'). String labels used directly.</dd>
145
+ </dl>
146
+
147
+ <h2>Infrastructure</h2>
148
+ <dl class="meta">
149
+ <dt>Hardware:</dt><dd>{hw.get('platform','?')}, {hw.get('cpu_count','?')} cores, {hw.get('ram_gb','?')} GB RAM</dd>
150
+ <dt>Encoder:</dt><dd>BAAI/bge-large-en-v1.5, ~1.2 GB on disk, {data.get('encoder_load_time_s','?')}s load time</dd>
151
+ <dt>Total pack:</dt><dd>{data.get('pack_total_bytes',0)/1024:.0f} KB (13 heads, no encoder)</dd>
152
+ <dt>Threads:</dt><dd>Recorded: device={hw.get('device','?')}, torch_threads={hw.get('torch_threads','?')}, OMP_NUM_THREADS={hw.get('omp_threads','?')}</dd>
153
+ <dt>Batch size:</dt><dd>1 (single-example inference measured)</dd>
154
+ <dt>Caching:</dt><dd>No embedding cache; each request embeds fresh</dd>
155
+ <dt>Truncation:</dt><dd>Default sentence-transformers truncation (512 tokens)</dd>
156
+ </dl>
157
+
158
+ <h2>Reproduce</h2>
159
+ <pre><code># Install
160
+ pip install sentence-transformers scikit-learn fastapi uvicorn
161
+
162
+ # Build model pack from source datasets
163
+ python -m jeffy.build_pack --out data/model_pack
164
+
165
+ # Run evaluation
166
+ python -m jeffy.evaluate --baselines --latency
167
+
168
+ # Start server with playground
169
+ python -m jeffy.server
170
+ # Open http://localhost:8400
171
+ </code></pre>
172
+
173
+ <h2>API</h2>
174
+ <pre><code># Predict with a pretrained head
175
+ curl -X POST http://localhost:8400/v1/predict \\
176
+ -H "Content-Type: application/json" \\
177
+ -d '{{"text": "I was charged twice", "task": "banking77"}}'
178
+
179
+ # List capabilities
180
+ curl http://localhost:8400/v1/capabilities
181
+
182
+ # Jeff-compatible endpoint (matches criteria to pretrained heads)
183
+ curl -X POST http://localhost:8400/v1/systemone \\
184
+ -H "Content-Type: application/json" \\
185
+ -d '{{"model":"jeffy","state":"I was charged twice",
186
+ "questions":{{"intent":{{"type":"choice","criteria":{{
187
+ "transaction_charged_twice":null,"request_refund":null,
188
+ "cancel_transfer":null}}}}}}}}'
189
+ </code></pre>
190
+
191
+ </div>
192
+ </body>
193
+ </html>"""
194
+
195
+ Path(output_path).write_text(html)
196
+ print(f"Benchmark page: {output_path}")
197
+
198
+
199
+ if __name__ == "__main__":
200
+ generate_html()
jeffy/build_pack.py ADDED
@@ -0,0 +1,355 @@
1
+ """Build a verified model pack from source datasets.
2
+
3
+ Trains each dataset independently, saves artifacts with explicit identity,
4
+ recovers label mappings from the dataset features, and verifies roundtrip
5
+ prediction fidelity.
6
+
7
+ Usage:
8
+ python -m jeffy.build_pack
9
+ python -m jeffy.build_pack --datasets banking77 ag_news
10
+ python -m jeffy.build_pack --out data/model_pack
11
+ """
12
+
13
+ import argparse
14
+ import json
15
+ import time
16
+ from pathlib import Path
17
+
18
+ import numpy as np
19
+ from sklearn.linear_model import LogisticRegression
20
+ from sklearn.metrics import accuracy_score, f1_score
21
+ from sklearn.preprocessing import LabelEncoder, StandardScaler
22
+
23
+ from .catalog import ENCODER, ENCODER_DIM, LICENSES
24
+ from .model_pack import ArtifactManifest, coef_hash, save_artifact, scaler_hash
25
+
26
+ # Dataset definitions with label recovery instructions
27
+ DATASETS = {
28
+ "banking77": {
29
+ "hf": "legacy-datasets/banking77",
30
+ "text": "text", "label": "label",
31
+ "description": "Banking customer service intent detection",
32
+ "task_type": "choice",
33
+ "label_source": "features", # labels from dataset.features['label'].names
34
+ },
35
+ "clinc_oos": {
36
+ "hf": "clinc_oos", "hf_config": "plus",
37
+ "text": "text", "label": "intent",
38
+ "description": "Intent detection with out-of-scope",
39
+ "task_type": "choice",
40
+ "label_source": "features",
41
+ },
42
+ "massive_intent": {
43
+ "hf": "mteb/amazon_massive_intent", "hf_config": "en",
44
+ "text": "text", "label": "label",
45
+ "description": "Amazon MASSIVE voice command intents",
46
+ "task_type": "choice",
47
+ "label_source": "string_values", # labels ARE the string values
48
+ },
49
+ "ag_news": {
50
+ "hf": "fancyzhx/ag_news",
51
+ "text": "text", "label": "label",
52
+ "description": "News article topic classification",
53
+ "task_type": "choice",
54
+ "label_source": "features",
55
+ },
56
+ "dbpedia": {
57
+ "hf": "fancyzhx/dbpedia_14",
58
+ "text": "content", "label": "label",
59
+ "description": "Wikipedia article ontology classification",
60
+ "task_type": "choice",
61
+ "label_source": "features",
62
+ },
63
+ "sst2": {
64
+ "hf": "stanfordnlp/sst2",
65
+ "text": "sentence", "label": "label",
66
+ "split_test": "validation",
67
+ "description": "Movie review sentiment (positive/negative)",
68
+ "task_type": "noul",
69
+ "label_source": "manual",
70
+ "manual_labels": {"0": "negative", "1": "positive"},
71
+ },
72
+ "emotion": {
73
+ "hf": "dair-ai/emotion",
74
+ "text": "text", "label": "label",
75
+ "description": "Text emotion detection",
76
+ "task_type": "choice",
77
+ "label_source": "features",
78
+ },
79
+ "imdb": {
80
+ "hf": "stanfordnlp/imdb",
81
+ "text": "text", "label": "label",
82
+ "description": "Movie review sentiment (positive/negative)",
83
+ "task_type": "noul",
84
+ "label_source": "manual",
85
+ "manual_labels": {"0": "negative", "1": "positive"},
86
+ },
87
+ "sms_spam": {
88
+ "hf": "ucirvine/sms_spam",
89
+ "text": "sms", "label": "label",
90
+ "description": "SMS spam detection",
91
+ "task_type": "noul",
92
+ "label_source": "manual",
93
+ "manual_labels": {"0": "ham", "1": "spam"},
94
+ },
95
+ "snli": {
96
+ "hf": "stanfordnlp/snli",
97
+ "text": ["premise", "hypothesis"], "label": "label",
98
+ "filter_label": -1,
99
+ "description": "Natural language inference",
100
+ "task_type": "choice",
101
+ "label_source": "manual",
102
+ "manual_labels": {"0": "entailment", "1": "neutral", "2": "contradiction"},
103
+ },
104
+ "tweet_eval_sentiment": {
105
+ "hf": "cardiffnlp/tweet_eval", "hf_config": "sentiment",
106
+ "text": "text", "label": "label",
107
+ "description": "Tweet sentiment analysis",
108
+ "task_type": "choice",
109
+ "label_source": "manual",
110
+ "manual_labels": {"0": "negative", "1": "neutral", "2": "positive"},
111
+ },
112
+ "tweet_eval_emotion": {
113
+ "hf": "cardiffnlp/tweet_eval", "hf_config": "emotion",
114
+ "text": "text", "label": "label",
115
+ "description": "Tweet emotion detection",
116
+ "task_type": "choice",
117
+ "label_source": "manual",
118
+ "manual_labels": {"0": "anger", "1": "joy", "2": "optimism", "3": "sadness"},
119
+ },
120
+ "tweet_eval_offensive": {
121
+ "hf": "cardiffnlp/tweet_eval", "hf_config": "offensive",
122
+ "text": "text", "label": "label",
123
+ "description": "Offensive language detection",
124
+ "task_type": "noul",
125
+ "label_source": "manual",
126
+ "manual_labels": {"0": "not_offensive", "1": "offensive"},
127
+ },
128
+ }
129
+
130
+
131
+ def recover_labels(ds, config) -> dict[str, str]:
132
+ """Recover human-readable label mapping from the dataset source."""
133
+ source = config["label_source"]
134
+
135
+ if source == "manual":
136
+ return config["manual_labels"]
137
+
138
+ if source == "features":
139
+ label_col = config["label"]
140
+ feat = ds["train"].features[label_col]
141
+ if hasattr(feat, "names"):
142
+ return {str(i): name for i, name in enumerate(feat.names)}
143
+ raise ValueError(f"No .names on feature {label_col}")
144
+
145
+ if source == "string_values":
146
+ # Labels are the string values themselves
147
+ train_labels = sorted(set(ds["train"][config["label"]]))
148
+ return {label: label for label in train_labels}
149
+
150
+ raise ValueError(f"Unknown label_source: {source}")
151
+
152
+
153
+ def load_and_split(ds_name, config, max_train=10000, max_test=2000):
154
+ """Load dataset, extract texts and string labels."""
155
+ from datasets import load_dataset
156
+ hf_config = config.get("hf_config")
157
+ ds = load_dataset(config["hf"], hf_config) if hf_config else load_dataset(config["hf"])
158
+
159
+ split_test = config.get("split_test", "test")
160
+ if split_test in ds:
161
+ train_ds, test_ds = ds["train"], ds[split_test]
162
+ elif "validation" in ds:
163
+ train_ds, test_ds = ds["train"], ds["validation"]
164
+ else:
165
+ split = ds["train"].train_test_split(test_size=0.2, seed=42)
166
+ train_ds, test_ds = split["train"], split["test"]
167
+
168
+ text_col = config["text"]
169
+ if isinstance(text_col, list):
170
+ train_texts = [" [SEP] ".join(train_ds[c][i] for c in text_col) for i in range(len(train_ds))]
171
+ test_texts = [" [SEP] ".join(test_ds[c][i] for c in text_col) for i in range(len(test_ds))]
172
+ else:
173
+ train_texts = list(train_ds[text_col])
174
+ test_texts = list(test_ds[text_col])
175
+
176
+ train_labels = [str(x) for x in train_ds[config["label"]]]
177
+ test_labels = [str(x) for x in test_ds[config["label"]]]
178
+
179
+ # Filter invalid
180
+ if "filter_label" in config:
181
+ fl = str(config["filter_label"])
182
+ mask = [l != fl for l in train_labels]
183
+ train_texts = [t for t, m in zip(train_texts, mask) if m]
184
+ train_labels = [l for l, m in zip(train_labels, mask) if m]
185
+ mask = [l != fl for l in test_labels]
186
+ test_texts = [t for t, m in zip(test_texts, mask) if m]
187
+ test_labels = [l for l, m in zip(test_labels, mask) if m]
188
+
189
+ # Cap sizes
190
+ if len(train_texts) > max_train:
191
+ np.random.seed(42)
192
+ idx = np.random.choice(len(train_texts), max_train, replace=False)
193
+ train_texts = [train_texts[i] for i in idx]
194
+ train_labels = [train_labels[i] for i in idx]
195
+
196
+ if len(test_texts) > max_test:
197
+ np.random.seed(42)
198
+ idx = np.random.choice(len(test_texts), max_test, replace=False)
199
+ test_texts = [test_texts[i] for i in idx]
200
+ test_labels = [test_labels[i] for i in idx]
201
+
202
+ labels = recover_labels(ds, config)
203
+ return train_texts, train_labels, test_texts, test_labels, labels
204
+
205
+
206
+ def main():
207
+ parser = argparse.ArgumentParser(description="Build verified model pack")
208
+ parser.add_argument("--datasets", nargs="+", default=list(DATASETS.keys()))
209
+ parser.add_argument("--out", default="data/model_pack")
210
+ parser.add_argument("--C", type=float, default=0.01)
211
+ parser.add_argument("--max-train", type=int, default=10000)
212
+ parser.add_argument("--max-test", type=int, default=2000)
213
+ args = parser.parse_args()
214
+
215
+ out_dir = Path(args.out)
216
+ out_dir.mkdir(parents=True, exist_ok=True)
217
+
218
+ print(f"Building model pack: {len(args.datasets)} datasets")
219
+ print(f"Output: {out_dir}")
220
+ print(f"Encoder: {ENCODER}")
221
+ print()
222
+
223
+ from sentence_transformers import SentenceTransformer
224
+ encoder = SentenceTransformer(ENCODER)
225
+ results = []
226
+
227
+ for ds_name in args.datasets:
228
+ if ds_name not in DATASETS:
229
+ print(f"Unknown dataset: {ds_name}")
230
+ continue
231
+
232
+ config = DATASETS[ds_name]
233
+ print(f"{'='*60}")
234
+ print(f"{ds_name}: {config['description']}")
235
+
236
+ try:
237
+ train_texts, train_labels, test_texts, test_labels, labels = \
238
+ load_and_split(ds_name, config, args.max_train, args.max_test)
239
+ except Exception as e:
240
+ print(f" FAILED to load: {e}")
241
+ continue
242
+
243
+ n_classes = len(set(train_labels))
244
+ print(f" Train: {len(train_texts)}, Test: {len(test_texts)}, Classes: {n_classes}")
245
+ print(f" Labels: {list(labels.items())[:5]}...")
246
+
247
+ # Embed
248
+ t0 = time.perf_counter()
249
+ train_emb = encoder.encode(train_texts, batch_size=256, show_progress_bar=False)
250
+ test_emb = encoder.encode(test_texts, batch_size=256, show_progress_bar=False)
251
+ t_embed = time.perf_counter() - t0
252
+
253
+ # Train
254
+ t1 = time.perf_counter()
255
+ scaler = StandardScaler()
256
+ X_train = scaler.fit_transform(train_emb)
257
+ X_test = scaler.transform(test_emb)
258
+
259
+ clf = LogisticRegression(C=args.C, max_iter=3000, solver="newton-cg", random_state=42)
260
+ clf.fit(X_train, train_labels)
261
+ t_train = time.perf_counter() - t1
262
+
263
+ # Evaluate
264
+ train_preds = clf.predict(X_train)
265
+ test_preds = clf.predict(X_test)
266
+ train_acc = accuracy_score(train_labels, train_preds)
267
+ test_acc = accuracy_score(test_labels, test_preds)
268
+ test_f1 = f1_score(test_labels, test_preds, average="macro", zero_division=0)
269
+
270
+ # Verify probability mapping roundtrip
271
+ probs = clf.predict_proba(X_test[:3])
272
+ for i in range(min(3, len(probs))):
273
+ pred_from_probs = clf.classes_[np.argmax(probs[i])]
274
+ pred_from_predict = test_preds[i]
275
+ assert pred_from_probs == pred_from_predict, \
276
+ f"Probability mapping mismatch at test[{i}]: {pred_from_probs} vs {pred_from_predict}"
277
+
278
+ # Verify all class IDs have label mappings
279
+ for cls_id in clf.classes_:
280
+ if str(cls_id) not in labels:
281
+ print(f" WARNING: class {cls_id} has no label mapping")
282
+
283
+ manifest = ArtifactManifest(
284
+ dataset=ds_name,
285
+ description=config["description"],
286
+ n_classes=n_classes,
287
+ task_type=config["task_type"],
288
+ encoder=ENCODER,
289
+ encoder_dim=ENCODER_DIM,
290
+ labels=labels,
291
+ class_order=list(clf.classes_),
292
+ hf_source=config["hf"],
293
+ hf_config=config.get("hf_config"),
294
+ hf_revision=None, # TODO: pin revision
295
+ license=LICENSES.get(ds_name, "see source"),
296
+ train_examples=len(train_texts),
297
+ train_accuracy=round(train_acc, 6),
298
+ test_examples=len(test_texts),
299
+ test_accuracy=round(test_acc, 6),
300
+ scaler_mean_hash=scaler_hash(scaler),
301
+ classifier_coef_hash=coef_hash(clf),
302
+ )
303
+
304
+ save_artifact(out_dir, ds_name, clf, scaler, manifest)
305
+
306
+ print(f" Train acc: {train_acc:.4f}")
307
+ print(f" Test acc: {test_acc:.4f}")
308
+ print(f" Test F1: {test_f1:.4f}")
309
+ print(f" Time: embed={t_embed:.1f}s train={t_train:.1f}s")
310
+ print(f" Saved to {out_dir / ds_name}/")
311
+
312
+ results.append({
313
+ "dataset": ds_name,
314
+ "n_classes": n_classes,
315
+ "train_size": len(train_texts),
316
+ "test_size": len(test_texts),
317
+ "train_acc": round(train_acc, 6),
318
+ "test_acc": round(test_acc, 6),
319
+ "test_f1": round(test_f1, 6),
320
+ "embed_time": round(t_embed, 1),
321
+ "train_time": round(t_train, 1),
322
+ })
323
+ print()
324
+
325
+ # Summary
326
+ print(f"\n{'='*80}")
327
+ print("MODEL PACK SUMMARY")
328
+ print(f"{'='*80}")
329
+ print(f"{'Dataset':<25s} {'Classes':>7s} {'Train':>7s} {'Test':>7s} {'Test Acc':>9s} {'F1':>8s}")
330
+ print("-" * 67)
331
+ for r in sorted(results, key=lambda x: -x["test_acc"]):
332
+ print(f"{r['dataset']:<25s} {r['n_classes']:>7d} {r['train_size']:>7d} "
333
+ f"{r['test_size']:>7d} {r['test_acc']:>9.4f} {r['test_f1']:>8.4f}")
334
+
335
+ avg_acc = np.mean([r["test_acc"] for r in results])
336
+ avg_f1 = np.mean([r["test_f1"] for r in results])
337
+ print(f"\n{'Average':<25s} {'':>7s} {'':>7s} {'':>7s} {avg_acc:>9.4f} {avg_f1:>8.4f}")
338
+
339
+ # Save pack manifest
340
+ with open(out_dir / "pack_manifest.json", "w") as f:
341
+ json.dump({
342
+ "encoder": ENCODER,
343
+ "encoder_dim": ENCODER_DIM,
344
+ "classifier": "LogisticRegression(C=0.01, solver=newton-cg)",
345
+ "datasets": results,
346
+ "build_time": time.strftime("%Y-%m-%d %H:%M:%S"),
347
+ "schema_version": 1,
348
+ }, f, indent=2)
349
+
350
+ print(f"\nPack manifest: {out_dir / 'pack_manifest.json'}")
351
+ print(f"Total: {len(results)} artifacts")
352
+
353
+
354
+ if __name__ == "__main__":
355
+ main()