| |
| """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() |
|
|