themis / phase1 /scripts /classify_treatments.py
vg15o2's picture
Moonley backend (HF Space build)
1d9bd9b
Raw
History Blame Contribute Delete
6.37 kB
#!/usr/bin/env python3
"""Classify citation treatments over edges_v2 contexts -> a REAL good-law citator.
The at-scale build hardcoded every edge treatment to "cited"; good_law_status is
"unknown" for 99.9% of the corpus, so the bad-law guardrail was theater. Each edge in
edges_v2.jsonl carries `para` — the citing court's sentence about the precedent — which
is exactly the input a treatment classifier needs.
Labels: relied_on | followed | distinguished | doubted | overruled | cited (neutral).
Model: deepseek-v4-flash (non-thinking pinned), temperature 0, batched 20 edges/call. Requires DEEPSEEK_API_KEY.
Run: python phase1/scripts/classify_treatments.py [data_dir] [--dry-run N] [--limit N]
Out: <data_dir>/edges_treatment.jsonl {from, target, treatment} (resumable)
<data_dir>/good_law_v2.jsonl per-target rollup {doc_id, good_law_status,
treatment_breakdown, cited_by}
Rollup rule: any overruled edge -> "overruled"; >=2 doubted/distinguished and no
follow/rely since -> "doubted"; else "good" if >=3 positive treatments else "unknown".
"""
import json, os, sys, time
from collections import Counter, defaultdict
data_dir = next((a for a in sys.argv[1:] if not a.startswith("--")), None) \
or os.environ.get("THEMIS_DATA", "phase1/data/thor_artifacts")
DRY = 0
if "--dry-run" in sys.argv:
i = sys.argv.index("--dry-run"); DRY = int(sys.argv[i + 1]) if len(sys.argv) > i + 1 else 3
LIMIT = int(sys.argv[sys.argv.index("--limit") + 1]) if "--limit" in sys.argv else None
SYS = ("You classify how an Indian Supreme Court judgment TREATS a precedent it cites, from the "
"sentence surrounding the citation. Labels: relied_on (the holding rests on it), followed "
"(applied approvingly), distinguished (held inapplicable on facts/law), doubted (correctness "
"questioned), overruled (expressly overruled/no longer good law), cited (neutral mention). "
"Input: JSON list of {i, context}. Output STRICT JSON: {\"labels\": [{\"i\": n, "
"\"treatment\": \"...\"}]} — one per input, label from the list only.")
def load_edges():
edges = []
for line in open(os.path.join(data_dir, "edges_v2.jsonl"), encoding="utf-8"):
e = json.loads(line)
if len((e.get("para") or "")) >= 60:
edges.append(e)
return edges
def rollup(treat_by_edge, edges):
per_tgt = defaultdict(Counter); indeg = Counter()
for e in edges: indeg[e["target"]] += 1
for (f_, t_), lab in treat_by_edge.items(): per_tgt[t_][lab] += 1
out = os.path.join(data_dir, "good_law_v2.jsonl")
with open(out, "w", encoding="utf-8") as f:
for tgt, cnt in per_tgt.items():
if cnt.get("overruled"): status = "overruled"
elif cnt.get("doubted", 0) + cnt.get("distinguished", 0) >= 2 and \
cnt.get("relied_on", 0) + cnt.get("followed", 0) == 0: status = "doubted"
elif cnt.get("relied_on", 0) + cnt.get("followed", 0) >= 3: status = "good"
else: status = "unknown"
f.write(json.dumps({"doc_id": tgt, "good_law_status": status,
"treatment_breakdown": dict(cnt),
"cited_by": indeg.get(tgt, 0)}, ensure_ascii=False) + "\n")
print(f"[treat] rollup -> {out} ({len(per_tgt)} targets)", flush=True)
def main():
edges = load_edges()
outp = os.path.join(data_dir, "edges_treatment.jsonl")
done = set()
if os.path.exists(outp):
for line in open(outp, encoding="utf-8"):
r = json.loads(line); done.add((r["from"], r["target"]))
todo = [e for e in edges if (e["from"], e["target"]) not in done]
if LIMIT: todo = todo[:LIMIT]
print(f"[treat] edges with context: {len(edges)} | done: {len(done)} | todo: {len(todo)}", flush=True)
if DRY:
for e in todo[:DRY]:
print(f"\n--- DRY {e['from']} -> {e['target']} ---\n {e['para'][:220]}", flush=True)
print(f"\n[treat] dry-run only ({DRY} shown); no API calls made.", flush=True)
return
key = os.environ.get("DEEPSEEK_API_KEY", "")
if not key:
env = os.path.join(os.path.dirname(os.path.abspath(__file__)), ".env")
if os.path.exists(env):
for l in open(env):
if l.startswith("DEEPSEEK_API_KEY="): key = l.split("=", 1)[1].strip()
if not key:
sys.exit("[treat] DEEPSEEK_API_KEY missing (env or phase1/scripts/.env) — aborting.")
import requests
B = 20
with open(outp, "a", encoding="utf-8") as f:
for s in range(0, len(todo), B):
batch = todo[s:s + B]
payload = [{"i": i, "context": e["para"][:400]} for i, e in enumerate(batch)]
try:
r = requests.post("https://api.deepseek.com/chat/completions",
headers={"Authorization": f"Bearer {key}"},
json={"model": os.environ.get("THEMIS_LLM_MODEL", "deepseek-v4-flash"), "temperature": 0,
"thinking": {"type": "disabled"},
"response_format": {"type": "json_object"},
"messages": [{"role": "system", "content": SYS},
{"role": "user", "content": json.dumps(payload)}]},
timeout=120)
labs = json.loads(r.json()["choices"][0]["message"]["content"]).get("labels", [])
VALID = {"relied_on","followed","distinguished","doubted","overruled","cited"}
for l in labs:
i = l.get("i"); t = l.get("treatment")
if isinstance(i, int) and 0 <= i < len(batch) and t in VALID:
e = batch[i]
f.write(json.dumps({"from": e["from"], "target": e["target"],
"treatment": t}, ensure_ascii=False) + "\n")
f.flush()
except Exception as ex:
print(f"[treat] batch {s}: {ex}", flush=True); time.sleep(3)
if (s // B) % 25 == 0:
print(f"[treat] {min(s+B,len(todo))}/{len(todo)}", flush=True)
treat = {}
for line in open(outp, encoding="utf-8"):
r = json.loads(line); treat[(r["from"], r["target"])] = r["treatment"]
rollup(treat, edges)
if __name__ == "__main__":
main()