Source code for nlp_shap.pipeline.runner

"""Public explain pipeline entry point."""

import asyncio
from pathlib import Path

from ..domain.conversation import ConversationSnapshot
from ..masking.partitions import TokenPartitioner
from ..plugins import PluginGroup, PluginRegistry, register_builtin_plugins
from ..plugins.backends import instantiate_backend
from ..runtime.progress import CoalitionProgress
from ..runtime.telemetry import NullObservabilitySink, ObservabilitySink
from .config import ExplainConfig
from .context import ExplainContext
from .orchestrator import ExplainOrchestrator
from .reanalyze import reanalyze_sync
from .result import ExplainRunOutput


[docs] class ExplainRunner: """Run the explain pipeline for a conversation snapshot.""" def __init__( self, config: ExplainConfig, registry: PluginRegistry | None = None, telemetry: ObservabilitySink | None = None, progress: CoalitionProgress | None = None, ) -> None: self._config = config self._registry = registry or _default_registry() self._telemetry = telemetry or NullObservabilitySink() self._progress = progress
[docs] async def explain(self, snapshot: ConversationSnapshot) -> ExplainRunOutput: """Execute the async explain pipeline.""" context = self._build_context(snapshot) backend = instantiate_backend(self._config.backend, self._registry) try: orchestrator = ExplainOrchestrator(context, backend) return await orchestrator.run() finally: aclose = getattr(backend, "aclose", None) if aclose is not None: await aclose()
[docs] def explain_sync(self, snapshot: ConversationSnapshot) -> ExplainRunOutput: """Execute the explain pipeline on the current event loop policy.""" return asyncio.run(self.explain(snapshot))
[docs] def reanalyze(self, archive_root: Path | str) -> ExplainRunOutput: """Rescore an archived run with the runner's current explain settings.""" return reanalyze_sync( Path(archive_root), self._config, self._registry, telemetry=self._telemetry, )
def _build_context(self, snapshot: ConversationSnapshot) -> ExplainContext: partitioner = self._registry.resolve( PluginGroup.PARTITIONS, self._config.explanation.players, ) if not isinstance(partitioner, TokenPartitioner): msg = "only token partitioners are supported by ExplainRunner" raise TypeError(msg) player_set = partitioner.partition(snapshot) run_id = f"{snapshot.snapshot_id}-{self._config.explanation.seed}" return ExplainContext( config=self._config, snapshot=snapshot, player_set=player_set, registry=self._registry, run_id=run_id, archive_root=self._resolve_archive_root(run_id), telemetry=self._telemetry, progress=self._progress, ) def _resolve_archive_root(self, run_id: str) -> Path | None: template = self._config.explanation.archive.path.strip() if not template: return None return Path(template.format(run_id=run_id))
def _default_registry() -> PluginRegistry: registry = PluginRegistry() register_builtin_plugins(registry) for group in ( PluginGroup.ESTIMATORS, PluginGroup.ESTIMANDS, PluginGroup.VALUE_FNS, PluginGroup.NORMALIZERS, PluginGroup.BACKENDS, PluginGroup.PARTITIONS, PluginGroup.ABSENCE_POLICIES, ): registry.load_entry_points(group) return registry