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.
- jeffy/__init__.py +7 -0
- jeffy/benchmark_page.py +200 -0
- jeffy/build_pack.py +355 -0
- jeffy/catalog.py +148 -0
- jeffy/compare_jeff.py +454 -0
- jeffy/engine.py +194 -0
- jeffy/evaluate.py +401 -0
- jeffy/examples/__init__.py +15 -0
- jeffy/examples/reviews.csv +25 -0
- jeffy/model_pack.py +145 -0
- jeffy/pack/ag_news/manifest.json +31 -0
- jeffy/pack/ag_news/model.npz +0 -0
- jeffy/pack/banking77/manifest.json +177 -0
- jeffy/pack/banking77/model.npz +0 -0
- jeffy/pack/clinc_oos/manifest.json +325 -0
- jeffy/pack/clinc_oos/model.npz +0 -0
- jeffy/pack/dbpedia/manifest.json +51 -0
- jeffy/pack/dbpedia/model.npz +0 -0
- jeffy/pack/emotion/manifest.json +35 -0
- jeffy/pack/emotion/model.npz +0 -0
- jeffy/pack/imdb/manifest.json +27 -0
- jeffy/pack/imdb/model.npz +0 -0
- jeffy/pack/massive_intent/manifest.json +143 -0
- jeffy/pack/massive_intent/model.npz +0 -0
- jeffy/pack/pack_manifest.json +152 -0
- jeffy/pack/sms_spam/manifest.json +27 -0
- jeffy/pack/sms_spam/model.npz +0 -0
- jeffy/pack/snli/manifest.json +29 -0
- jeffy/pack/snli/model.npz +0 -0
- jeffy/pack/sst2/manifest.json +27 -0
- jeffy/pack/sst2/model.npz +0 -0
- jeffy/pack/tweet_eval_emotion/manifest.json +31 -0
- jeffy/pack/tweet_eval_emotion/model.npz +0 -0
- jeffy/pack/tweet_eval_offensive/manifest.json +27 -0
- jeffy/pack/tweet_eval_offensive/model.npz +0 -0
- jeffy/pack/tweet_eval_sentiment/manifest.json +29 -0
- jeffy/pack/tweet_eval_sentiment/model.npz +0 -0
- jeffy/server.py +532 -0
- jeffy/train.py +281 -0
- jeffy_classify-0.1.0a8.dist-info/METADATA +325 -0
- jeffy_classify-0.1.0a8.dist-info/RECORD +44 -0
- jeffy_classify-0.1.0a8.dist-info/WHEEL +4 -0
- jeffy_classify-0.1.0a8.dist-info/entry_points.txt +6 -0
- jeffy_classify-0.1.0a8.dist-info/licenses/LICENSE +21 -0
jeffy/__init__.py
ADDED
jeffy/benchmark_page.py
ADDED
|
@@ -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()
|