AngeloUNIMI's picture
Document Exam Trainer v5.0.3: written and oral answers
3035446 verified
Raw History Blame Contribute Delete
12.7 kB
from __future__ import annotations
from dataclasses import asdict
import hashlib
import json
from pathlib import Path
import secrets
import shutil
import uuid
import numpy as np
from .config import SETTINGS
from .ingestion import IngestionError, stage_documents, parse_documents
from .retrieval import CorpusIndex, Embedder, pack_context
from .schemas import Chunk, Concept, DocumentInfo, Evidence, ParsedCorpus, Rubric, Topic
from .sessions import Session, SessionStore, SessionError
from .validation import OutputError, evidence_is_valid, validate_assessment, validate_draft
from .workers import CPUWorker
class TrainerService:
def __init__(self, backend, store: SessionStore | None = None, demo: bool = SETTINGS.demo):
self.backend = backend
self.store = store or SessionStore()
self.documents_worker = CPUWorker("documents")
self.speech_worker = CPUWorker("speech")
self.embedder = Embedder(self.documents_worker, demo=demo)
self.demo = demo
def close(self):
self.documents_worker.close()
self.speech_worker.close()
if hasattr(self.backend, "close"):
self.backend.close()
def build(self, session: Session, inputs: list[tuple[Path, str]], progress=None) -> ParsedCorpus:
stage = session.path / f"build-{secrets.token_hex(8)}"
stage.mkdir()
try:
staged, warnings = stage_documents(inputs, stage / "raw", self.store.settings)
if progress:
progress("Extracting document text and finding sections...")
if self.demo:
parsed = parse_documents(staged, warnings, self.store.settings)
else:
data = self.documents_worker.request("parse", {"files": [(str(p), n, r) for p, n, r in staged], "warnings": warnings},
timeout=300, progress=progress)
parsed = ParsedCorpus(documents=[DocumentInfo(**d) for d in data["documents"]],
chunks=[Chunk(**c) for c in data["chunks"]], topics=[Topic(**t) for t in data["topics"]],
warnings=data["warnings"], fingerprint=data["fingerprint"])
# Originals are not needed by inference and are not retained internally.
shutil.rmtree(stage / "raw", ignore_errors=True)
if session.index is not None and session.index.parsed.fingerprint == parsed.fingerprint:
shutil.rmtree(stage, ignore_errors=True)
if progress:
progress("Reusing this session's unchanged corpus index.")
session.rubric = session.last_result = None
session.last_key = ""
return session.index.parsed
vectors = self.embedder.encode([c.text for c in parsed.chunks], progress=progress)
index = CorpusIndex(parsed, vectors)
# These local files are never uploaded to a Hub repository or exposed via Gradio.
(stage / "manifest.json").write_text(json.dumps(asdict(parsed), ensure_ascii=False), encoding="utf-8")
np.save(stage / "embeddings.npy", vectors, allow_pickle=False)
for old in session.path.glob("build-*"):
if old != stage:
shutil.rmtree(old, ignore_errors=True)
session.index = index
session.version = uuid.uuid4().hex
session.rubric = session.last_result = None
session.last_key = ""
session.history.clear()
session.course_id = session.transcript = session.transcript_question = ""
return parsed
except Exception:
shutil.rmtree(stage, ignore_errors=True)
raise
@staticmethod
def require_index(session: Session) -> CorpusIndex:
if session.index is None:
raise SessionError("Upload and process primary documents first.")
return session.index
def question(self, session: Session, topic_id: str, focus: str, style: str) -> Rubric:
index = self.require_index(session)
if len(focus) > 400:
raise ValueError("Keep the optional focus to 400 characters or fewer.")
choices = {"General description", "Definitions and properties", "Procedure / operation", "Limitations and extensions"}
if style not in choices:
raise ValueError("Invalid question style.")
topic, chunks = index.question_context(topic_id, focus, self.embedder)
_, admitted = pack_context(chunks, 10000)
if not admitted:
raise ValueError("No usable primary passages are available for that topic.")
payload = {"topic_label": topic.label, "focus": focus, "style": style,
"avoid_repeating": session.history[-4:],
"sources": [c.as_dict() for c in admitted]}
raw = self.backend.call("question", payload)
rubric = validate_draft(raw, admitted)
rubric.corpus_version = session.version
rubric.topic_id = topic.id
rubric.context_ids = [c.id for c in admitted]
session.rubric = rubric
session.transcript = session.transcript_question = ""
session.last_result = None
session.last_key = ""
session.history.append(rubric.question)
session.history = session.history[-12:]
return rubric
def evaluate(self, session: Session, answer: str, supporting_explanations: bool = False, progress=None):
index = self.require_index(session)
rubric = session.rubric
if rubric is None or rubric.corpus_version != session.version:
raise SessionError("Generate a question for the current corpus first.")
answer = answer.strip()
if not answer:
raise ValueError("Type an answer or transcribe a recording first.")
if len(answer) > self.store.settings.max_answer_chars:
raise ValueError(f"Use no more than {self.store.settings.max_answer_chars} characters per answer.")
key = hashlib.sha256((rubric.model_dump_json() + answer + str(supporting_explanations)).encode()).hexdigest()
if key == session.last_key and session.last_result is not None:
if progress:
progress("Reusing the result for this unchanged answer (no extra GPU call).")
return session.last_result
if progress:
progress("Retrieving primary evidence for the frozen question rubric...")
exact_ids = list(dict.fromkeys(e.source_id for c in rubric.concepts for e in c.evidence))
required = [index.by_id[x] for x in exact_ids]
query = rubric.question # NEVER retrieve primary expectations from the learner's answer.
vector = self.embedder.encode([query], progress=progress)[0]
topic_ids = set(index.topics[rubric.topic_id].source_ids)
hits = index.hits(vector, role="primary", ids=topic_ids, query=query, top_k=4)
_, context = pack_context(required + hits, 12500)
if not set(exact_ids) <= {c.id for c in context}:
raise OutputError("The rubric evidence exceeds the context budget. Generate a narrower question.")
if progress:
progress("Waiting for inference / checking semantic coverage and omissions...")
# Deliberately NO supporting passages here. They cannot introduce grading requirements.
raw = self.backend.call("grade", {"question": rubric.question, "answer": answer,
"rubric": rubric.model_dump(), "primary_sources": [c.as_dict() for c in context]})
result = validate_assessment(raw, rubric, answer)
gaps = [c for c in result.checks if c.status in ("missing", "partial", "incorrect")]
if gaps and any(c.role == "supporting" for c in index.chunks):
if progress:
progress("Finding optional supporting references for confirmed gaps...")
by_id = {c.id: c for c in rubric.concepts}
queries = [by_id[c.concept_id].name + ": " + by_id[c.concept_id].description for c in gaps]
embeddings = self.embedder.encode(queries, progress=progress)
selected_by_concept = {}
for check, query, vector in zip(gaps, queries, embeddings):
refs = index.hits(vector, role="supporting", query=query, top_k=1)
# A search hit is a suggested reference, not a new requirement or
# proof. Explanations have an additional evidence check below.
selected_by_concept[check.concept_id] = refs
result.supporting_ids[check.concept_id] = [c.id for c in refs]
if supporting_explanations:
if progress:
progress("Preparing optional supporting explanations (additional GPU request)...")
try:
subset = gaps[:3]
data = {"gaps": [{"concept": by_id[g.concept_id].model_dump(),
"missing_detail": g.missing_detail,
"primary_evidence": [index.by_id[e.source_id].as_dict() for e in by_id[g.concept_id].evidence],
"supporting": [c.as_dict() for c in selected_by_concept[g.concept_id]]} for g in subset]}
notes = self.backend.call("explain", data).get("notes", [])
for note in notes:
cid = note.get("concept_id")
if cid not in selected_by_concept or not isinstance(note.get("explanation"), str):
continue
allowed = {c.id: c for c in selected_by_concept[cid]}
evidence = [Evidence.model_validate(e) for e in note.get("evidence", [])]
if evidence and all(evidence_is_valid(e, allowed, "supporting") for e in evidence):
result.supporting_hints[cid] = note["explanation"][:1400]
result.supporting_ids[cid] = [e.source_id for e in evidence]
except Exception:
# An optional hint must never erase completed primary-only feedback.
result.warnings.append("Optional supporting explanations were unavailable; the primary-source answer check is complete.")
session.last_key, session.last_result = key, result
return result
def transcribe(self, session: Session, audio_path: Path, language: str, vocabulary: str, progress=None) -> dict:
if audio_path.stat().st_size > self.store.settings.max_audio_bytes:
raise ValueError("The recording exceeds the configured upload limit.")
work = session.path / f"audio-{secrets.token_hex(12)}{audio_path.suffix[:10]}"
shutil.copyfile(audio_path, work)
try:
return self.speech_worker.request("transcribe", {"path": str(work), "language": language,
"vocabulary": vocabulary[:500]},
timeout=600, progress=progress)
finally:
work.unlink(missing_ok=True)
def followup(self, session: Session) -> Rubric:
self.require_index(session)
if session.last_result is None or session.rubric is None:
raise ValueError("Evaluate an answer before requesting a follow-up.")
by_id = {c.id: c for c in session.rubric.concepts}
rank = {"essential": 0, "important": 1, "minor": 2}
gaps = [c for c in session.last_result.checks if c.status in ("missing", "partial", "incorrect")]
if not gaps:
raise ValueError("No confirmed gaps to follow up. Generate a new question instead.")
gap = sorted(gaps, key=lambda g: rank[by_id[g.concept_id].importance])[0]
concept = by_id[gap.concept_id].model_copy(deep=True)
concept.aspect = concept.name
concept.relevance = "This follow-up asks specifically for the previously identified primary concept."
old = session.rubric
session.rubric = Rubric(question=f"Describe {concept.name} as presented in the primary material.",
asked_aspects=[concept.name], concepts=[concept], corpus_version=session.version,
topic_id=old.topic_id, context_ids=list(dict.fromkeys(e.source_id for e in concept.evidence)))
session.last_result, session.last_key = None, ""
session.transcript = session.transcript_question = ""
session.history.append(session.rubric.question)
return session.rubric