constructelligence's picture
Publish mf-0.2: completed GPU fine-tune (val 46.6%, 2.3x mf-0.1) + new model card with metrics/ONNX
b774c08 verified
Raw History Blame Contribute Delete
6.64 kB
"""Classify construction text into MasterFormat level-2 groups (171 classes, 32 divisions).
Examples
--------
pip install transformers torch
python predict.py "4000 psi concrete slab on grade" "TPO roofing, 60 mil, fully adhered"
python predict.py --top-k 5 --divisions "8\" CMU wall, grout filled at 32\" o.c."
echo "Cat 6 data cabling and jacks" | python predict.py -
python predict.py --file items.txt --json > out.json
ONNX (no torch, ~34 MB int8 weights):
pip install onnxruntime transformers
python predict.py --onnx onnx/model_quantized.onnx "Panelboards 208/120V 42 circuit"
Without arguments it classifies one built-in example.
"""
import argparse
import json
import sys
MODEL = "constructelligence/masterformat-classifier"
# Division code -> name. Kept local so --divisions needs no extra package.
DIVISIONS = {
"01": "General Requirements", "02": "Existing Conditions", "03": "Concrete",
"04": "Masonry", "05": "Metals", "06": "Wood, Plastics, and Composites",
"07": "Thermal and Moisture Protection", "08": "Openings", "09": "Finishes",
"10": "Specialties", "11": "Equipment", "12": "Furnishings",
"13": "Special Construction", "14": "Conveying Equipment", "21": "Fire Suppression",
"22": "Plumbing", "23": "HVAC", "25": "Integrated Automation", "26": "Electrical",
"27": "Communications", "28": "Electronic Safety and Security", "31": "Earthwork",
"32": "Exterior Improvements", "33": "Utilities", "34": "Transportation",
"35": "Waterway and Marine Construction", "40": "Process Interconnections",
"41": "Material Processing and Handling Equipment",
"43": "Process Gas and Liquid Handling, Purification, and Storage Equipment",
"44": "Pollution and Waste Control Equipment", "46": "Water and Wastewater Equipment",
"48": "Electrical Power Generation",
}
EXAMPLE = "4000 psi concrete slab on grade"
def split_label(label):
"""'03 30 00 Cast-in-Place Concrete' -> ('03 30 00', 'Cast-in-Place Concrete')."""
code, name = label[:8].strip(), label[8:].strip()
return code, (name or label)
def division_rollup(preds):
"""Sum level-2 probabilities by the first two digits (division)."""
totals = {}
for p in preds:
div = p["code"][:2]
totals[div] = totals.get(div, 0.0) + p["score"]
return [
{"code": d, "name": DIVISIONS.get(d, d), "score": s}
for d, s in sorted(totals.items(), key=lambda kv: -kv[1])
]
def hf_scorer(model, top_k):
from transformers import pipeline
clf = pipeline("text-classification", model=model, top_k=top_k)
return lambda texts: [[d for d in out] for out in clf(texts)]
def onnx_scorer(path, top_k):
"""Run onnx/model*.onnx directly. Mirrors scripts/export_onnx.py: inputs
input_ids / attention_mask / token_type_ids, output logits [batch, num_labels]."""
import numpy as np
import onnxruntime as ort
from pathlib import Path
from transformers import AutoTokenizer
# The labels and tokenizer live beside onnx/ in the repo; fall back to the Hub.
root = Path(path).resolve().parent.parent
src = str(root) if (root / "tokenizer.json").exists() or (root / "vocab.txt").exists() else MODEL
tok = AutoTokenizer.from_pretrained(src)
id2label = None
try:
c = json.loads((root / "config.json").read_text())
id2label = {int(k): v for k, v in c.get("id2label", {}).items()}
except Exception:
id2label = None
sess = ort.InferenceSession(path, providers=["CPUExecutionProvider"])
names = {i.name for i in sess.get_inputs()}
def score(texts):
enc = tok(texts, truncation=True, max_length=128, padding=True, return_tensors="np")
feed = {n: enc[n].astype(np.int64) for n in ("input_ids", "attention_mask", "token_type_ids") if n in names}
logits = sess.run(None, feed)[0]
idx = np.argsort(-logits, 1)[:, :top_k]
shifted = logits - logits.max(1, keepdims=True)
probs = np.exp(shifted) / np.exp(shifted).sum(1, keepdims=True)
out = []
for row, cols in zip(probs, idx):
out.append([{"label": id2label.get(int(c), str(int(c))), "score": float(row[c])} for c in cols])
return out
return score
def predict(texts, score, top_k, want_div):
results = []
for text, preds in zip(texts, score(texts)):
pl = []
for p in preds:
code, name = split_label(p["label"])
pl.append({"code": code, "label": p["label"], "name": name, "score": round(float(p["score"]), 4)})
item = {"text": text, "predictions": pl}
if want_div:
item["divisions"] = [dict(d, score=round(d["score"], 4)) for d in division_rollup(pl)]
results.append(item)
return results
def main(argv=None):
ap = argparse.ArgumentParser(description="MasterFormat level-2 classifier (171 groups, 32 divisions).")
ap.add_argument("text", nargs="*", help="text to classify; use '-' to read lines from stdin")
ap.add_argument("--model", default=MODEL, help="HF repo id or local checkpoint dir")
ap.add_argument("--onnx", metavar="PATH", help="classify with an ONNX model instead of PyTorch")
ap.add_argument("--top-k", type=int, default=3, help="number of level-2 predictions to show")
ap.add_argument("--divisions", action="store_true", help="also show the level-1 (division) roll-up")
ap.add_argument("--file", help="read newline-separated inputs from a file")
ap.add_argument("--json", action="store_true", help="emit JSON instead of a table")
a = ap.parse_args(argv)
texts = list(a.text)
if a.file:
texts += [l.rstrip("\n") for l in open(a.file, encoding="utf-8") if l.strip()]
if "-" in texts:
texts = [t for t in texts if t != "-"] + [l.rstrip("\n") for l in sys.stdin if l.strip()]
if not texts:
texts = [EXAMPLE]
texts = [t for t in texts if t.strip()]
if not texts:
ap.error("no input text")
score = onnx_scorer(a.onnx, a.top_k) if a.onnx else hf_scorer(a.model, a.top_k)
results = predict(texts, score, a.top_k, a.divisions)
if a.json:
json.dump(results, sys.stdout, indent=2, ensure_ascii=False)
print()
return
for r in results:
print(f"\n{r['text']}")
for p in r["predictions"]:
print(f" {p['code']} {p['name']:<48.48} {p['score']:.3f}")
if a.divisions:
print(" -- divisions --")
for d in r.get("divisions", []):
print(f" {d['code']} {d['name']:<48.48} {d['score']:.3f}")
if __name__ == "__main__":
main()