ny-solar-siting-rag / scripts /run_own_eval.py
leandersen's picture
Main Project Commit
79ecbd3 verified
Raw
History Blame Contribute Delete
7.54 kB
"""Run the 75-item project eval set through the RAG pipeline (--mode rag)
or the bare generator (--mode norag). Preambles and scoring match the
Check-In 4 notebook. PSL s141 is excluded from the corpus in code
(transmission-only; see the sources.csv notes), so the Qwen3-4B arms were
rerun with this script rather than reusing the Check-In 4 numbers.
Usage: python run_own_eval.py --mode rag --k 5
"""
import argparse, os, json
from bench_common import (Retriever, generate, save_outputs, cleanup,
GEN_MODEL, MODEL_TAG)
DATA_DIR = os.environ.get("DATA_DIR", "RAG Data")
SCRATCH = os.environ.get("SCRATCH", ".") # results/ and emb/ land here
EVAL_PATH = os.environ.get("EVAL_PATH", "eval_set_v1.jsonl")
# Verbatim from the Check-In 4 notebook (cell: own-eval definitions).
SYSTEM_PREAMBLE = """You are a regulatory assistant for early-stage solar development in New York State. Answer the developer's question using ONLY the context passages below. Cite the supporting document and section for every claim in the form [Document, Section]. If the context does not contain the answer, reply: "The provided sources do not answer this question" and briefly say what source likely would."""
NORAG_PREAMBLE = """You are a regulatory assistant for early-stage solar development in New York State. Answer the developer's question from your own knowledge. Cite the governing document and section for every claim in the form [Document, Section]. If you do not know, say so."""
DECLINE_MARKERS = ['do not answer', 'does not answer', 'cannot answer',
'not contain', 'do not know', 'no information']
# Docs kept out of the index at notebook level (D-09) plus the Q-11 fix.
EXCLUDE_FILES = [
'NYSERDA Solar Guidebook for Local Governments (effective November 2025).pdf',
]
EXCLUDE_SECTION_FILES = ['psl_a8_s141.txt'] # Q-11: transmission twin of s143
def load_corpus(data_dir=DATA_DIR):
"""Load sources.csv rows with ingest == true, minus EXCLUDE_FILES."""
import pandas as pd
from pypdf import PdfReader
manifest = pd.read_csv(os.path.join(data_dir, 'sources.csv'))
m = manifest[manifest['ingest'].astype(str).str.lower() == 'true']
m = m[~m['filename'].isin(EXCLUDE_FILES)]
docs = []
for _, row in m.iterrows():
path = os.path.join(data_dir, row['filename'])
if row['filename'].endswith('/'):
for fname in sorted(os.listdir(path)):
if not fname.endswith('.txt'):
continue
if fname in EXCLUDE_SECTION_FILES:
print(f"NOTE: skipping {fname} (Q-11: transmission-only "
f"section; excluded in code)")
continue
with open(os.path.join(path, fname), encoding='utf-8') as fh:
docs.append({'text': fh.read(), 'source': fname,
'family': row['document_family']})
elif path.endswith('.md') or path.endswith('.txt'):
with open(path, encoding='utf-8') as fh:
docs.append({'text': fh.read(), 'source': row['filename'],
'family': row['document_family']})
else:
reader = PdfReader(path)
text = '\n'.join(p.extract_text() or '' for p in reader.pages)
docs.append({'text': text, 'source': row['filename'],
'family': row['document_family']})
n_words = sum(len(d['text'].split()) for d in docs)
print(f"{len(docs)} documents, {n_words:,} words")
return docs
def chunk_corpus(docs):
from langchain_text_splitters import RecursiveCharacterTextSplitter
splitter = RecursiveCharacterTextSplitter(
chunk_size=1500, chunk_overlap=150,
separators=['\n## ', '\n\n', '\n', '. ', ' '])
chunks = []
for d in docs:
for c in splitter.split_text(d['text']):
chunks.append({'text': c, 'id': d['source'],
'family': d['family']})
print(f"{len(chunks)} chunks")
return chunks
def build_retriever(chunks):
# v2 cache name: chunk count changed when s141 was dropped, and the
# Retriever asserts on row count, so a stale cache fails loudly.
return Retriever(chunks, f"{SCRATCH}/emb/own_corpus_v2.npy")
def load_eval_set(path=EVAL_PATH):
with open(path, encoding='utf-8') as fh:
items = [json.loads(l) for l in fh if l.strip()]
from collections import Counter
print(f"{len(items)} items", Counter(i['stratum'] for i in items))
return items
def score_item(item, response):
resp = response.lower()
declined = any(m in resp for m in DECLINE_MARKERS)
cited_gold = any(tok.lower() in resp
for tok in item.get('gold_citation_keys', []))
if item['stratum'] == 'unanswerable':
return {'declined': declined, 'over_answer': not declined,
'over_decline': False, 'cited_gold': None}
return {'declined': declined, 'over_answer': False,
'over_decline': declined,
'cited_gold': cited_gold and not declined}
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--mode", choices=["rag", "norag"], required=True)
ap.add_argument("--k", type=int, default=5)
args = ap.parse_args()
print(f"generator: {GEN_MODEL} tag: {MODEL_TAG or '(none)'}")
items = load_eval_set()
retriever = None
if args.mode == "rag":
retriever = build_retriever(chunk_corpus(load_corpus()))
samples = []
for i, item in enumerate(items):
if retriever is not None:
hits = retriever.search(item['question'], k=args.k)
context = '\n\n'.join(
f'[{j+1}] ({h["id"]}): "{h["text"]}"'
for j, (h, _) in enumerate(hits))
block = f"Context:\n{context}\n\nQuestion: {item['question']}"
resp = generate(block, SYSTEM_PREAMBLE)
sources = [h['id'] for h, _ in hits]
else:
resp = generate(f"Question: {item['question']}", NORAG_PREAMBLE)
sources = []
row = {'id': item['id'], 'stratum': item['stratum'],
'family': item['family'], 'question': item['question'],
'response': resp, 'retrieved': sources}
row.update(score_item(item, resp))
if sources and item.get('gold_sources'):
row['retrieval_hit'] = any(s in item['gold_sources'] for s in sources)
samples.append(row)
print(f" [{i+1}/{len(items)}] {item['id']} "
f"{'DECLINED' if row['declined'] else 'answered'}")
unans = [s for s in samples if s['stratum'] == 'unanswerable']
ans = [s for s in samples if s['stratum'] != 'unanswerable']
cited = [s['cited_gold'] for s in ans]
metrics = {
'over_answer_rate': (sum(s['over_answer'] for s in unans) / len(unans))
if unans else None,
'over_decline_rate': (sum(s['over_decline'] for s in ans) / len(ans))
if ans else None,
'gold_citation_rate': (sum(cited) / len(cited)) if cited else None,
}
hits = [s['retrieval_hit'] for s in samples if 'retrieval_hit' in s]
if hits:
metrics['retrieval_hit_rate'] = sum(hits) / len(hits)
save_outputs(f"{SCRATCH}/results", "own", args.mode, metrics, samples,
extra={'k': args.k, 'corpus': 'v2 (s141 excluded, Q-11)',
'eval_set': EVAL_PATH})
cleanup()
if __name__ == "__main__":
main()