Livrare LOT 1 - Didi
This commit is contained in:
commit
5380c3fc63
990 changed files with 133308 additions and 0 deletions
258
ai_platform/modules/didi_brain/lint/runner.py
Normal file
258
ai_platform/modules/didi_brain/lint/runner.py
Normal file
|
|
@ -0,0 +1,258 @@
|
|||
"""Orchestrator for the Lint pass.
|
||||
|
||||
Top-level flow:
|
||||
|
||||
1. Pull every Type/Claim atom id from Atomic
|
||||
2. For each claim, ask Atomic for its nearest neighbors → candidate pairs
|
||||
3. Drop pairs we've already evaluated at the current prompt version
|
||||
4. Fetch the full body of each unique atom referenced by the new pairs
|
||||
5. Run Qwen NLI pair classification with bounded parallelism
|
||||
6. Record every verdict in the state file; save every PERIODIC_SAVE_EVERY
|
||||
pairs so a crash mid-run doesn't lose everything
|
||||
|
||||
The runner is safe to re-run at any time — the state file makes it fully
|
||||
idempotent, and we ALWAYS save on exit even on exception.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from shared.atomic_api import AtomicClient
|
||||
from shared.config import settings
|
||||
from shared.llm_client import LlmClient
|
||||
from shared.logging import get_logger
|
||||
from shared.taxonomy import TagResolver
|
||||
|
||||
# Cross-module reuse: the same function the extractor/brain_api use to
|
||||
# pull (claim_text, stance, parent_id) out of a Type/Claim atom body.
|
||||
from brain_api.services.mapping import parse_claim_atom_body
|
||||
|
||||
from lint._state import LintState, LintStats, PairVerdict, now_iso, pair_hash
|
||||
from lint.detector import PROMPT_VERSION, classify_pair
|
||||
from lint.pairs import CandidatePair, generate_candidates
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
# Bounded parallelism for NLI calls. We have two llama.cpp backends behind
|
||||
# the router; anything more than that just queues at the backend and eats
|
||||
# our per-call timeout. See brain_api/services/nli.py for the same reasoning.
|
||||
MAX_PARALLEL = 2
|
||||
|
||||
# Write the state file every N verdicts so a crash or Ctrl-C doesn't erase
|
||||
# the whole run. Saving is cheap (small JSON file).
|
||||
PERIODIC_SAVE_EVERY = 25
|
||||
|
||||
# Chunked atom fetch to avoid hammering /api/atoms/{id} with one huge gather.
|
||||
FETCH_CHUNK = 20
|
||||
|
||||
|
||||
async def _load_claim_atom_ids(
|
||||
atomic: AtomicClient,
|
||||
*,
|
||||
type_claim_id: str,
|
||||
limit: int | None,
|
||||
) -> list[str]:
|
||||
"""Page through /api/atoms?tag_id=<Type/Claim> and collect ids."""
|
||||
page_size = 100
|
||||
offset = 0
|
||||
out: list[str] = []
|
||||
while True:
|
||||
result = await atomic.list_atoms(
|
||||
limit=page_size, offset=offset, tag_id=type_claim_id
|
||||
)
|
||||
atoms = result.get("atoms") or (result if isinstance(result, list) else [])
|
||||
if not atoms:
|
||||
break
|
||||
for a in atoms:
|
||||
aid = a.get("id")
|
||||
if aid:
|
||||
out.append(aid)
|
||||
if limit and len(out) >= limit:
|
||||
return out
|
||||
if len(atoms) < page_size:
|
||||
break
|
||||
offset += page_size
|
||||
return out
|
||||
|
||||
|
||||
async def _fetch_atoms_bulk(
|
||||
atomic: AtomicClient, atom_ids: list[str], *, chunk: int = FETCH_CHUNK
|
||||
) -> dict[str, dict]:
|
||||
"""Fetch full atom bodies in parallel chunks. Missing atoms are dropped."""
|
||||
out: dict[str, dict] = {}
|
||||
total = len(atom_ids)
|
||||
for i in range(0, total, chunk):
|
||||
ids = atom_ids[i : i + chunk]
|
||||
results = await asyncio.gather(
|
||||
*[atomic.get_atom(a) for a in ids], return_exceptions=True
|
||||
)
|
||||
for a, r in zip(ids, results, strict=True):
|
||||
if isinstance(r, dict):
|
||||
out[a] = r
|
||||
if (i + chunk) % 200 == 0 or (i + chunk) >= total:
|
||||
log.info("fetch_progress", fetched=len(out), total=total)
|
||||
return out
|
||||
|
||||
|
||||
async def run_lint_pass(
|
||||
*,
|
||||
limit_atoms: int | None = None,
|
||||
force: bool = False,
|
||||
) -> LintStats:
|
||||
if not settings.atomic_token:
|
||||
raise RuntimeError("ATOMIC_TOKEN missing — can't talk to brain")
|
||||
|
||||
resolver = TagResolver()
|
||||
type_claim_id = resolver.require("Type/Claim")
|
||||
|
||||
state = LintState()
|
||||
state.set_prompt_version(PROMPT_VERSION)
|
||||
stats = LintStats()
|
||||
|
||||
async with AtomicClient() as atomic, LlmClient() as llm:
|
||||
# ------------------------------------------------------ 1. load atoms
|
||||
t0 = time.perf_counter()
|
||||
claim_ids = await _load_claim_atom_ids(
|
||||
atomic, type_claim_id=type_claim_id, limit=limit_atoms
|
||||
)
|
||||
stats.atoms_seen = len(claim_ids)
|
||||
log.info(
|
||||
"lint_atoms_loaded",
|
||||
count=stats.atoms_seen,
|
||||
elapsed_s=round(time.perf_counter() - t0, 1),
|
||||
)
|
||||
if not claim_ids:
|
||||
return stats
|
||||
|
||||
# ------------------------------------------ 2. candidate pair generation
|
||||
t0 = time.perf_counter()
|
||||
candidates = await generate_candidates(atomic, source_atom_ids=claim_ids)
|
||||
stats.candidates_generated = len(candidates)
|
||||
log.info(
|
||||
"lint_candidates_built",
|
||||
pairs=stats.candidates_generated,
|
||||
elapsed_s=round(time.perf_counter() - t0, 1),
|
||||
)
|
||||
if not candidates:
|
||||
return stats
|
||||
|
||||
# ----------------------------------- 3. filter out already-evaluated pairs
|
||||
to_eval: list[CandidatePair] = []
|
||||
for c in candidates:
|
||||
h = pair_hash(c.atom_a_id, c.atom_b_id)
|
||||
if not force and state.already_evaluated(h, PROMPT_VERSION):
|
||||
stats.pairs_skipped_cached += 1
|
||||
continue
|
||||
to_eval.append(c)
|
||||
log.info(
|
||||
"lint_filter_done",
|
||||
new_pairs=len(to_eval),
|
||||
cached=stats.pairs_skipped_cached,
|
||||
)
|
||||
|
||||
if not to_eval:
|
||||
state.save()
|
||||
return stats
|
||||
|
||||
# -------------------------------------- 4. pre-fetch full atom bodies
|
||||
needed_ids: set[str] = set()
|
||||
for c in to_eval:
|
||||
needed_ids.add(c.atom_a_id)
|
||||
needed_ids.add(c.atom_b_id)
|
||||
t0 = time.perf_counter()
|
||||
full_atoms = await _fetch_atoms_bulk(atomic, list(needed_ids))
|
||||
log.info(
|
||||
"atoms_fetched",
|
||||
count=len(full_atoms),
|
||||
needed=len(needed_ids),
|
||||
elapsed_s=round(time.perf_counter() - t0, 1),
|
||||
)
|
||||
|
||||
# Cache parsed claim bodies so each one is parsed once, not per-pair
|
||||
parsed_by_id: dict[str, tuple[str, str]] = {}
|
||||
for atom_id, full in full_atoms.items():
|
||||
content = full.get("content") or ""
|
||||
text, _stance, _parent_id = parse_claim_atom_body(content)
|
||||
parent_url = (full.get("source_url") or "").split("#", 1)[0]
|
||||
if text:
|
||||
parsed_by_id[atom_id] = (text, parent_url)
|
||||
|
||||
# -------------------------------------- 5. classify pairs (parallel)
|
||||
sem = asyncio.Semaphore(MAX_PARALLEL)
|
||||
total = len(to_eval)
|
||||
|
||||
# Mutable counters so we can log inside the coroutine
|
||||
progress = {"completed": 0}
|
||||
|
||||
async def _worker(pair: CandidatePair) -> PairVerdict | None:
|
||||
a = parsed_by_id.get(pair.atom_a_id)
|
||||
b = parsed_by_id.get(pair.atom_b_id)
|
||||
if not a or not b:
|
||||
return None
|
||||
text_a, url_a = a
|
||||
text_b, url_b = b
|
||||
async with sem:
|
||||
result = await classify_pair(llm, claim_a=text_a, claim_b=text_b)
|
||||
progress["completed"] += 1
|
||||
if progress["completed"] % 20 == 0:
|
||||
log.info(
|
||||
"lint_progress",
|
||||
done=progress["completed"],
|
||||
total=total,
|
||||
pct=round(progress["completed"] / total * 100, 1),
|
||||
)
|
||||
return PairVerdict(
|
||||
pair_hash=pair_hash(pair.atom_a_id, pair.atom_b_id),
|
||||
atom_a_id=pair.atom_a_id,
|
||||
atom_b_id=pair.atom_b_id,
|
||||
atom_a_url=url_a,
|
||||
atom_b_url=url_b,
|
||||
atom_a_claim=text_a,
|
||||
atom_b_claim=text_b,
|
||||
label=result.label,
|
||||
confidence=result.confidence,
|
||||
similarity=pair.similarity,
|
||||
detected_at=now_iso(),
|
||||
prompt_version=PROMPT_VERSION,
|
||||
error=result.error,
|
||||
)
|
||||
|
||||
t0 = time.perf_counter()
|
||||
tasks = [_worker(c) for c in to_eval]
|
||||
# Process results as they arrive so we can save periodically and
|
||||
# keep the ledger fresh even during long runs.
|
||||
try:
|
||||
for fut in asyncio.as_completed(tasks):
|
||||
verdict = await fut
|
||||
if verdict is None:
|
||||
stats.errors += 1
|
||||
continue
|
||||
state.upsert(verdict)
|
||||
stats.pairs_evaluated += 1
|
||||
if verdict.error is not None:
|
||||
stats.errors += 1
|
||||
elif verdict.label == "CONTRADICTORY":
|
||||
stats.contradictory += 1
|
||||
elif verdict.label == "EQUIVALENT":
|
||||
stats.equivalent += 1
|
||||
else:
|
||||
stats.incomparable += 1
|
||||
if stats.pairs_evaluated % PERIODIC_SAVE_EVERY == 0:
|
||||
state.save()
|
||||
finally:
|
||||
# ALWAYS save — even on Ctrl-C / exception, we keep what we got
|
||||
state.save()
|
||||
|
||||
log.info(
|
||||
"lint_classify_done",
|
||||
evaluated=stats.pairs_evaluated,
|
||||
contradictory=stats.contradictory,
|
||||
equivalent=stats.equivalent,
|
||||
incomparable=stats.incomparable,
|
||||
errors=stats.errors,
|
||||
elapsed_s=round(time.perf_counter() - t0, 1),
|
||||
)
|
||||
|
||||
return stats
|
||||
Loading…
Add table
Add a link
Reference in a new issue