Instructions to use vishwr/claim_drafter with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use vishwr/claim_drafter with PEFT:
from peft import PeftModel from transformers import AutoModelForCausalLM base_model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3.5-9B") model = PeftModel.from_pretrained(base_model, "vishwr/claim_drafter") - Notebooks
- Google Colab
- Kaggle
| #!/usr/bin/env python3 | |
| """Validate a claim-drafter JSONL dataset before spending training credits. | |
| python3 scripts/validate_dataset.py data/sft | |
| Checks structural integrity, claim-format correctness, dependency validity, | |
| input/target leakage, and length distribution. | |
| Usage: | |
| python3 validate_dataset.py <dir_with_train.jsonl_and_val.jsonl> | |
| """ | |
| import json | |
| import os | |
| import re | |
| import statistics | |
| import sys | |
| def ngrams(text, n=8): | |
| w = re.findall(r"[a-z]+", text.lower()) | |
| return set(tuple(w[i : i + n]) for i in range(len(w) - n + 1)) | |
| def check(path): | |
| rows = [] | |
| for i, line in enumerate(open(path), 1): | |
| try: | |
| rows.append(json.loads(line)) | |
| except json.JSONDecodeError as e: | |
| print(f" FAIL malformed JSON at line {i}: {e}") | |
| return None | |
| print(f"\n=== {os.path.basename(path)}: {len(rows)} examples ===") | |
| bad_roles = bad_first = fwd_ref = dep_no_antecedent = 0 | |
| ulens, alens, ncl, leaks = [], [], [], [] | |
| for r in rows: | |
| m = r.get("messages", []) | |
| if [x.get("role") for x in m] != ["system", "user", "assistant"]: | |
| bad_roles += 1 | |
| continue | |
| user, asst = m[1]["content"], m[2]["content"] | |
| claims = asst.split("\n\n") | |
| ulens.append(len(user)) | |
| alens.append(len(asst)) | |
| ncl.append(len(claims)) | |
| if not claims[0].lstrip().startswith("1."): | |
| bad_first += 1 | |
| for i, c in enumerate(claims, 1): | |
| if not c.lstrip().startswith(f"{i}."): | |
| bad_first += 1 | |
| break | |
| # A dependent claim must reference a strictly earlier claim. | |
| for i, c in enumerate(claims, 1): | |
| for mm in re.finditer(r"\bclaims?\s+(\d+)", c): | |
| if int(mm.group(1)) >= i: | |
| fwd_ref += 1 | |
| break | |
| else: | |
| continue | |
| break | |
| # Leakage: share of claim-1 8-grams already present in the prompt. | |
| cg = ngrams(claims[0]) | |
| if len(cg) >= 20: | |
| leaks.append(len(ngrams(user) & cg) / len(cg)) | |
| print(f" role sequence errors: {bad_roles}") | |
| print(f" claim numbering errors: {bad_first}") | |
| print(f" invalid forward references: {fwd_ref} ({fwd_ref/max(len(rows),1):.1%})") | |
| if leaks: | |
| print(f" claim-1 8-gram leakage: median {statistics.median(leaks):.1%} " | |
| f"mean {statistics.mean(leaks):.1%} >=25%: {sum(1 for x in leaks if x>=.25)/len(leaks):.1%}") | |
| print(f" claims per example: min {min(ncl)} median {int(statistics.median(ncl))} max {max(ncl)}") | |
| print(f" prompt chars: median {int(statistics.median(ulens))} max {max(ulens)}") | |
| print(f" target chars: median {int(statistics.median(alens))} max {max(alens)}") | |
| approx = [(u + a) / 4 for u, a in zip(ulens, alens)] | |
| print(f" approx tokens/example: median {int(statistics.median(approx))} " | |
| f"p99 {int(sorted(approx)[int(len(approx)*.99)])} max {int(max(approx))}") | |
| return {"n": len(rows), "errors": bad_roles + bad_first} | |
| def main(): | |
| if len(sys.argv) < 2: | |
| sys.exit(__doc__) | |
| d = sys.argv[1] | |
| total_err = 0 | |
| for name in ("train.jsonl", "val.jsonl"): | |
| p = os.path.join(d, name) | |
| if not os.path.exists(p): | |
| print(f"missing {p}") | |
| continue | |
| res = check(p) | |
| if res: | |
| total_err += res["errors"] | |
| # No patent may appear in both splits. | |
| man = os.path.join(d, "manifest.jsonl") | |
| if os.path.exists(man): | |
| splits = {} | |
| for line in open(man): | |
| r = json.loads(line) | |
| splits.setdefault(r["patent_number"], set()).add(r["split"]) | |
| overlap = [p for p, s in splits.items() if len(s) > 1] | |
| print(f"\ntrain/val patent overlap: {len(overlap)}") | |
| total_err += len(overlap) | |
| print(f"\n{'PASS' if total_err == 0 else 'ISSUES: ' + str(total_err)}") | |
| if __name__ == "__main__": | |
| main() | |