Source code for stormlog.infer.analysis

"""Analysis helpers for inference profiling artifacts."""

from __future__ import annotations

import json
from collections import Counter
from collections.abc import Iterable
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any

from ..session import SESSION_STATUS_COMPLETED
from .arrival_report import arrival_lines, arrival_summary, latency_from_intended_ms
from .cache_state import cache_lines, cache_summary
from .correlation_accounting import AlignedTimestamp
from .errors import InferInputError
from .host_clock import is_boot_qualified
from .latency_report import latency_summary, streaming_summary
from .populations import (
    CasePopulation,
    MeasuredInterval,
    case_populations,
    goodput,
    rate,
)
from .report_stats import int_value as _int_value
from .report_stats import is_number as _is_number
from .report_stats import number_values as _number_values
from .report_stats import percentile as _percentile
from .server_clock import (
    AMBIGUOUS,
    UNCOVERED,
    SampleAlignment,
    ServerClock,
    align_samples,
    artifact_alignments,
    build_server_clock,
    client_clock_domain,
)
from .server_group import members as group_members
from .server_group import membership_issue
from .slo import SloSpec, require_measured_window, slo_from_artifact
from .telemetry import ServerIdentity, TelemetrySample, load_telemetry
from .vllm_analysis import (
    JoinedSpans,
    joined_span_attributes,
    load_external_spans,
    vllm_case_lines,
    vllm_lines,
    vllm_report,
)
from .vllm_execution_report import execution_lines, execution_report
from .vllm_spans import VllmSpanRecord
from .workload_report import (
    length_summary,
    prompt_lines,
    prompt_summary,
    workload_lines,
    workload_summary,
)


