Text Classification
Transformers
ONNX
Safetensors
bert
social-media
content-filtering
engagement-bait
distillation
text-embeddings-inference
Instructions to use selftaughtdev/engagement-farm-classifier with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use selftaughtdev/engagement-farm-classifier with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="selftaughtdev/engagement-farm-classifier")# Load model directly from transformers import AutoTokenizer, AutoModelForSequenceClassification tokenizer = AutoTokenizer.from_pretrained("selftaughtdev/engagement-farm-classifier") model = AutoModelForSequenceClassification.from_pretrained("selftaughtdev/engagement-farm-classifier", device_map="auto") - Notebooks
- Google Colab
- Kaggle
File size: 3,313 Bytes
48a842f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 | """Evaluate the trained ONNX model on the validation split with a threshold sweep.
Usage: python eval.py [--thresholds 0.5,0.4,0.3,0.2] [--min-confidence 0.85]
Mirrors train.py's split (same seed, filter, dedup) so the rows match training.
The serving threshold is a deployment choice: lower it to hide more bait at the
cost of hiding more genuine tweets.
"""
import argparse
import glob
import json
import os
import random
import numpy as np
from optimum.onnxruntime import ORTModelForSequenceClassification
from transformers import AutoTokenizer
SEED = 13
def val_rows(rows, min_conf, full=False):
data = [
r
for r in rows
if r.get("label") is not None and r.get("text") and r.get("confidence", 1.0) >= min_conf
]
seen, deduped = set(), []
for r in data:
if r["text"] not in seen:
seen.add(r["text"])
deduped.append(r)
random.Random(SEED).shuffle(deduped)
return deduped if full else deduped[int(len(deduped) * 0.9) :]
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--dir", default="model_out")
ap.add_argument("--labeled", default="labeled.json")
ap.add_argument("--thresholds", default="0.5,0.4,0.3,0.2")
ap.add_argument("--min-confidence", type=float, default=0.85)
ap.add_argument("--full", action="store_true", help="evaluate all rows, skip the 90/10 val split")
args = ap.parse_args()
thresholds = [float(t) for t in args.thresholds.split(",")]
d = args.dir
files = glob.glob(os.path.join(d, "onnx-int8", "*.onnx"))
if not files:
raise SystemExit(f"no artifacts in {d}/onnx-int8, run train.py first")
val = val_rows(json.load(open(args.labeled)), args.min_confidence, args.full)
onnx_file = max(files, key=os.path.getsize)
model = ORTModelForSequenceClassification.from_pretrained(
os.path.dirname(onnx_file), file_name=os.path.basename(onnx_file)
)
tok = AutoTokenizer.from_pretrained(d)
labels = np.array([int(bool(r["label"])) for r in val])
probs = np.array(
[
float(
model(**tok(r["text"], return_tensors="pt", truncation=True, max_length=128)).logits.softmax(-1)[0, 1]
)
for r in val
]
)
print(f"{d} (min_conf {args.min_confidence}): val {len(val)} rows, {int(labels.sum())} farming")
sweep(labels, probs, thresholds)
clean = [(l, p) for r, l, p in zip(val, labels, probs) if r.get("src") == "timeline"]
if clean:
clabels = np.array([l for l, _ in clean])
cprobs = np.array([p for _, p in clean])
print(f"clean-timeline-only: {len(clean)} rows, {int(clabels.sum())} farming")
sweep(clabels, cprobs, thresholds)
def sweep(labels, probs, thresholds):
for t in thresholds:
preds = (probs >= t).astype(int)
tp = int(((preds == 1) & (labels == 1)).sum())
fp = int(((preds == 1) & (labels == 0)).sum())
fn = int(((preds == 0) & (labels == 1)).sum())
prec = tp / (tp + fp) if tp + fp else 0.0
rec = tp / (tp + fn) if tp + fn else 0.0
f1 = 2 * prec * rec / (prec + rec) if prec + rec else 0.0
print(f" t={t}: TP={tp} FP={fp} FN={fn} precision={prec:.2f} recall={rec:.2f} f1={f1:.2f}")
if __name__ == "__main__":
main()
|