nwhite-systems's picture
Add training, inference and verification scripts
33947ea verified
Raw
History Blame Contribute Delete
1.92 kB
"""Export the fitted scikit-learn pipeline for local browser inference."""
from __future__ import annotations
import json
from pathlib import Path
import joblib
ROOT = Path(__file__).resolve().parents[1]
def export_web_model(model_path: Path, output_path: Path) -> dict[str, object]:
model = joblib.load(model_path)
vectorizer = model.named_steps["tfidf"]
classifier = model.named_steps["classifier"]
vocabulary = dict(sorted(vectorizer.vocabulary_.items(), key=lambda item: item[1]))
payload: dict[str, object] = {
"format": "nwhite-tfidf-logistic-regression-v1",
"classes": [str(label) for label in classifier.classes_],
"vectorizer": {
"vocabulary": vocabulary,
"idf": [float(value) for value in vectorizer.idf_],
"ngram_range": [int(value) for value in vectorizer.ngram_range],
"lowercase": bool(vectorizer.lowercase),
"sublinear_tf": bool(vectorizer.sublinear_tf),
"norm": vectorizer.norm,
},
"classifier": {
"coef": [[float(value) for value in row] for row in classifier.coef_],
"intercept": [float(value) for value in classifier.intercept_],
},
}
with output_path.open("w", encoding="utf-8", newline="\n") as handle:
json.dump(payload, handle, ensure_ascii=False, separators=(",", ":"))
handle.write("\n")
return payload
def main() -> None:
payload = export_web_model(ROOT / "model.joblib", ROOT / "web_model.json")
print(
json.dumps(
{
"status": "exported",
"format": payload["format"],
"classes": len(payload["classes"]),
"features": len(payload["vectorizer"]["idf"]),
"output": str(ROOT / "web_model.json"),
},
indent=2,
)
)
if __name__ == "__main__":
main()