[docs] def analyze_inference_events( path: str | Path, *, server_telemetry_paths: Iterable[str | Path] = (), direct_server: bool = False, clock_offset_ns: int | None = None, clock_uncertainty_ns: int | None = None, vllm_span_paths: Iterable[str | Path] = (), slo: SloSpec | None = None, slo_source: str = "flags", ) -> dict[str, Any]: """Analyze an inference profiling JSONL artifact. ``slo`` overrides the policy the artifact recorded, if it recorded one. """ records = _load_jsonl(path) policy = _policy(records, slo, slo_source) requests, samples = _partition_inference_records(records) server_samples = _load_server_samples(server_telemetry_paths) external_spans = _external_spans(records, vllm_span_paths) vllm = _vllm_telemetry(records, external_spans) spans = _joined_spans(records, external_spans) join, members = _server_join( records, server_samples, direct_server=direct_server, clock_offset_ns=clock_offset_ns, clock_uncertainty_ns=clock_uncertainty_ns, ) timelines = [(member, _member_timeline(member)) for member in members] ok_requests = [record for record in requests if record.get("status") == "ok"] cases = _case_reports( records, requests, samples, timelines, "group" in join, spans, policy, _server_admitted_ids(records, spans, external_spans), ) if timelines: join["case_coverage"] = _coverage_counts(cases) failed = [record for record in requests if record.get("status") != "ok"] return { # 2: per-case populations and intervals; throughput divides by the # case's declared interval (throughput.interval_seconds) instead of # the span of its successful requests (the old duration_seconds). "analysis_version": ANALYSIS_VERSION, "slo": None if policy is None else policy.to_record(), "summary": { "total_requests": len(requests), "successful_requests": len(ok_requests), "failed_requests": len(failed), "failure_rate": (len(failed) / len(requests)) if requests else 0.0, "failures_by_status": _failures_by_status(failed), "case_count": len(cases), "session_status": _session_status(records), }, "cases": cases, "workload": workload_summary(records), "telemetry": { "client_observation_scope": "client_local", "server_join": join, "server_targets": _server_targets(server_samples), "vllm": vllm, "execution": execution_report(records), }, }
ANALYSIS_VERSION = 2 @dataclass(frozen=True) class _Policy: """The SLO policy a report judges cases against, and where it came from.""" spec: SloSpec source: str # The policies the artifact recorded that this one replaces. overrides: tuple[dict[str, Any], ...] = () def to_record(self) -> dict[str, Any]: return { "name": self.spec.name, "digest": self.spec.digest(), "source": self.source, "policy": self.spec.to_record(), "overrides": [dict(item) for item in self.overrides], } def _policy( records: list[dict[str, Any]], slo: SloSpec | None, source: str ) -> _Policy | None: """The given policy, which replaces any the artifact recorded, or that one.""" if slo is not None: return _Policy(slo, source, overrides=_recorded_policies(records)) recorded = slo_from_artifact(records) if recorded is None: return None return _Policy( require_measured_window(recorded, "the artifact's infer.slo record"), "artifact", ) def _recorded_policies(records: list[dict[str, Any]]) -> tuple[dict[str, Any], ...]: """Each ``infer.slo`` record's name and digest, as the run recorded them.""" found = [] for record in records: if record.get("event_type") != "infer.slo": continue policy = record.get("slo") name = policy.get("name") if isinstance(policy, dict) else None found.append({"name": name, "digest": record.get("digest")}) return tuple(found)
[docs] def replaced_policies(report: dict[str, Any]) -> list[dict[str, Any]]: """The recorded policies a report's own policy replaced, if it differs.""" slo = report.get("slo") if not isinstance(slo, dict): return [] return [ item for item in slo.get("overrides") or [] if item.get("digest") != slo.get("digest") ]
def _external_spans( records: list[dict[str, Any]], span_paths: Iterable[str | Path] ) -> list[VllmSpanRecord]: """Span files given on the command line; one that cannot be read is invalid.""" try: return load_external_spans(records, span_paths) except (OSError, ValueError) as exc: raise InferInputError(f"vLLM telemetry: {_reason(exc)}") from exc def _vllm_telemetry( records: list[dict[str, Any]], external_spans: list[VllmSpanRecord] ) -> dict[str, Any]: """The vLLM block; a span record that cannot be read is an input error.""" try: return vllm_report(records, external_spans) except (OSError, ValueError) as exc: raise InferInputError(f"vLLM telemetry: {_reason(exc)}") from exc def _joined_spans( records: list[dict[str, Any]], external_spans: list[VllmSpanRecord] ) -> JoinedSpans: try: return joined_span_attributes(records, external_spans) except ValueError as exc: raise InferInputError(f"vLLM telemetry: {_reason(exc)}") from exc def _server_admitted_ids( records: list[dict[str, Any]], spans: JoinedSpans, external_spans: list[VllmSpanRecord], ) -> set[str] | None: """The requests the server confirmed it saw, or None with no server source. A request with a joined span reached the engine, and so did one whose spans conflict: the conflict is in what they say, not in whether they exist. The execution hook's records confirm a request by its ``X-Request-Id`` too. A span receiver that got nothing is a source that confirms nothing; no source at all leaves the count unknown, never 0. """ hook_ids = _hook_request_ids(records) if not (hook_ids or external_spans or _span_source(records)): return None return set(spans.by_request) | set(spans.quarantined) | hook_ids def _hook_request_ids(records: list[dict[str, Any]]) -> set[str]: """The run's requests the execution hook saw, by ``X-Request-Id``.""" ids = set() for record in records: if ( record.get("event_type") != "infer.request" or record.get("schema_version") != 2 ): continue x_request_id = (record.get("metadata") or {}).get("x_request_id") if x_request_id: ids.add(str(x_request_id)) return ids def _span_source(records: list[dict[str, Any]]) -> bool: """Whether the run received spans, or ran a receiver for them.""" for record in records: if record.get("event_type") == "infer.vllm_span": return True config = record.get("config") if record.get("event_type") == "infer.session" and isinstance(config, dict): if config.get("vllm_spans") is not None: return True return False def _case_reports( records: list[dict[str, Any]], requests: list[dict[str, Any]], samples: list[dict[str, Any]], timelines: list[tuple[_Member, _ServerTimeline]], grouped: bool, spans: JoinedSpans, policy: _Policy | None, server_ids: set[str] | None, ) -> dict[str, dict[str, Any]]: by_case: dict[str, list[dict[str, Any]]] = {} for record in requests: by_case.setdefault(str(record.get("case_id", "unknown")), []).append(record) windows = _measured_windows(records) cache_states = _case_records(records, "infer.cache_state") populations = case_populations(records, server_admitted_ids=server_ids) cases = {} for case_id, case_requests in sorted(by_case.items()): cases[case_id] = _case_report( case_requests, samples, timelines, grouped, windows.get(case_id), populations[case_id], ) cases[case_id]["cache"] = cache_summary(cache_states.get(case_id)) cases[case_id]["latency"] = latency_summary(case_requests, spans=spans) cases[case_id]["streaming"] = streaming_summary(case_requests) if policy is not None: cases[case_id]["slo"] = goodput( case_requests, policy.spec, populations[case_id].intervals.rate, spans=spans, slo_source=policy.source, cohort=populations[case_id].population, ).to_record() return cases def _case_report( case_requests: list[dict[str, Any]], samples: list[dict[str, Any]], timelines: list[tuple[_Member, _ServerTimeline]], grouped: bool, window: dict[str, Any] | None, population: CasePopulation, ) -> dict[str, Any]: """Summarize one case: latency from its completed requests, arrivals from all. Rates divide by the case's rate interval; its populations count every measured request. """ ok = [record for record in case_requests if record.get("status") == "ok"] report = _summarize_requests( ok, samples=_samples_for_request_window(samples, ok), interval=population.intervals.rate, ) report["population"] = population.population.to_record() report["intervals"] = population.intervals.to_record() report["arrivals"] = arrival_summary(case_requests, window) report["prompts"] = prompt_summary(case_requests, window) report["lengths"] = length_summary(ok) report["memory"].update(_server_case_memory(timelines, ok, grouped)) return report def _case_records( records: list[dict[str, Any]], event_type: str ) -> dict[str, dict[str, Any]]: return { str(record.get("case_id")): record for record in records if record.get("event_type") == event_type } def _measured_windows(records: list[dict[str, Any]]) -> dict[str, dict[str, Any]]: return { str(record.get("case_id")): record for record in records if record.get("event_type") == "infer.phase_window" and record.get("phase") == "measured" } def _session_status(records: list[dict[str, Any]]) -> str | None: """How the run ended: the last status its session records.""" statuses = [ str(record["status"]) for record in records if record.get("event_type") == "infer.session" and record.get("status") ] return statuses[-1] if statuses else None def _failures_by_status(failed: list[dict[str, Any]]) -> dict[str, int]: counts = Counter(str(record.get("status")) for record in failed) return dict(sorted(counts.items()))
[docs] def format_analysis_text(report: dict[str, Any]) -> str: """Render an inference analysis report as text.""" summary = report.get("summary", {}) lines = [ "Inference Profile Analysis", "-" * 28, f"Total requests: {summary.get('total_requests', 0)}", f"Successful requests: {summary.get('successful_requests', 0)}", f"Failed requests: {summary.get('failed_requests', 0)}" + _failure_breakdown(summary.get("failures_by_status")), f"Failure rate: {float(summary.get('failure_rate', 0.0)):.2%}", ] status = summary.get("session_status") if status not in (None, SESSION_STATUS_COMPLETED): lines.insert(2, f"Session status: {status}") cases = report.get("cases", {}) telemetry = report.get("telemetry", {}) join = telemetry.get("server_join", {}) vllm = telemetry.get("vllm") vllm_cases = vllm.get("cases", {}) if isinstance(vllm, dict) else {} lines.extend(workload_lines(report.get("workload"))) lines.extend(_policy_lines(report)) lines.append("Memory observations: client-local") lines.extend(_server_status_lines(join)) lines.extend(vllm_lines(vllm)) lines.extend(execution_lines(telemetry.get("execution"))) if isinstance(cases, dict) and cases: lines.append("") lines.append("Cases:") for case_id, case in cases.items(): lines.extend(_case_lines(case_id, case)) lines.extend(vllm_case_lines(vllm_cases.get(case_id))) return "\n".join(lines)
def _failure_breakdown(by_status: Any) -> str: if not isinstance(by_status, dict) or not by_status: return "" parts = ", ".join(f"{status} {count}" for status, count in by_status.items()) return f" ({parts})" def _server_status_lines(join: dict[str, Any]) -> list[str]: if join.get("status") == "joined": return _joined_status_lines(join) if join.get("status") not in {None, "not_configured"}: return [f"Server telemetry: unjoined ({join.get('reason')})"] return [] def _joined_status_lines(join: dict[str, Any]) -> list[str]: lines = [ "Server telemetry: joined to case windows (declared direct route, " f"{join.get('clock_alignment_evidence')} clock evidence)" ] counts = join.get("case_coverage") if isinstance(counts, dict): lines.append( f"Server coverage: {counts.get('observed', 0)} observed, " f"{counts.get('partial', 0)} partial, {counts.get('empty', 0)} empty" ) group = join.get("group") if isinstance(group, dict): lines.append( f"Server group: {group.get('group_id')} " f"({group.get('world_size')} members)" ) for label, invalidation in _invalidations(join): lines.append( f"Server identity ended{label} ({invalidation.get('detail')}); " "case windows after the last confirmed poll are not joined" ) return lines def _invalidations(join: dict[str, Any]) -> list[tuple[str, dict[str, Any]]]: group = join.get("group") if not isinstance(group, dict): invalidation = join.get("invalidation") return [("", invalidation)] if isinstance(invalidation, dict) else [] return [ (f" for rank {member.get('rank')}", member["invalidation"]) for member in group.get("members", []) if isinstance(member.get("invalidation"), dict) ] def _case_lines(case_id: str, case: Any) -> list[str]: latency = case.get("latency_ms", {}) if isinstance(case, dict) else {} throughput = case.get("throughput", {}) if isinstance(case, dict) else {} lines = [ f"- {case_id}: " f"p50 E2E={_fmt(latency.get('e2e_p50'))} ms, " f"p95 E2E={_fmt(latency.get('e2e_p95'))} ms, " f"p50 TTFT={_fmt(latency.get('ttft_p50'))} ms, " f"output={_fmt(throughput.get('output_tokens_per_second'))} tok/s, " f"requests={_fmt(throughput.get('requests_per_second'))} req/s" ] if isinstance(case, dict): lines.extend( _interval_lines(throughput, case.get("population"), case.get("intervals")) ) lines.extend(_slo_lines(case.get("slo"))) lines.extend(arrival_lines(case.get("arrivals"), case.get("latency_ms"))) lines.extend(prompt_lines(case.get("prompts"))) lines.extend(cache_lines(case.get("cache"))) memory = case.get("memory", {}) if isinstance(case, dict) else {} lines.extend(_server_case_lines(memory)) return lines def _interval_lines(throughput: Any, population: Any, intervals: Any) -> list[str]: """What the rates divide by, and whether the request cohort is whole.""" lines = [] if isinstance(throughput, dict) and throughput.get("interval_kind"): cohort = str(throughput.get("numerator_cohort", "")).replace("_", " ") lines.append( f" rates per {throughput['interval_kind'].replace('_', ' ')} of " f"{_fmt(throughput.get('interval_seconds'))} s ({cohort})" ) elif isinstance(intervals, dict) and intervals.get("rate_reason"): lines.append(f" no rate interval ({intervals['rate_reason']})") if isinstance(population, dict) and not population.get("cohort_valid", True): issues = ", ".join(str(issue) for issue in population.get("issues", [])) lines.append(f" cohort invalid: {issues}") return lines def _policy_lines(report: dict[str, Any]) -> list[str]: slo = report.get("slo") if not isinstance(slo, dict): return [] lines = [ f"SLO policy: {slo.get('name')} from {slo.get('source')}, " f"digest {str(slo.get('digest'))[:12]}" ] lines.extend( f" replaces the recorded policy {item.get('name')} " f"(digest {str(item.get('digest'))[:12]})" for item in replaced_policies(report) ) return lines def _slo_lines(slo: Any) -> list[str]: """Attainment and goodput as bounds, or why the policy could not judge.""" if not isinstance(slo, dict): return [] name = slo.get("slo_name") if slo.get("status") != "evaluated": return [f" SLO {name}: unmeasurable ({slo.get('reason')})"] low, high = slo.get("attainment_lower"), slo.get("attainment_upper") attainment = ( _fmt_share(low) if low == high else f"{_fmt_share(low)}-{_fmt_share(high)}" ) rates = (slo.get("goodput_lower_rps"), slo.get("goodput_upper_rps")) goodput_text = ( _fmt(rates[0]) if rates[0] == rates[1] else f"{_fmt(rates[0])}-{_fmt(rates[1])}" ) return [ f" SLO {name}: attainment {attainment} of {slo.get('offered')} offered, " f"goodput {goodput_text} req/s, {slo.get('unknown')} unknown" ] def _fmt_share(value: Any) -> str: return f"{float(value):.1%}" if isinstance(value, (int, float)) else "-" def _server_case_lines(memory: Any) -> list[str]: if not isinstance(memory, dict): return [] if isinstance(memory.get("server_members"), list): return [ line for member in memory["server_members"] for line in _server_member_lines(member) ] coverage = memory.get("server_coverage") or {} if coverage.get("status") == "empty": return [f" server telemetry: none ({coverage.get('reason')})"] lines = [] if coverage.get("status") == "partial": lines.append(f" server telemetry: partial ({coverage.get('reason')})") for metric, observation in (memory.get("server_observations") or {}).items(): lines.append(_server_metric_line(metric, observation)) return lines def _server_member_lines(member: dict[str, Any]) -> list[str]: where = member.get("device_uuid") or member.get("host") label = f" rank {member.get('rank')} ({where})" coverage = member.get("coverage") or {} if coverage.get("status") == "empty": return [f"{label}: none ({coverage.get('reason')})"] lines = [] if coverage.get("status") == "partial": lines.append(f"{label}: partial ({coverage.get('reason')})") for metric, observation in (member.get("observations") or {}).items(): lines.append(_server_metric_line(metric, observation, prefix=f"{label} ")) return lines def _server_metric_line( metric: str, observation: dict[str, Any], prefix: str = " " ) -> str: label = f"{prefix}{metric} ({observation.get('observation_scope')})" value = observation.get("maximum_recorded_bytes") if value is None: missing = observation.get("missing_samples", 0) return f"{label}: no valid samples ({missing} missing)" return ( f"{label}: max recorded {value} bytes, " f"{observation.get('valid_samples')} samples" ) # Every profile writes a session record first; requests follow. _ARTIFACT_EVENT_TYPES = frozenset({"infer.session", "infer.request"}) def _load_jsonl(path: str | Path) -> list[dict[str, Any]]: """Read an inference artifact; a file that cannot be read is invalid input.""" try: records = _read_jsonl(Path(path)) except (OSError, ValueError) as exc: raise InferInputError(f"{path}: {_reason(exc)}") from exc if not any(record.get("event_type") in _ARTIFACT_EVENT_TYPES for record in records): raise InferInputError( f"{path}: not an inference artifact (no infer.session or " "infer.request records)" ) return records def _read_jsonl(path: Path) -> list[dict[str, Any]]: records: list[dict[str, Any]] = [] with path.open("r", encoding="utf-8") as handle: for line_number, raw_line in enumerate(handle, start=1): line = raw_line.strip() if not line: continue payload = json.loads(line) if not isinstance(payload, dict): raise ValueError(f"Line {line_number} is not a JSON object") records.append(payload) return records def _reason(exc: Exception) -> str: if isinstance(exc, OSError) and exc.strerror: return exc.strerror return str(exc) def _summarize_requests( requests: list[dict[str, Any]], *, samples: list[dict[str, Any]], interval: MeasuredInterval | None, ) -> dict[str, Any]: e2e = _number_values(requests, "e2e_latency_ms") from_intended = latency_from_intended_ms(requests) ttft = _number_values(requests, "ttft_ms") first_chunk = _number_values(requests, "first_chunk_latency_ms") output_tokens = sum(_int_value(record.get("output_tokens")) for record in requests) total_tokens = sum(_int_value(record.get("total_tokens")) for record in requests) request_count = len(requests) peak_device_used = _peak_sample_value(samples, "device_used_bytes") peak_process_rss = _peak_sample_value(samples, "process_rss_bytes") return { "request_count": request_count, "latency_ms": { "e2e_p50": _percentile(e2e, 50), "e2e_p95": _percentile(e2e, 95), "e2e_p99": _percentile(e2e, 99), "e2e_from_intended_p50": _percentile(from_intended, 50), "e2e_from_intended_p95": _percentile(from_intended, 95), "e2e_from_intended_p99": _percentile(from_intended, 99), "ttft_p50": _percentile(ttft, 50), "ttft_p95": _percentile(ttft, 95), "ttft_p99": _percentile(ttft, 99), "first_chunk_p50": _percentile(first_chunk, 50), "first_chunk_p95": _percentile(first_chunk, 95), }, "throughput": _throughput(request_count, output_tokens, total_tokens, interval), "tokens": { "output_tokens": output_tokens, "total_tokens": total_tokens, "output_token_sources": sorted( { str(record.get("output_token_source", "unknown")) for record in requests } ), }, "memory": { "observation_scope": "client_local", "peak_device_used_bytes": peak_device_used, "peak_process_rss_bytes": peak_process_rss, }, } def _throughput( requests: int, output_tokens: int, total_tokens: int, interval: MeasuredInterval | None, ) -> dict[str, Any]: """Successful requests and their tokens per second of the rate interval. The rates are null when the interval has no length, instead of zero. """ return { "interval_seconds": interval.seconds if interval is not None else None, "interval_kind": interval.kind if interval is not None else None, "numerator_cohort": ( interval.numerator_cohort if interval is not None else None ), "requests_per_second": rate(requests, interval), "output_tokens_per_second": rate(output_tokens, interval), "total_tokens_per_second": rate(total_tokens, interval), } def _samples_for_request_window( samples: list[dict[str, Any]], requests: list[dict[str, Any]], ) -> list[dict[str, Any]]: request_window = _request_time_window(requests) if request_window is None: return [] start_ns, end_ns = request_window return [ sample for sample in samples if _is_number(sample.get("timestamp_ns")) and start_ns <= _int_value(sample.get("timestamp_ns")) <= end_ns ] def _request_time_window(requests: list[dict[str, Any]]) -> tuple[int, int] | None: bounds: list[tuple[int, int]] = [] for record in requests: started_at = record.get("started_at_ns") ended_at = record.get("ended_at_ns") if not _is_number(started_at) or not _is_number(ended_at): continue bounds.append((_int_value(started_at), _int_value(ended_at))) if not bounds: return None return min(start for start, _end in bounds), max(end for _start, end in bounds) def _peak_sample_value(samples: list[dict[str, Any]], field: str) -> int | None: return max( ( _int_value(sample.get(field)) for sample in samples if _is_number(sample.get(field)) ), default=None, ) def _fmt(value: Any) -> str: if isinstance(value, (int, float)): return f"{float(value):.2f}" return "-" def _partition_inference_records( records: list[dict[str, Any]], ) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: requests = [ record for record in records if record.get("event_type") == "infer.request" and record.get("phase") == "measured" ] samples = [ record for record in records if record.get("event_type") == "infer.system_sample" ] return requests, samples def _load_server_samples(paths: Iterable[str | Path]) -> list[TelemetrySample]: """Load every artifact; drop exact duplicates, e.g. a file passed twice.""" return list( dict.fromkeys(sample for path in paths for sample in _server_samples(path)) ) def _server_samples(path: str | Path) -> list[TelemetrySample]: try: return load_telemetry(path) except (OSError, ValueError) as exc: raise InferInputError(f"--server-telemetry {path}: {_reason(exc)}") from exc @dataclass(frozen=True) class _Member: """One joined server identity, its clock and its samples on the client clock.""" identity: ServerIdentity samples: list[TelemetrySample] clock: ServerClock placed: SampleAlignment def _server_join( records: list[dict[str, Any]], samples: list[TelemetrySample], *, direct_server: bool, clock_offset_ns: int | None, clock_uncertainty_ns: int | None, ) -> tuple[dict[str, Any], list[_Member]]: """Return a case-window join for one server or one declared group. The second value lists the joined members, ordered by rank. """ if not samples: return {"status": "not_configured"}, [] artifact = _artifact_record(records) if artifact is None: return _unjoined("missing_run_identity"), [] issue = _run_id_issue(artifact, samples) or _identity_issue(samples, direct_server) if issue is not None: return issue, [] by_identity = group_members(samples) clocks = _member_clocks( artifact, records, list(by_identity), clock_offset_ns, clock_uncertainty_ns ) members = clocks if isinstance(clocks, str) else _place(by_identity, clocks) if isinstance(members, str): return _unjoined(members), [] return _joined(members), members def _place( by_identity: dict[ServerIdentity, list[TelemetrySample]], clocks: dict[ServerIdentity, ServerClock], ) -> list[_Member] | str: """Place each member's samples; every member needs at least one placed.""" members = [ _Member( identity, found, clocks[identity], align_samples(found, clocks[identity]) ) for identity, found in by_identity.items() ] for member in members: if not member.placed.aligned: ambiguous = member.placed.unaligned[AMBIGUOUS] return f"clock_alignment_{'ambiguous' if ambiguous else 'uncovered'}" return members def _unjoined(reason: str) -> dict[str, Any]: return {"status": "unjoined", "reason": reason} def _identity_issue( samples: list[TelemetrySample], direct_server: bool ) -> dict[str, Any] | None: membership = membership_issue(samples) if membership is not None: return _unjoined(membership) route_issue = _route_issue(samples, direct_server) return _unjoined(route_issue) if route_issue else None def _member_clocks( artifact: dict[str, Any], records: list[dict[str, Any]], identities: list[ServerIdentity], offset_ns: int | None, uncertainty_ns: int | None, ) -> dict[ServerIdentity, ServerClock] | str: """Build one clock per server domain; the flags describe one remote domain.""" client_domain = client_clock_domain(artifact) if client_domain is None: return "missing_client_clock_domain" domains = sorted({identity.clock_domain for identity in identities}) flags_domain = _flags_domain(domains, client_domain, offset_ns) if flags_domain is None: return "clock_flags_ambiguous" run_id = str(artifact["context"].get("run_id")) recorded = artifact_alignments(records, run_id) clocks: dict[str, ServerClock] = {} for domain in domains: flags = (offset_ns, uncertainty_ns) if domain == flags_domain else (None, None) clock = build_server_clock( run_id=run_id, server_domain=domain, client_domain=client_domain, recorded=recorded, offset_ns=flags[0], uncertainty_ns=flags[1], ) if isinstance(clock, str): return clock clocks[domain] = clock return {identity: clocks[identity.clock_domain] for identity in identities} def _flags_domain( domains: list[str], client_domain: str, offset_ns: int | None ) -> str | None: """The server domain the clock flags describe, or None if that is unclear. Members on the client's host and boot share its clock; the flags then belong to the single remote domain. A member with the client's hostname but no boot ID counts as remote, because only the flags can align it. With several remote hosts, each needs its own ``infer.clock_alignment`` record. """ remote = [ domain for domain in domains if domain != client_domain or not is_boot_qualified(domain) ] if offset_ns is not None and len(remote) > 1: return None return remote[0] if remote else client_domain def _joined(members: list[_Member]) -> dict[str, Any]: return { "status": "joined", "route_evidence": "operator_declared_direct", **_membership_report(members), **_clock_report(members), "attribution": "case_window_observation_only", } def _membership_report(members: list[_Member]) -> dict[str, Any]: first = members[0].identity if first.group_id is None: return { "identity": asdict(first), "invalidation": _invalidation(members[0].samples), } return { "group": { "group_id": first.group_id, "world_size": first.world_size, "members": [ { "rank": member.identity.rank, "identity": asdict(member.identity), "clock_alignment_evidence": member.clock.evidence, "invalidation": _invalidation(member.samples), } for member in members ], } } def _clock_report(members: list[_Member]) -> dict[str, Any]: applied = _merge_applied( [item for member in members for item in member.placed.applied] ) offsets = {item["offset_ns"] for item in applied} report: dict[str, Any] = { "clock_alignment_evidence": _evidence(members), # One offset when a single alignment placed every joined sample. "clock_offset_ns": next(iter(offsets)) if len(offsets) == 1 else None, "clock_uncertainty_ns": _max_uncertainty(members), "clock_alignments": applied, "unaligned_samples": _unaligned(members), } for key, field in ( ("overridden_clock_alignments", "overridden"), ("ignored_clock_alignments", "ignored"), ): event_ids = _record_ids(members, field) if event_ids: report[key] = event_ids return report def _record_ids(members: list[_Member], field: str) -> list[str]: """Event IDs a member's clock replaced or ignored, across the whole join.""" return sorted({item for member in members for item in getattr(member.clock, field)}) def _max_uncertainty(members: list[_Member]) -> int: return max( placed.uncertainty_ns for member in members for placed in member.placed.aligned.values() ) def _unaligned(members: list[_Member]) -> dict[str, int]: return { key: sum(member.placed.unaligned[key] for member in members) for key in (UNCOVERED, AMBIGUOUS) } def _evidence(members: list[_Member]) -> str: evidence = sorted({member.clock.evidence for member in members}) return evidence[0] if len(evidence) == 1 else "mixed" def _merge_applied(items: list[dict[str, Any]]) -> list[dict[str, Any]]: """Combine one alignment used by several members into one entry.""" merged: dict[tuple[Any, ...], dict[str, Any]] = {} for item in items: key = tuple(value for name, value in item.items() if name != "aligned_samples") if key in merged: merged[key]["aligned_samples"] += item["aligned_samples"] else: merged[key] = dict(item) return list(merged.values()) def _run_id_issue( artifact: dict[str, Any], samples: list[TelemetrySample] ) -> dict[str, Any] | None: """A different run ID leaves the server data out; the client report stays.""" artifact_run_id = artifact["context"].get("run_id") telemetry_run_ids = sorted({sample.run_id for sample in samples}) if telemetry_run_ids == [artifact_run_id]: return None return { "status": "unjoined", "reason": "run_id_mismatch", "artifact_run_id": artifact_run_id, "telemetry_run_ids": telemetry_run_ids, } def _route_issue(samples: list[TelemetrySample], direct_server: bool) -> str | None: if not direct_server: return "route_not_declared" if not any(sample.state == "valid" for sample in samples): return "no_valid_server_samples" return None def _invalidation(samples: list[TelemetrySample]) -> dict[str, Any] | None: """Locate where the observed process or GPU stopped being the original one. Samples before the first ``invalid`` sample describe the original identity, so an invalidation only affects case windows that could reach past the last poll that still confirmed it. """ first = min( (sample for sample in samples if sample.state == "invalid"), key=lambda sample: sample.observed_at_ns, default=None, ) if first is None: return None confirmed = [ sample.observed_at_ns for sample in samples if sample.observed_at_ns < first.observed_at_ns ] return { "observed_at_ns": first.observed_at_ns, "detail": first.detail, "last_confirmed_at_ns": max(confirmed, default=None), } def _artifact_record(records: list[dict[str, Any]]) -> dict[str, Any] | None: artifacts = [r for r in records if r.get("event_type") == "infer.artifact"] if len(artifacts) != 1 or not isinstance(artifacts[0].get("context"), dict): return None return artifacts[0] def _server_targets(samples: list[TelemetrySample]) -> list[dict[str, Any]]: grouped: dict[ServerIdentity, list[TelemetrySample]] = {} for sample in samples: grouped.setdefault(sample.identity, []).append(sample) targets: list[dict[str, Any]] = [] for identity in sorted( grouped, key=lambda item: ( item.host, item.pid, item.process_start_ns, item.device_uuid or "", ), ): selected = grouped[identity] states = Counter(s.state for s in selected) targets.append( { "identity": asdict(identity), "metrics": sorted({s.metric for s in selected}), "sample_states": { state: states[state] for state in ("valid", "missing", "stale", "invalid") }, } ) return targets @dataclass(frozen=True) class _ServerTimeline: """A joined collector's polls in client time, and how long it can be trusted. Each sample keeps the uncertainty of the alignment that placed it, so one imprecise alignment does not widen the margins of every other sample. """ placed: dict[TelemetrySample, AlignedTimestamp] # Samples no single alignment placed, and the offsets that placed others. unplaced: dict[TelemetrySample, str] offsets: tuple[int, ...] slack_ns: int first_poll_ns: int last_poll_ns: int invalidated: bool trusted_until_ns: int | None def _member_timeline(member: _Member) -> _ServerTimeline: aligned = member.placed.aligned values = [placed.value_ns for placed in aligned.values()] invalidation = _invalidation(member.samples) return _ServerTimeline( placed=aligned, unplaced=member.placed.unplaced, offsets=tuple(sorted({item["offset_ns"] for item in member.placed.applied})), slack_ns=max(sample.interval_ms for sample in member.samples) * 1_000_000, first_poll_ns=min(values), last_poll_ns=max(values), invalidated=invalidation is not None, trusted_until_ns=_trusted_until(aligned, invalidation), ) def _trusted_until( aligned: dict[TelemetrySample, AlignedTimestamp], invalidation: dict[str, Any] | None, ) -> int | None: """Client time by which the last confirming poll had certainly happened.""" if invalidation is None: return None confirmed = [ # A poll could have happened up to its own uncertainty earlier. placed.value_ns - placed.uncertainty_ns for sample, placed in aligned.items() if sample.observed_at_ns < invalidation["observed_at_ns"] ] return max(confirmed) if confirmed else None def _server_case_memory( timelines: list[tuple[_Member, _ServerTimeline]], requests: list[dict[str, Any]], declared_group: bool, ) -> dict[str, Any]: """Server values for one case: one server's, or one entry per group member. Values from different members are never combined: separate collectors sample at different instants, so no per-case total would be well defined. """ if not timelines: return {"server_observations": {}, "server_coverage": {"status": "not_joined"}} views = [] for member, timeline in timelines: observations, coverage = _server_case_view(member.samples, requests, timeline) views.append((member, observations, coverage)) if not declared_group: _member, observations, coverage = views[0] return {"server_observations": observations, "server_coverage": coverage} return { "server_observations": {}, "server_coverage": _group_coverage([coverage for _, _, coverage in views]), "server_members": [ _member_view(member, observations, coverage) for member, observations, coverage in views ], } def _member_view( member: _Member, observations: dict[str, Any], coverage: dict[str, Any] ) -> dict[str, Any]: identity = member.identity return { "rank": identity.rank, "host": identity.host, "pid": identity.pid, "device_uuid": identity.device_uuid, "gpu_instance_id": identity.gpu_instance_id, "observations": observations, "coverage": coverage, } def _group_coverage(coverages: list[dict[str, Any]]) -> dict[str, Any]: statuses = {coverage["status"] for coverage in coverages} if statuses == {"observed"}: return {"status": "observed", "reason": None} if statuses == {"empty"}: return {"status": "empty", "reason": "no_member_observed"} return {"status": "partial", "reason": "some_members_not_fully_observed"} def _server_case_view( samples: list[TelemetrySample], requests: list[dict[str, Any]], timeline: _ServerTimeline, ) -> tuple[dict[str, Any], dict[str, Any]]: """Summarize server samples in one case window and say how well it is covered. A sample counts only if its own clock uncertainty cannot move it outside the window, so a short window or a collector that was not running can leave a joined case without server values; the coverage record says which. """ window = _request_time_window(requests) if window is None: return {}, {"status": "empty", "reason": "no_request_window"} start_ns, end_ns = window in_window = [ sample for sample, placed in timeline.placed.items() if start_ns + placed.uncertainty_ns <= placed.value_ns <= end_ns - placed.uncertainty_ns ] margin = _case_margin(timeline, start_ns, end_ns) low, high = start_ns + margin, end_ns - margin gap = _unplaced_reason(timeline, start_ns, end_ns) coverage = _case_coverage(end_ns, low, high, bool(in_window), gap, timeline) if coverage["status"] == "empty": return {}, coverage return _metric_summaries(samples, in_window), coverage def _case_margin(timeline: _ServerTimeline, start_ns: int, end_ns: int) -> int: """The clock uncertainty that describes one case window's coverage. It is the smallest uncertainty among samples placed inside the window, or among all samples when none are, so the counted window is the widest any sample could qualify for. """ inside = [ placed.uncertainty_ns for placed in timeline.placed.values() if start_ns <= placed.value_ns <= end_ns ] return min(inside or [placed.uncertainty_ns for placed in timeline.placed.values()]) def _unplaced_reason( timeline: _ServerTimeline, start_ns: int, end_ns: int ) -> str | None: """Say why a case has no samples when unplaced polls probably fell inside it. Unplaced samples have no client time; any offset that placed other samples gives the best estimate of where they would land. """ near = [ reason for sample, reason in timeline.unplaced.items() if any( start_ns <= sample.observed_at_ns + offset <= end_ns for offset in timeline.offsets ) ] if not near: return None return f"clock_alignment_{AMBIGUOUS if AMBIGUOUS in near else UNCOVERED}" def _case_coverage( end_ns: int, low: int, high: int, has_samples: bool, unplaced_reason: str | None, timeline: _ServerTimeline, ) -> dict[str, Any]: empty_reason = _empty_reason( end_ns, low, high, has_samples, unplaced_reason, timeline ) if empty_reason is not None: return { "status": "empty", "reason": empty_reason, "counted_window_ns": [low, high] if high >= low else None, } partial_reason = _partial_reason(low, high, timeline) return { "status": "observed" if partial_reason is None else "partial", "reason": partial_reason, "counted_window_ns": [low, high], } def _empty_reason( end_ns: int, low: int, high: int, has_samples: bool, unplaced_reason: str | None, timeline: _ServerTimeline, ) -> str | None: if high < low: return "window_shorter_than_uncertainty" if timeline.invalidated and ( timeline.trusted_until_ns is None or end_ns > timeline.trusted_until_ns ): return "identity_invalidated" if not has_samples: return unplaced_reason or "no_collector_coverage" return None def _partial_reason(low: int, high: int, timeline: _ServerTimeline) -> str | None: if timeline.first_poll_ns > low + timeline.slack_ns: return "collector_started_after_window_start" if timeline.last_poll_ns < high - timeline.slack_ns: return "collector_stopped_before_window_end" return None def _coverage_counts(cases: dict[str, Any]) -> dict[str, int]: statuses = Counter( case["memory"]["server_coverage"]["status"] for case in cases.values() ) return {status: statuses[status] for status in ("observed", "partial", "empty")} def _metric_summaries( samples: list[TelemetrySample], in_window: list[TelemetrySample] ) -> dict[str, Any]: by_metric: dict[str, TelemetrySample] = {} for sample in samples: by_metric.setdefault(sample.metric, sample) window_by_metric: dict[str, list[TelemetrySample]] = {} for sample in in_window: window_by_metric.setdefault(sample.metric, []).append(sample) return { metric: _summarize_server_metric(first, window_by_metric.get(metric, [])) for metric, first in sorted(by_metric.items()) } def _summarize_server_metric( first: TelemetrySample, window_samples: list[TelemetrySample], ) -> dict[str, Any]: """Describe one metric using only the samples inside the counted window.""" states = Counter(sample.state for sample in window_samples) valid_values = [ sample.value_bytes for sample in window_samples if sample.state == "valid" and sample.value_bytes is not None ] return { "observation_scope": first.scope, "counter_owner": first.counter_owner, "provenance": sorted({sample.provenance for sample in window_samples}), "sources": sorted({sample.source for sample in window_samples}), "maximum_recorded_bytes": max(valid_values, default=None), "valid_samples": len(valid_values), "missing_samples": states["missing"], "stale_samples": states["stale"], "invalid_samples": states["invalid"], "intervals_ms": sorted({sample.interval_ms for sample in window_samples}), }