Coverage for src/lilbee/data/ingest/fanout.py: 100%
199 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-09-28 17:20 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-09-28 17:20 +0000
1"""One ingest worker process per GPU, each over its own slice of the corpus."""
3from __future__ import annotations
5import asyncio
6import contextlib
7import logging
8import multiprocessing
9import os
10import queue
11import sys
12import time
13from dataclasses import dataclass
14from pathlib import Path
15from typing import TYPE_CHECKING, Literal
17from rich.progress import (
18 BarColumn,
19 MofNCompleteColumn,
20 Progress,
21 SpinnerColumn,
22 TimeElapsedColumn,
23)
25from lilbee.core.config import active_config
26from lilbee.data.ingest.errors import error_reason
27from lilbee.data.types import ShardId, SyncResult
28from lilbee.runtime.cpu import available_cpu_count, cpu_quota
29from lilbee.runtime.engine_lock import ENGINE_DIR_ENV
30from lilbee.runtime.progress import (
31 BatchProgressEvent,
32 BatchStatus,
33 DetailedProgressCallback,
34 EventType,
35 ProgressEvent,
36)
37from lilbee.runtime.progress.columns import literal_text_column
39if TYPE_CHECKING:
40 from collections.abc import Sequence
41 from multiprocessing.process import BaseProcess
42 from multiprocessing.queues import Queue
43 from multiprocessing.synchronize import Event
45 from lilbee.core.config.model import Config
46 from lilbee.runtime.cancellation import CancelSignal
48log = logging.getLogger(__name__)
50# Per-worker state (store, engine slots, log) under the parent data root; the
51# skip records stay the corpus's, at the parent data root itself.
52SHARDS_DIRNAME = "shards"
53_DATA_ROOT_ENV = "LILBEE_DATA"
54_CPU_QUOTA_ENV = "LILBEE_CPU_QUOTA"
56# Below this many files on disk a fan-out costs more than it saves: every worker
57# pays a fresh interpreter, its own engine and a store of its own.
58_MIN_FILES_FOR_FANOUT = 2000
60# Under two workers there is nothing to fan out to.
61_MIN_FANOUT_WORKERS = 2
63# How often a worker reports its counters to the parent.
64_REPORT_INTERVAL_S = 0.25
66# How long the parent sleeps between drains of the worker message queue.
67_DRAIN_INTERVAL_S = 0.1
69# Grace for the queue's feeder thread to flush a dead worker's last messages.
70_FINAL_DRAIN_S = 1.0
72# How long a worker gets to exit on its own before it is killed.
73_WORKER_EXIT_GRACE_S = 30.0
75# Where a worker's console output lands, under its own data root.
76WORKER_LOG_NAME = "sync.log"
79@dataclass(frozen=True)
80class ShardSpec:
81 """One worker's slice, its card, and the private state it owns."""
83 shard: ShardId
84 device: int
85 config: Config
86 engine_dir: Path
87 cpu_share: int
88 visible_devices: dict[str, str]
91@dataclass(frozen=True)
92class ShardOptions:
93 """What every worker of one fan-out is told about the run it belongs to."""
95 parent_pid: int
96 force_rebuild: bool = False
99@dataclass(frozen=True)
100class ShardProgress:
101 """A worker's counters as it works."""
103 kind: Literal["progress"]
104 index: int
105 done: int
106 planned: int
107 file: str
108 status: BatchStatus
111@dataclass(frozen=True)
112class ShardDone:
113 """A worker's verdict; *error* set means it produced no usable shard."""
115 kind: Literal["done"]
116 index: int
117 result: SyncResult | None
118 error: str | None
121ShardMessage = ShardProgress | ShardDone
124def resolve_process_count(devices: int) -> int:
125 """Ingest worker processes for this run; 1 keeps ingest in this process.
127 Auto (``ingest_processes = 0``) is one worker per visible card. An explicit
128 count is honored past the card count, since two workers on one card is a
129 legitimate configuration; they share that card's engine slot rather than
130 putting a second fleet on it.
131 """
132 configured = active_config().ingest_processes
133 if configured:
134 return max(1, configured)
135 return devices
138def plan_fanout() -> list[ShardSpec]:
139 """The workers for this sync, empty when it runs in this process."""
140 from lilbee.data.ingest.discovery import corpus_has_at_least
141 from lilbee.providers.fleet.gpu_env import apply_fleet_gpu_env
142 from lilbee.providers.fleet.replicas import gpu_device_count
144 # Applied before the cards are counted, so a gpu_devices pin is the space the
145 # workers are dealt in: without it they would be dealt cards the pin excludes.
146 apply_fleet_gpu_env()
147 devices = gpu_device_count()
148 processes = resolve_process_count(devices)
149 if processes < _MIN_FANOUT_WORKERS or not corpus_has_at_least(_MIN_FILES_FOR_FANOUT):
150 return []
151 return shard_specs(active_config(), processes, devices)
154def shard_specs(config: Config, processes: int, devices: int) -> list[ShardSpec]:
155 """One spec per worker, dividing the corpus, the cards and the CPU pools."""
156 from lilbee.providers.fleet.gpu_env import shard_visible_devices
158 cpu_share = max(1, cpu_quota() // processes)
159 plan_share = max(1, available_cpu_count() // processes)
160 root = config.data_root / SHARDS_DIRNAME
161 return [
162 ShardSpec(
163 shard=ShardId(index=index, count=processes, records_root=config.data_root),
164 device=index % devices,
165 config=_shard_config(config, root / f"w{index}", plan_share, processes),
166 # Keyed by card, not by worker: workers sharing a card share one
167 # fleet, workers on different cards never see each other's.
168 engine_dir=root / f"gpu{index % devices}" / "engine",
169 cpu_share=cpu_share,
170 visible_devices=shard_visible_devices(index % devices),
171 )
172 for index in range(processes)
173 ]
176def _shard_config(config: Config, root: Path, plan_share: int, processes: int) -> Config:
177 """*config* with a private data root and this worker's share of the CPU pools.
179 ``documents_dir`` and ``linked_roots`` are inherited: every worker reads the
180 one shared corpus and only its own state is private.
181 """
182 threads = config.extraction_threads
183 return config.model_copy(
184 update={
185 "data_root": root,
186 "lancedb_dir": root / "data" / "lancedb",
187 "ingest_workers": plan_share,
188 "extraction_threads": max(1, threads // processes) if threads else 0,
189 }
190 )
193def _apply_shard_env(spec: ShardSpec) -> None:
194 """Pin this process to the worker's card, engine slot, CPU share and log."""
195 os.environ.update(spec.visible_devices)
196 os.environ[ENGINE_DIR_ENV] = str(spec.engine_dir)
197 os.environ[_DATA_ROOT_ENV] = str(spec.config.data_root)
198 os.environ[_CPU_QUOTA_ENV] = str(spec.cpu_share)
199 _redirect_output(spec.config.data_root / WORKER_LOG_NAME)
202def _redirect_output(path: Path) -> None:
203 """Send this process's console output to *path*.
205 At the file descriptor, so the engine this worker spawns follows it: N
206 workers logging onto the parent's terminal is the pile of log files the one
207 aggregated bar exists to replace.
208 """
209 path.parent.mkdir(parents=True, exist_ok=True)
210 with path.open("ab", buffering=0) as handle:
211 os.dup2(handle.fileno(), sys.stdout.fileno())
212 os.dup2(handle.fileno(), sys.stderr.fileno())
215class _ShardReporter:
216 """Throttled relay of a worker's counters onto the parent's queue.
218 The counters are the pipeline's own: how much of this worker's slice is done
219 and how big that slice is. Counting files here instead would only re-derive
220 the first, and the per-file events carry no slice size -- FILE_START's total
221 is the plan so far, which grows all run.
222 """
224 def __init__(self, index: int, messages: Queue[ShardMessage]) -> None:
225 self._index = index
226 self._messages = messages
227 self._done = 0
228 self._planned = 0
229 self._last_sent = 0.0
231 def __call__(self, event_type: EventType, data: ProgressEvent) -> None:
232 if event_type is not EventType.BATCH_PROGRESS or not isinstance(data, BatchProgressEvent):
233 return
234 self._done = data.current
235 self._planned = data.total
236 now = time.monotonic()
237 if now - self._last_sent < _REPORT_INTERVAL_S:
238 return
239 self._last_sent = now
240 self._send(data.file, data.status)
242 def flush(self) -> None:
243 """Send the final counters past the throttle, so the bar lands on its total."""
244 self._send("", BatchStatus.INGESTED)
246 def _send(self, file: str, status: BatchStatus) -> None:
247 self._messages.put(
248 ShardProgress(
249 kind="progress",
250 index=self._index,
251 done=self._done,
252 planned=self._planned,
253 file=file,
254 status=status,
255 )
256 )
259class _Aggregate:
260 """Every worker's latest counters, as one set of totals."""
262 def __init__(self, on_progress: DetailedProgressCallback) -> None:
263 self._latest: dict[int, ShardProgress] = {}
264 self._on_progress = on_progress
266 def update(self, message: ShardProgress) -> tuple[int, int]:
267 """Record *message* and return the corpus-wide (done, planned)."""
268 self._latest[message.index] = message
269 done = sum(p.done for p in self._latest.values())
270 planned = sum(p.planned for p in self._latest.values())
271 self._on_progress(
272 EventType.BATCH_PROGRESS,
273 BatchProgressEvent(
274 file=message.file, status=message.status, current=done, total=planned
275 ),
276 )
277 return done, planned
280def _drain(messages: Queue[ShardMessage]) -> list[ShardMessage]:
281 """Every message queued right now, without blocking."""
282 drained: list[ShardMessage] = []
283 with contextlib.suppress(queue.Empty):
284 while True:
285 drained.append(messages.get_nowait())
286 return drained
289def _shard_progress_bar(quiet: bool) -> Progress:
290 """The one bar a fan-out reports on, disabled when the caller wants no output."""
291 return Progress(
292 SpinnerColumn(),
293 literal_text_column("{task.description}", style="progress.description"),
294 BarColumn(),
295 MofNCompleteColumn(),
296 TimeElapsedColumn(),
297 disable=quiet,
298 )
301async def _supervise(
302 workers: Sequence[BaseProcess],
303 messages: Queue[ShardMessage],
304 stop: Event,
305 *,
306 quiet: bool,
307 on_progress: DetailedProgressCallback,
308 cancel: CancelSignal | None,
309) -> dict[int, ShardDone]:
310 """Drain worker messages until every worker has reported, keeping one bar current."""
311 verdicts: dict[int, ShardDone] = {}
312 aggregate = _Aggregate(on_progress)
313 with _shard_progress_bar(quiet) as progress:
314 task = progress.add_task(f"Ingesting on {len(workers)} workers", total=None)
315 while len(verdicts) < len(workers):
316 for message in _drain(messages):
317 if message.kind == "done":
318 verdicts[message.index] = message
319 else:
320 done, planned = aggregate.update(message)
321 progress.update(task, completed=done, total=planned or None)
322 if cancel is not None and cancel.is_set():
323 stop.set()
324 if not any(worker.is_alive() for worker in workers):
325 verdicts.update(_final_verdicts(workers, messages, verdicts))
326 break
327 await asyncio.sleep(_DRAIN_INTERVAL_S)
328 return verdicts
331def _final_verdicts(
332 workers: Sequence[BaseProcess],
333 messages: Queue[ShardMessage],
334 verdicts: dict[int, ShardDone],
335) -> dict[int, ShardDone]:
336 """Verdicts still in flight once every worker has exited, plus one per silent death.
338 A worker the kernel killed (out of memory is the usual reason) reports
339 nothing, so its shard is recorded as failed rather than silently missing from
340 the merge.
341 """
342 time.sleep(_FINAL_DRAIN_S)
343 late = {m.index: m for m in _drain(messages) if m.kind == "done"}
344 for index, worker in enumerate(workers):
345 if index in verdicts or index in late:
346 continue
347 late[index] = ShardDone(
348 kind="done",
349 index=index,
350 result=None,
351 error=f"worker exited with code {worker.exitcode} before reporting",
352 )
353 return late
356def _stop_workers(workers: Sequence[BaseProcess], stop: Event) -> None:
357 """Ask every live worker to stop, then wait for it, then insist.
359 A worker owns a GPU fleet, and its teardown can outlast a TERM; a plain join
360 would hang the sync behind it instead of returning a result it already has.
361 """
362 stop.set()
363 for worker in workers:
364 if worker.is_alive():
365 worker.terminate()
366 worker.join(_WORKER_EXIT_GRACE_S)
367 if worker.is_alive():
368 log.warning("Ingest worker %s did not exit; killing it", worker.name)
369 worker.kill()
370 worker.join()
373async def run_workers(
374 specs: list[ShardSpec],
375 *,
376 options: ShardOptions,
377 quiet: bool,
378 on_progress: DetailedProgressCallback,
379 cancel: CancelSignal | None,
380) -> list[ShardDone]:
381 """Run every worker to completion and return their verdicts, in shard order."""
382 context = multiprocessing.get_context("spawn")
383 messages: Queue[ShardMessage] = context.Queue()
384 stop = context.Event()
385 workers = [
386 context.Process(
387 target=run_shard,
388 args=(spec, options, messages, stop),
389 name=f"lilbee-shard-{spec.shard.index}",
390 )
391 for spec in specs
392 ]
393 log.warning("Ingesting across %d worker processes, one per GPU", len(workers))
394 for worker in workers:
395 worker.start()
396 try:
397 verdicts = await _supervise(
398 workers, messages, stop, quiet=quiet, on_progress=on_progress, cancel=cancel
399 )
400 finally:
401 _stop_workers(workers, stop)
402 return [verdicts[index] for index in sorted(verdicts)]
405def aggregate_results(verdicts: list[ShardDone]) -> SyncResult:
406 """The one result a fan-out reports, unioned from every worker's."""
407 results = [verdict.result for verdict in verdicts if verdict.result is not None]
408 return SyncResult(
409 added=[name for r in results for name in r.added],
410 updated=[name for r in results for name in r.updated],
411 relocated=[name for r in results for name in r.relocated],
412 failed=[name for r in results for name in r.failed],
413 skipped=[name for r in results for name in r.skipped],
414 skipped_ocr={name: ocr for r in results for name, ocr in r.skipped_ocr.items()},
415 removed=[name for r in results for name in r.removed],
416 held_out=[held for r in results for held in r.held_out],
417 unchanged=sum(r.unchanged for r in results),
418 truncated=sum(r.truncated for r in results),
419 )
422def run_shard(
423 spec: ShardSpec, options: ShardOptions, messages: Queue[ShardMessage], stop: Event
424) -> None:
425 """Ingest this worker's slice in a fresh process, reporting onto *messages*."""
426 from lilbee.app.services import build_services, services_scope
427 from lilbee.core.config.context import config_scope
428 from lilbee.data.ingest.pipeline import sync
429 from lilbee.providers.fleet.child_guard import bind_lifetime_to_parent
431 bind_lifetime_to_parent(options.parent_pid)
432 _apply_shard_env(spec)
433 index = spec.shard.index
434 reporter = _ShardReporter(index, messages)
435 try:
436 with config_scope(spec.config), services_scope(build_services(spec.config)):
437 result = asyncio.run(
438 sync(
439 force_rebuild=options.force_rebuild,
440 quiet=True,
441 on_progress=reporter,
442 cancel=stop,
443 shard=spec.shard,
444 )
445 )
446 reporter.flush()
447 messages.put(ShardDone(kind="done", index=index, result=result, error=None))
448 except (Exception, asyncio.CancelledError) as exc:
449 messages.put(ShardDone(kind="done", index=index, result=None, error=error_reason(exc)))