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

1"""One ingest worker process per GPU, each over its own slice of the corpus.""" 

2 

3from __future__ import annotations 

4 

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 

16 

17from rich.progress import ( 

18 BarColumn, 

19 MofNCompleteColumn, 

20 Progress, 

21 SpinnerColumn, 

22 TimeElapsedColumn, 

23) 

24 

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 

38 

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 

44 

45 from lilbee.core.config.model import Config 

46 from lilbee.runtime.cancellation import CancelSignal 

47 

48log = logging.getLogger(__name__) 

49 

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" 

55 

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 

59 

60# Under two workers there is nothing to fan out to. 

61_MIN_FANOUT_WORKERS = 2 

62 

63# How often a worker reports its counters to the parent. 

64_REPORT_INTERVAL_S = 0.25 

65 

66# How long the parent sleeps between drains of the worker message queue. 

67_DRAIN_INTERVAL_S = 0.1 

68 

69# Grace for the queue's feeder thread to flush a dead worker's last messages. 

70_FINAL_DRAIN_S = 1.0 

71 

72# How long a worker gets to exit on its own before it is killed. 

73_WORKER_EXIT_GRACE_S = 30.0 

74 

75# Where a worker's console output lands, under its own data root. 

76WORKER_LOG_NAME = "sync.log" 

77 

78 

79@dataclass(frozen=True) 

80class ShardSpec: 

81 """One worker's slice, its card, and the private state it owns.""" 

82 

83 shard: ShardId 

84 device: int 

85 config: Config 

86 engine_dir: Path 

87 cpu_share: int 

88 visible_devices: dict[str, str] 

89 

90 

91@dataclass(frozen=True) 

92class ShardOptions: 

93 """What every worker of one fan-out is told about the run it belongs to.""" 

94 

95 parent_pid: int 

96 force_rebuild: bool = False 

97 

98 

99@dataclass(frozen=True) 

100class ShardProgress: 

101 """A worker's counters as it works.""" 

102 

103 kind: Literal["progress"] 

104 index: int 

105 done: int 

106 planned: int 

107 file: str 

108 status: BatchStatus 

109 

110 

111@dataclass(frozen=True) 

112class ShardDone: 

113 """A worker's verdict; *error* set means it produced no usable shard.""" 

114 

115 kind: Literal["done"] 

116 index: int 

117 result: SyncResult | None 

118 error: str | None 

119 

120 

121ShardMessage = ShardProgress | ShardDone 

122 

123 

124def resolve_process_count(devices: int) -> int: 

125 """Ingest worker processes for this run; 1 keeps ingest in this process. 

126 

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 

136 

137 

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 

143 

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) 

152 

153 

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 

157 

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 ] 

174 

175 

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. 

178 

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 ) 

191 

192 

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) 

200 

201 

202def _redirect_output(path: Path) -> None: 

203 """Send this process's console output to *path*. 

204 

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()) 

213 

214 

215class _ShardReporter: 

216 """Throttled relay of a worker's counters onto the parent's queue. 

217 

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 """ 

223 

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 

230 

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) 

241 

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) 

245 

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 ) 

257 

258 

259class _Aggregate: 

260 """Every worker's latest counters, as one set of totals.""" 

261 

262 def __init__(self, on_progress: DetailedProgressCallback) -> None: 

263 self._latest: dict[int, ShardProgress] = {} 

264 self._on_progress = on_progress 

265 

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 

278 

279 

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 

287 

288 

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 ) 

299 

300 

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 

329 

330 

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. 

337 

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 

354 

355 

356def _stop_workers(workers: Sequence[BaseProcess], stop: Event) -> None: 

357 """Ask every live worker to stop, then wait for it, then insist. 

358 

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() 

371 

372 

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)] 

403 

404 

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 ) 

420 

421 

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 

430 

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)))