From a4e4deba6bd3b931207a594849b6489bd4af3371 Mon Sep 17 00:00:00 2001 From: aknvda Date: Mon, 5 Oct 2026 14:33:16 -0700 Subject: [PATCH 1/3] Add optional NVTX tracing for framework pin lifecycles Signed-off-by: aknvda --- docs/profiling.md | 141 ++++++++++++++ examples/nvtx_pin_capture.py | 195 +++++++++++++++++++ pyproject.toml | 3 + src/kvcr/_nvtx.py | 235 +++++++++++++++++++++++ src/kvcr/remote_fw_dram.py | 78 +++++++- tests/unit/test_nvtx.py | 355 +++++++++++++++++++++++++++++++++++ uv.lock | 43 +++++ 7 files changed, 1043 insertions(+), 7 deletions(-) create mode 100644 docs/profiling.md create mode 100644 examples/nvtx_pin_capture.py create mode 100644 src/kvcr/_nvtx.py create mode 100644 tests/unit/test_nvtx.py diff --git a/docs/profiling.md b/docs/profiling.md new file mode 100644 index 0000000..be18ddb --- /dev/null +++ b/docs/profiling.md @@ -0,0 +1,141 @@ + + +# NVTX framework-pin tracing + +This first instrumentation checkpoint traces framework pin requests used by the +remote-memory path. It separates the synchronous `request_pin` callback from the +asynchronous wait for its result. A shared pin has one lifetime and separate +associations to each waiting source operation. + +It does **not** yet trace NIXL submission/completion, target processing, router +hints, or caller polling. Complete request/session correlation and performance +validation remain separate work. Review a decoded pin capture before expanding +the hooks. + +## Enable tracing + +On Linux, install the optional Python bindings and structured-payload dependency: + +```bash +uv sync --extra profiling +``` + +`KVCR_NVTX_LEVEL` is read when a KVCR remote-memory backend is constructed: + +| Value | Behavior | +| --- | --- | +| `off` | No NVTX/NumPy import, pin trace objects, or payload construction. | +| `low` | Pin callback scopes, registration, waiter associations, terminal events. Default when the profiling dependencies are available. | +| `medium` | Low detail plus individual waiter-detachment events. | + +Without the optional dependencies, tracing is a no-op. An explicit request for +tracing with an unavailable backend warns once. An unsupported level warns once +and disables tracing. There is no `high` level in this checkpoint. + +Tracing is independent of KVCR telemetry and is not gated on profiler attachment. +Consequently, low/medium payload construction also costs work outside a profiler. +Use `off` as the baseline when measuring overhead. Annotation failures are kept +out of KVCR's result and resource-ownership paths. + +## Capture a real pin and transfer + +The example starts two real KVCR/NIXL UCX agents in **one process**, with registered +host memory and a small framework binding that completes pins after a configured +delay. It checks transferred bytes and pin release. It requires Linux and NIXL; +it does not run a model, Dynamo, vLLM, or a GPU transfer. + +```bash +KVCR_NVTX_LEVEL=low nsys profile \ + --trace=nvtx,cuda --sample=none --cpuctxsw=none \ + --output=kvcr-pin-success \ + uv run --extra profiling python examples/nvtx_pin_capture.py \ + --blocks 2 --delay-ms 20 + +nsys export --type=sqlite --include-json=true \ + --output=kvcr-pin-success.sqlite kvcr-pin-success.nsys-rep +``` + +For controlled failed results and timeouts, use `--scenario failure` and +`--scenario timeout` with distinct report names. Timeout deliberately delays the +pin beyond the 10-second operation deadline. These are correctness examples, +not throughput benchmarks. Keep the reports outside the source checkout. + +Extended payloads require a collector/viewer with support for them. The initial +probe used Python 3.12, `nvtx` 0.2.16, NumPy 2.5.3, and Nsight Systems 2025.3.2. +Verify decoding in your environment before relying on a capture. The CUDA runtime +package `nvidia-nvtx` does not supply the Python `nvtx` annotation API. + +In the SQLite export, inspect `NVTX_EVENTS.jsonText` (enabled by `--include-json`) +and join its `textId` to `StringIds.id` for the event name. For example: + +```sql +SELECT e.start, e.end, s.value AS event, e.jsonText +FROM NVTX_EVENTS AS e JOIN StringIds AS s ON s.id = e.textId +WHERE s.value LIKE 'source.pin.%' +ORDER BY e.start; +``` + +## Interpretation + +All annotations use the `KVCR` domain and `framework_pin` category. Event names +are bounded static strings. IDs and counts are payloads, never registered names. + +| Event | Meaning | +| --- | --- | +| `source.pin.framework` | Same-thread push/pop around the framework's `request_pin` callback, including exceptions. | +| `source.pin.registered` | A new physical request was accepted into KVCR's pending-pin state. | +| `source.pin.waiter` | Association between that pin and one source operation/target operation handle. | +| `source.pin.completed` | KVCR observed a usable/failed result, deadline, cancellation, or shutdown. Exactly one terminal observation per pin trace. | +| `source.pin.detached` | One source operation stopped waiting; other waiters may remain. Medium detail only. | + +Registration-to-completion measures the observed asynchronous wait, including +delay until KVCR polls the framework. It is not the framework's internal execution +time. A cancellation marker reports KVCR's decision; it does not establish native +transfer quiescence or permission to reuse a buffer. Late results are still +discarded/released by the existing lifecycle and do not emit a second completion. + +Schema version 1 uses a fresh structured NumPy payload for every event: + +| Field | Interpretation | +| --- | --- | +| `instance_hi`, `instance_lo` | Two uint64 halves of a UUID assigned to this KVCR tracer instance. | +| `pin_id` | Independent uint64 sequence within that instance; distinguishes framework request-ID reuse. | +| `pin_request_known`, `pin_request_id` | Signed int64 framework request ID and availability flag. The flag is zero before the callback returns or when its Python integer is outside int64 range; the independent pin identity and lifecycle events are still recorded. Zero is a valid request ID when the flag is set. | +| `source_op_id`, `op_handle` | Signed int64 source and target operation handles. Zero denotes unavailable context in direct helper calls. The source path carries these even if the pin callback fails before registration. | +| `requested_blocks`, `completed_blocks` | Requested count and count of non-missing entries in an accepted result. `-1` means the completed count is unavailable. | +| `status`, `reason` | Bounded codes below. | +| `fw_dram_utilization_known` | Always zero in this checkpoint: no framework utilization source is exposed by the current bindings. | + +Join pin events by `(instance_hi, instance_lo, pin_id)`, not by framework request +ID alone. Operation handles alone are not globally unique across target agents. +This schema is a source-side pin checkpoint, not a cross-worker identity scheme. +No request, session, or parent-session identity is inferred from those handles. + +Status codes: `1=pending`, `2=success`, `3=partial`, `4=failed`, `5=timeout`, +`6=cancelled`. Partial reports how many blocks were available, without asserting +why others were missing. + +Reason codes: `0=unknown`, `1=none`, `2=callback_error`, `3=duplicate_request`, +`4=invalid_result`, `5=deadline`, `6=cancelled`, `7=shutdown`, `8=no_waiters`. +A framework result of `None` gives no failure cause: it is **not** evidence of +cache pressure. Callback errors identify the stage, not an underlying diagnosis. +Exception strings are not recorded. + +## Integration gaps on the reviewed baseline + +On `fb4f264`, `submit_hint()` accepts a `request_id`; the hint parser retains only +source endpoint and block hashes. The target pull retains the request ID, but its +`start_write` message does not transport request/session/parent-session context. +`KVCRBindings` exposes pin callbacks, not framework-cache utilization. Capturing +those inputs needs a verified integration source and potentially a separately +scoped API or wire change. KVCR local-cache occupancy is not a substitute. + +The pin hooks can also be encountered by remote fetch. They do not provide full +fetch fan-out correlation. No production-overhead acceptance threshold has been +established by this example. + +References: [NVTX Python best practices](https://nvidia.github.io/NVTX/python/best_practices.html) +and [extended payload examples](https://nvidia.github.io/NVTX/python/annotation_attributes.html). diff --git a/examples/nvtx_pin_capture.py b/examples/nvtx_pin_capture.py new file mode 100644 index 0000000..7408ea9 --- /dev/null +++ b/examples/nvtx_pin_capture.py @@ -0,0 +1,195 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +"""Capture framework pinning through two real KVCR/NIXL agents on Linux. + +The framework binding is deliberately small: it returns registered host-memory +blocks after a configurable delay. This validates the pin instrumentation, not +Dynamo/vLLM integration, GPU transfers, or production performance. +""" + +import argparse +import ctypes +import json +import socket +import time +from contextlib import ExitStack + +from kvcr import KVCR, KVCRBindings +from kvcr.config import KVCRBackendConfigs, KVCRConfig, RemoteFWDramOptions +from kvcr.control_channels import ZmqPeerControlChannel +from kvcr.types import BlockKey, MemoryRef, PinRequestId, RegionDescriptor + + +def control_channel(): + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + port = sock.getsockname()[1] + return ZmqPeerControlChannel("127.0.0.1", port, "127.0.0.1") + + +class Framework: + def __init__(self, name, keys, delay, fail): + self.name, self.keys, self.delay, self.fail = name, keys, delay, fail + self.pending = {} + self.requests, self.releases, self.cancelled = 0, [], [] + + def request_pin(self, keys): + request = PinRequestId(self.requests) + self.requests += 1 + self.pending[request] = (time.monotonic() + self.delay, tuple(keys)) + return request + + def poll_pin_results(self): + ready = [] + for request, (deadline, keys) in list(self.pending.items()): + if time.monotonic() < deadline: + continue + del self.pending[request] + result = ( + None + if self.fail + else ( + f"pin-{request}", + { + key: [ + MemoryRef( + end_point_name=self.name, + element_index=self.keys.index(key), + ) + ] + for key in keys + }, + ) + ) + ready.append((request, result)) + return ready + + def release_pin(self, handle): + self.releases.append(handle) + return True + + def cancel_pin_request(self, request): + self.cancelled.append(request) + self.pending.pop(request, None) + + +def run(blocks, delay_ms, scenario): + size = 4096 + keys = tuple(BlockKey(f"block-{i}".encode()) for i in range(blocks)) + expected = b"".join(bytes([i % 255 + 1]) * size for i in range(blocks)) + source_memory = ctypes.create_string_buffer(expected, len(expected)) + target_memory = ctypes.create_string_buffer(len(expected)) + source_framework = Framework( + "nvtx-source", + keys, + 20.0 if scenario == "timeout" else delay_ms / 1000, + scenario == "failure", + ) + target_framework = Framework("nvtx-target", keys, 0, False) + channels = [control_channel(), control_channel()] + with ExitStack() as stack: + workers = [] + for name, memory, framework, control in zip( + ("nvtx-source", "nvtx-target"), + (source_memory, target_memory), + (source_framework, target_framework), + channels, + ): + worker = KVCR( + KVCRConfig( + nixl_agent_name=name, + nixl_listen_port=0, + pool_layouts=[("", size)], + operation_timeout_ms=10_000, + abandon_timeout_ms=20_000, + ), + KVCRBindings( + framework.request_pin, + framework.poll_pin_results, + framework.release_pin, + cancel_pin_request=framework.cancel_pin_request, + framework_control=control, + ), + KVCRBackendConfigs( + framework_regions=[ + RegionDescriptor( + addr=ctypes.addressof(memory), + size=size, + count=blocks, + ) + ], + remote_fw_dram=RemoteFWDramOptions(eager_ctrl_connect=False), + ), + ) + stack.callback(worker.close) + workers.append(worker) + source, target = workers + target.submit_hint( + { + "protocol_version": "0.1", + "actions": [ + { + "action_type": "kv.fetch", + "action_version": "1.0", + "payload": { + "source_control_endpoint": channels[0].endpoint, + "block_hashes": list(range(blocks)), + }, + } + ], + }, + request_id="nvtx-pin-example", + ) + operation = target.deliver( + { + key: [MemoryRef(end_point_name="nvtx-target", element_index=i)] + for i, key in enumerate(keys) + }, + request_id="nvtx-pin-example", + ) + deadline = time.monotonic() + 30 + results = {} + while time.monotonic() < deadline: + source.poll_completed() + results.update(target.poll_completed()) + if operation in results and not source_framework.pending: + # Let source-side release follow target notification processing. + if scenario != "success" or source_framework.releases: + break + time.sleep(0.001) + assert operation in results, "delivery did not complete" + assert source_framework.requests == 1, "framework pin path was not exercised" + success = all(item.success for item in results[operation].values()) + assert success == (scenario == "success"), results + if success: + assert target_memory.raw == expected, "transferred bytes differ" + assert source_framework.releases == ["pin-0"], "pin was not released once" + if scenario == "timeout": + assert source_framework.cancelled == [0] + print( + json.dumps( + { + "scenario": scenario, + "blocks": blocks, + "bytes": len(expected), + "success": success, + "pin_requests": source_framework.requests, + "pin_releases": source_framework.releases, + "cancelled": source_framework.cancelled, + "operation": operation, + } + ) + ) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--blocks", type=int, default=2) + parser.add_argument("--delay-ms", type=float, default=20) + parser.add_argument( + "--scenario", choices=("success", "failure", "timeout"), default="success" + ) + args = parser.parse_args() + if args.blocks < 1 or args.delay_ms < 0: + parser.error("blocks must be positive and delay-ms nonnegative") + run(args.blocks, args.delay_ms, args.scenario) diff --git a/pyproject.toml b/pyproject.toml index e051190..7ace00e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -27,6 +27,9 @@ dependencies = [ # cursor operations. It is loaded on first use, so importing kvcr without it # works; running a KVCR-Service or claiming one of its pools does not. +[project.optional-dependencies] +profiling = ["nvtx>=0.2.16,<0.3", "numpy>=1.26,<3"] + [dependency-groups] dev = [ "pytest>=8,<9", diff --git a/src/kvcr/_nvtx.py b/src/kvcr/_nvtx.py new file mode 100644 index 0000000..816cda9 --- /dev/null +++ b/src/kvcr/_nvtx.py @@ -0,0 +1,235 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +"""Optional, best-effort NVTX annotations for framework pin lifetimes. + +Ranges cover a single synchronous callback. Marks carry an independent pin +identity so shared waits and reused framework request IDs remain distinguishable. +Neither profiler presence nor telemetry configuration controls payload work. +""" + +import logging +import os +from dataclasses import dataclass +from enum import IntEnum +from importlib import import_module +from itertools import count +from uuid import uuid4 + +_LOGGER = logging.getLogger(__name__) +_WARNED: set[str] = set() +_NAMES = ( + "source.pin.framework", + "source.pin.registered", + "source.pin.waiter", + "source.pin.completed", + "source.pin.detached", +) + + +class Status(IntEnum): + PENDING = 1 + SUCCESS = 2 + PARTIAL = 3 + FAILED = 4 + TIMEOUT = 5 + CANCELLED = 6 + + +class Reason(IntEnum): + UNKNOWN = 0 + NONE = 1 + CALLBACK_ERROR = 2 + DUPLICATE_REQUEST = 3 + INVALID_RESULT = 4 + DEADLINE = 5 + CANCELLED = 6 + SHUTDOWN = 7 + NO_WAITERS = 8 + + +def _warn_once(key: str, message: str) -> None: + if key not in _WARNED: + _WARNED.add(key) + _LOGGER.warning(message) + + +def _load_backend(): + nvtx = import_module("nvtx") + numpy = import_module("numpy") + return nvtx.get_domain("KVCR"), numpy + + +def pin_result_blocks(result): + """Keep diagnostic traversal of a framework-supplied mapping best effort.""" + try: + return sum(refs is not None for refs in result[1].values()) + except Exception: + return -1 + + +def create_tracer(): + """Return None before optional imports when off; absence is harmless.""" + configured = os.getenv("KVCR_NVTX_LEVEL") + level = configured if configured is not None else "low" + if level == "off": + return None + if level not in ("low", "medium"): + _warn_once("level", "Invalid KVCR_NVTX_LEVEL; expected off, low or medium.") + return None + try: + domain, numpy = _load_backend() + return _PinTracer(domain, numpy, level) + except Exception: + if configured is not None: + _warn_once( + "backend", + "KVCR NVTX unavailable; install the profiling extra to enable it.", + ) + return None + + +class _PinTracer: + def __init__(self, domain, numpy, level): + self.domain, self.numpy, self.level = domain, numpy, level + instance = uuid4().int + self.instance_hi, self.instance_lo = instance >> 64, instance & (2**64 - 1) + self._ids = count(1) + self.category = domain.get_category_id("framework_pin") + for name in _NAMES: + domain.get_registered_string(name) + self.dtype = numpy.dtype( + [ + ("schema_version", "u2"), + ("instance_hi", "u8"), + ("instance_lo", "u8"), + ("pin_id", "u8"), + ("pin_request_known", "u1"), + ("pin_request_id", "i8"), + ("source_op_id", "i8"), + ("op_handle", "i8"), + ("requested_blocks", "u8"), + ("completed_blocks", "i8"), + ("status", "u1"), + ("reason", "u1"), + ("fw_dram_utilization_known", "u1"), + ] + ) + + def begin(self, block_count, *, source_op_id=0, op_handle=0): + return _PinTrace(self, next(self._ids), block_count, source_op_id, op_handle) + + +@dataclass +class _PinTrace: + tracer: _PinTracer + pin_id: int + requested_blocks: int + source_op_id: int + op_handle: int + request_id: int | None = None + finished: bool = False + + def _attributes( + self, + name, + status=Status.PENDING, + reason=Reason.NONE, + completed_blocks=-1, + source_op_id=None, + op_handle=None, + ): + tracer = self.tracer + # Framework IDs are unrestricted Python ints. An unrepresentable ID must + # not suppress the independent pin identity or terminal status event. + request_known = ( + self.request_id is not None and -(2**63) <= self.request_id < 2**63 + ) + # Each event owns its buffer and attributes. Never mutate shared payloads. + payload = tracer.numpy.array( + ( + 1, + tracer.instance_hi, + tracer.instance_lo, + self.pin_id, + request_known, + self.request_id if request_known else 0, + self.source_op_id if source_op_id is None else source_op_id, + self.op_handle if op_handle is None else op_handle, + self.requested_blocks, + completed_blocks, + status, + reason, + 0, + ), + dtype=tracer.dtype, + ) + return tracer.domain.get_event_attributes( + message=name, + category=tracer.category, + payload=payload, + ) + + def _mark(self, name, **fields): + try: + self.tracer.domain.mark(self._attributes(name, **fields)) + except Exception: + pass # Annotation errors never alter KVCR results or resources. + + def push(self): + try: + self.tracer.domain.push_range(self._attributes("source.pin.framework")) + return True + except Exception: + return False + + def pop(self, pushed): + if pushed: + try: + self.tracer.domain.pop_range() + except Exception: + pass + + def registered(self, request_id): + self.request_id = request_id + self._mark("source.pin.registered") + + def waiter(self, source_op_id, op_handle): + self._mark("source.pin.waiter", source_op_id=source_op_id, op_handle=op_handle) + + def detached(self, source_op_id, op_handle): + if self.tracer.level == "medium": + self._mark( + "source.pin.detached", + source_op_id=source_op_id, + op_handle=op_handle, + ) + + def finish(self, result, *, reason=None, completed_blocks=-1): + if self.finished: + return + self.finished = True + status = { + "success": Status.SUCCESS, + "failed": Status.FAILED, + "timeout": Status.TIMEOUT, + "cancelled": Status.CANCELLED, + }[result] + if status is Status.SUCCESS and 0 <= completed_blocks < self.requested_blocks: + status = Status.PARTIAL + if reason is None: + reason = { + "success": Reason.NONE, + "failed": Reason.UNKNOWN, + "timeout": Reason.DEADLINE, + "cancelled": Reason.CANCELLED, + }[result] + elif isinstance(reason, str): + reason = Reason[reason.upper()] + if status is Status.PARTIAL and reason is Reason.NONE: + reason = Reason.UNKNOWN + self._mark( + "source.pin.completed", + status=status, + reason=reason, + completed_blocks=completed_blocks, + ) diff --git a/src/kvcr/remote_fw_dram.py b/src/kvcr/remote_fw_dram.py index 2ed5013..36039dd 100644 --- a/src/kvcr/remote_fw_dram.py +++ b/src/kvcr/remote_fw_dram.py @@ -22,6 +22,7 @@ import msgspec +from . import _nvtx from .config import KeyAdapter, RemoteFWDramOptions from .core import ( DURATION_METRIC, @@ -490,6 +491,7 @@ class _PendingPinWait: keys: tuple[BlockKey, ...] started_at: float | None op_ids: set[_OpId] = field(default_factory=set) + trace: _nvtx._PinTrace | None = None class _RemoteFWDram: @@ -508,6 +510,7 @@ def __init__( self._kvcr = kvcr self._options = options self._key_adapter = key_adapter + self._nvtx = _nvtx.create_tracer() # Main-thread state: request hints, framework pins, and progress state. self._closed = False @@ -1415,6 +1418,8 @@ def _process_pending_pin_results(self) -> None: and request in op.pending_pin_ids ] if not ops: + if wait is not None and wait.trace is not None: + wait.trace.finish("cancelled", reason=_nvtx.Reason.NO_WAITERS) self._discard_pin_result(result, wait.keys if wait is not None else ()) continue @@ -1433,7 +1438,20 @@ def _process_pending_pin_results(self) -> None: if result is not None and wait is not None: pin_handle = self._install_framework_pin(wait.keys, result) self._record_pending_pin_wait( - wait, "success" if pin_handle is not None else "failed" + wait, + "success" if pin_handle is not None else "failed", + reason=( + _nvtx.Reason.NONE + if pin_handle is not None + else _nvtx.Reason.UNKNOWN + if result is None + else _nvtx.Reason.INVALID_RESULT + ), + completed_blocks=( + _nvtx.pin_result_blocks(result) + if pin_handle is not None and wait.trace is not None + else -1 + ), ) if pin_handle is not None: for _, op in active_ops: @@ -1503,7 +1521,7 @@ def _resume_source_pin(self, op_id: _OpId, op: _SourcePinOp) -> None: ): return op.framework_acquire_attempted = True - framework_sources = self._acquire_framework_sources(unresolved_keys) + framework_sources = self._acquire_framework_sources(unresolved_keys, op=op) if isinstance(framework_sources, _PendingFrameworkSources): op.framework_pins.update(framework_sources.framework_pins) for request in framework_sources.pending_pins: @@ -1533,10 +1551,13 @@ def _register_pending_pin(self, request: PinRequestId, op_id: _OpId) -> None: request, ) return + new_waiter = op_id not in wait.op_ids wait.op_ids.add(op_id) op = self._source_pin_ops.get(op_id) if op is not None: op.pending_pin_ids.add(request) + if new_waiter and wait.trace is not None: + wait.trace.waiter(op_id[1], op.op_handle) def _remove_pending_pin_state( self, request_id: PinRequestId @@ -1585,6 +1606,8 @@ def _cancel_pending_pin_for_op( if wait is None: continue wait.op_ids.discard(op_id) + if wait.trace is not None: + wait.trace.detached(op_id[1], op.op_handle) if not wait.op_ids: wait = self._remove_pending_pin_state(request_id) if wait is not None: @@ -1602,9 +1625,20 @@ def _cancel_pending_pin(self, request: PinRequestId) -> None: ) def _record_pending_pin_wait( - self, wait: _PendingPinWait | None, result: str + self, + wait: _PendingPinWait | None, + result: str, + *, + reason: _nvtx.Reason | None = None, + completed_blocks: int = -1, ) -> None: if wait is not None: + if wait.trace is not None: + wait.trace.finish( + result, + reason=_nvtx.Reason.SHUTDOWN if self._closed else reason, + completed_blocks=completed_blocks, + ) self._kvcr._record_duration("framework_pin_wait", wait.started_at, result) def _discard_pin_result( @@ -1625,25 +1659,53 @@ def _discard_pin_result( # Framework pin ownership. - def _pin_framework_keys(self, keys: Collection[BlockKey]) -> PinRequestId | None: + def _pin_framework_keys( + self, + keys: Collection[BlockKey], + *, + op: _SourcePinOp | None = None, + ) -> PinRequestId | None: kvcr = self._kvcr if not keys: return None keys = tuple(keys) started_at = kvcr._timer() result = "failed" + trace = None + if self._nvtx is not None: + try: + trace = self._nvtx.begin( + len(keys), + source_op_id=op.op_id[1] if op is not None else 0, + op_handle=op.op_handle if op is not None else 0, + ) + except Exception: + pass try: - request = kvcr._request_pin_callback(keys) + pushed = trace.push() if trace is not None else False + try: + request = kvcr._request_pin_callback(keys) + finally: + if trace is not None: + trace.pop(pushed) + if trace is not None: + trace.request_id = request if request in self._pending_pin_ops: + if trace is not None: + trace.finish("failed", reason=_nvtx.Reason.DUPLICATE_REQUEST) logger.warning("KVCR reused pin request id %d", request) return None - wait = _PendingPinWait(request, keys, kvcr._timer()) + wait = _PendingPinWait(request, keys, kvcr._timer(), trace=trace) self._pending_pin_ops[request] = wait for key in keys: self._pending_pin_keys.setdefault(key, set()).add(request) result = "pending" + if trace is not None: + trace.registered(request) return request except Exception: + if trace is not None: + trace.finish("failed", reason=_nvtx.Reason.CALLBACK_ERROR) return None finally: kvcr._record_duration("source_acquire", started_at, result) @@ -1688,6 +1750,8 @@ def _install_framework_pin( def _acquire_framework_sources( self, keys: tuple[BlockKey, ...], + *, + op: _SourcePinOp | None = None, ) -> ( tuple[dict[BlockKey, list[_TransferRef]], set[PinHandle]] | _PendingFrameworkSources @@ -1713,7 +1777,7 @@ def _acquire_framework_sources( } pending_pins, covered_keys = self._find_pending_pins(keys_to_pin) uncovered_keys = [key for key in keys_to_pin if key not in covered_keys] - pin_request = self._pin_framework_keys(uncovered_keys) + pin_request = self._pin_framework_keys(uncovered_keys, op=op) if pin_request is not None: pending_pins.append(pin_request) if pending_pins: diff --git a/tests/unit/test_nvtx.py b/tests/unit/test_nvtx.py new file mode 100644 index 0000000..94b8263 --- /dev/null +++ b/tests/unit/test_nvtx.py @@ -0,0 +1,355 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +"""NVTX must describe pin ownership without changing it.""" + +import threading +from concurrent.futures import ThreadPoolExecutor +from types import SimpleNamespace + +import msgspec +import pytest +from _kvcr_test_utils import ( + FakeBytesControl, + FakeNixlAgent, + PendingPrimaryPinning, + _new_kvcr, + _poll_until, + _start_write_message, +) + +from kvcr import _nvtx +from kvcr.types import BlockKey + + +class RecordingDomain: + def __init__(self): + self.events = [] + self.registered = set() + self.stack = [] + self.fail = None + + def get_registered_string(self, message): + self.registered.add(message) + return message + + def get_category_id(self, category): + return 1 + + def get_event_attributes(self, **kwargs): + if self.fail == "attributes": + raise RuntimeError("annotation failure") + return SimpleNamespace(**kwargs) + + def _record(self, kind, attrs): + if self.fail == kind: + raise RuntimeError("annotation failure") + self.events.append((kind, attrs.message, attrs.payload, threading.get_ident())) + + def mark(self, attrs): + self._record("mark", attrs) + + def push_range(self, attrs): + self._record("push", attrs) + self.stack.append(threading.get_ident()) + + def pop_range(self): + assert self.stack.pop() == threading.get_ident() + if self.fail == "pop": + raise RuntimeError("annotation failure") + + +@pytest.fixture +def recording(monkeypatch): + np = pytest.importorskip("numpy") + domain = RecordingDomain() + monkeypatch.setenv("KVCR_NVTX_LEVEL", "low") + monkeypatch.setattr(_nvtx, "_load_backend", lambda: (domain, np)) + return domain + + +def payloads(domain, name): + return [p for _, message, p, _ in domain.events if message == name] + + +def make_source(domain, handles=(9,), *, request_id=0): + agent = FakeNixlAgent(metadata=b"source-md") + pinning, control = PendingPrimaryPinning(), FakeBytesControl() + pinning._next_request_id = request_id + source = _new_kvcr(agent, pinning, control, name="source") + for handle in handles: + control.incoming.append(_start_write_message(handle, BlockKey(b"k0"))) + _poll_until( + source, + lambda _: len(source._core._remote_fw_dram._source_pin_ops) == len(handles), + ) + return source, agent, pinning + + +def test_off_does_not_load_optional_dependencies(monkeypatch): + monkeypatch.setenv("KVCR_NVTX_LEVEL", "off") + + def forbidden(): + pytest.fail("off mode imported the optional backend") + + monkeypatch.setattr(_nvtx, "_load_backend", forbidden) + assert _nvtx.create_tracer() is None + + +@pytest.mark.parametrize("level", [None, "low", "medium", "invalid"]) +def test_missing_binding_is_safe_and_explicit_request_warns_once( + monkeypatch, caplog, level +): + _nvtx._WARNED.clear() + if level is None: + monkeypatch.delenv("KVCR_NVTX_LEVEL", raising=False) + else: + monkeypatch.setenv("KVCR_NVTX_LEVEL", level) + + def absent(): + raise ImportError("no nvtx") + + monkeypatch.setattr(_nvtx, "_load_backend", absent) + assert _nvtx.create_tracer() is None + assert _nvtx.create_tracer() is None + assert len(caplog.records) == (0 if level is None else 1) + + +def test_default_low_and_fresh_payloads(recording, monkeypatch): + monkeypatch.delenv("KVCR_NVTX_LEVEL") + tracer = _nvtx.create_tracer() + first = tracer.begin(2, source_op_id=3, op_handle=-(2**60)) + second = tracer.begin(5, source_op_id=4, op_handle=2**60) + first.registered(0) + second.registered(0) + first.finish("success", completed_blocks=2) + second.finish("failed", reason="unknown") + first.finish("cancelled") + registered = payloads(recording, "source.pin.registered") + completed = payloads(recording, "source.pin.completed") + assert [int(p["requested_blocks"]) for p in registered] == [2, 5] + assert [int(p["op_handle"]) for p in registered] == [-(2**60), 2**60] + assert registered[0]["pin_id"] != registered[1]["pin_id"] + assert len(completed) == 2 + assert all(int(p["fw_dram_utilization_known"]) == 0 for p in completed) + assert not recording.stack + + +@pytest.mark.parametrize("failure", ["attributes", "push", "mark", "pop"]) +def test_annotation_failures_preserve_pin_callback_and_release(recording, failure): + recording.fail = failure + source, agent, pinning = make_source(recording) + assert pinning.searches == [(BlockKey(b"k0"),)] + agent.state = "DONE" + pinning.complete(0) + _poll_until(source, lambda _: pinning.unpins == ["pin"]) + assert not recording.stack + + +def test_shared_pin_has_one_lifetime_and_two_waiter_associations(recording): + source, agent, pinning = make_source(recording, (9, 10)) + assert len(pinning.searches) == 1 + agent.state = "DONE" + pinning.complete(0) + _poll_until(source, lambda _: pinning.unpins == ["pin"]) + registered = payloads(recording, "source.pin.registered") + waiters = payloads(recording, "source.pin.waiter") + completed = payloads(recording, "source.pin.completed") + assert len(registered) == len(completed) == 1 + assert {int(p["op_handle"]) for p in waiters} == {9, 10} + assert {int(p["pin_id"]) for p in waiters} == {int(registered[0]["pin_id"])} + assert int(completed[0]["completed_blocks"]) == 1 + assert int(completed[0]["status"]) == int(_nvtx.Status.SUCCESS) + assert not recording.stack + + +@pytest.mark.parametrize("finish", ["success", "failed", "timeout", "cancelled"]) +def test_terminal_pin_event_is_not_duplicated_by_late_result(recording, finish): + source, agent, pinning = make_source(recording) + backend = source._core._remote_fw_dram + agent.state = "DONE" + if finish == "success": + pinning.complete(0) + _poll_until(source, lambda _: pinning.unpins == ["pin"]) + elif finish == "failed": + pinning.completed.append((0, None)) + _poll_until(source, lambda _: not backend._pending_pin_ops) + elif finish == "timeout": + source._core._clock = lambda: float("inf") + _poll_until(source, lambda _: not backend._pending_pin_ops) + else: + source.close() + # The late result is discarded/released, never a second terminal event. + pinning.completed.append((0, None)) + backend._process_pending_pin_results() + completed = payloads(recording, "source.pin.completed") + assert len(completed) == 1 + assert int(completed[0]["status"]) == int( + { + "success": _nvtx.Status.SUCCESS, + "failed": _nvtx.Status.FAILED, + "timeout": _nvtx.Status.TIMEOUT, + "cancelled": _nvtx.Status.CANCELLED, + }[finish] + ) + if finish == "failed": + assert int(completed[0]["reason"]) == int(_nvtx.Reason.UNKNOWN) + assert not recording.stack + + +@pytest.mark.parametrize("level,expected", [("low", 0), ("medium", 1)]) +def test_detail_level_gates_waiter_detach(recording, monkeypatch, level, expected): + monkeypatch.setenv("KVCR_NVTX_LEVEL", level) + source, agent, _ = make_source(recording, (9, 10)) + backend = source._core._remote_fw_dram + op_id, op = next(iter(backend._source_pin_ops.items())) + backend._cancel_pending_pin_for_op(op_id, op) + assert len(payloads(recording, "source.pin.detached")) == expected + assert not payloads(recording, "source.pin.completed") + agent.state = "DONE" + + +def test_reused_framework_request_id_gets_new_trace_identity(recording): + source, agent, pinning = make_source(recording) + agent.state = "DONE" + pinning.complete(0) + _poll_until(source, lambda _: pinning.unpins == ["pin"]) + pinning._next_request_id = 0 + source._core.framework_control.incoming.append( + _start_write_message(10, BlockKey(b"k1")) + ) + _poll_until(source, lambda _: len(pinning.searches) == 2) + pinning.complete(0) + _poll_until(source, lambda _: len(pinning.unpins) == 2) + records = payloads(recording, "source.pin.registered") + assert [int(p["pin_request_id"]) for p in records] == [0, 0] + assert records[0]["pin_id"] != records[1]["pin_id"] + + +@pytest.mark.parametrize( + "request_id", [-(2**63) - 1, -(2**63), -1, 0, 2**63 - 1, 2**63, 2**128] +) +def test_framework_request_id_range_cannot_drop_lifecycle_events(recording, request_id): + source, agent, pinning = make_source(recording, request_id=request_id) + agent.state = "DONE" + pinning.complete(request_id) + _poll_until(source, lambda _: pinning.unpins == ["pin"]) + representable = -(2**63) <= request_id < 2**63 + identities = set() + for name in ("registered", "waiter", "completed"): + events = payloads(recording, f"source.pin.{name}") + assert len(events) == 1 + event = events[0] + assert bool(event["pin_request_known"]) == representable + assert int(event["pin_request_id"]) == (request_id if representable else 0) + identities.add( + tuple(int(event[f]) for f in ("instance_hi", "instance_lo", "pin_id")) + ) + assert len(identities) == 1 + assert int(payloads(recording, "source.pin.completed")[0]["status"]) == int( + _nvtx.Status.SUCCESS + ) + assert not recording.stack + + +def test_callback_failure_is_correlated_before_registration(recording): + source, agent, pinning = make_source(recording) + agent.state = "DONE" + pinning.complete(0) + _poll_until(source, lambda _: pinning.unpins == ["pin"]) + + def fails(keys): + raise RuntimeError("framework error is not a capacity diagnosis") + + source._core._request_pin_callback = fails + source._core.framework_control.incoming.append( + _start_write_message(11, BlockKey(b"other")) + ) + _poll_until(source, lambda _: len(payloads(recording, "source.pin.completed")) == 2) + event = payloads(recording, "source.pin.completed")[-1] + assert int(event["op_handle"]) == 11 + assert int(event["source_op_id"]) > 0 + assert not int(event["pin_request_known"]) + assert int(event["reason"]) == int(_nvtx.Reason.CALLBACK_ERROR) + assert not recording.stack + + +def test_partial_result_counts_available_blocks_without_guessing_cause(recording): + agent, pinning, control = ( + FakeNixlAgent(), + PendingPrimaryPinning(), + FakeBytesControl(), + ) + source = _new_kvcr(agent, pinning, control, name="source") + message = msgspec.msgpack.decode(_start_write_message(9, BlockKey(b"k0"))) + message["keys"].append(b"k1") + message["dst_descriptors"] *= 2 + control.incoming.append(msgspec.msgpack.encode(message)) + _poll_until(source, lambda _: bool(pinning.pending)) + agent.state = "DONE" + pinning.complete(0, missing_indices=(1,)) + _poll_until(source, lambda _: pinning.unpins == ["pin"]) + event = payloads(recording, "source.pin.completed")[0] + assert int(event["requested_blocks"]) == 2 + assert int(event["completed_blocks"]) == 1 + assert int(event["status"]) == int(_nvtx.Status.PARTIAL) + assert int(event["reason"]) == int(_nvtx.Reason.UNKNOWN) + + +def test_invalid_result_has_bounded_reason_and_releases_handle(recording): + source, agent, pinning = make_source(recording) + agent.state = "DONE" + pinning.completed.append((0, ("pin", {}))) + _poll_until(source, lambda _: pinning.unpins == ["pin"]) + event = payloads(recording, "source.pin.completed")[0] + assert int(event["status"]) == int(_nvtx.Status.FAILED) + assert int(event["reason"]) == int(_nvtx.Reason.INVALID_RESULT) + + +def test_payload_counting_failure_cannot_interrupt_pin_processing(recording): + class ValuesUnavailable(dict): + def values(self): + raise RuntimeError("diagnostic-only traversal failed") + + source, agent, pinning = make_source(recording) + agent.state = "DONE" + pinning.complete(0) + request, (handle, mapping) = pinning.completed.pop() + pinning.completed.append((request, (handle, ValuesUnavailable(mapping)))) + _poll_until(source, lambda _: pinning.unpins == ["pin"]) + event = payloads(recording, "source.pin.completed")[0] + assert int(event["status"]) == int(_nvtx.Status.SUCCESS) + assert int(event["completed_blocks"]) == -1 + + +def test_off_records_no_pin_state_or_events(recording, monkeypatch): + monkeypatch.setenv("KVCR_NVTX_LEVEL", "off") + source, agent, pinning = make_source(recording) + wait = next(iter(source._core._remote_fw_dram._pending_pin_ops.values())) + assert wait.trace is None + agent.state = "DONE" + pinning.complete(0) + _poll_until(source, lambda _: pinning.unpins == ["pin"]) + assert recording.events == [] + + +def test_concurrent_annotations_have_independent_identities_and_payloads(recording): + tracer = _nvtx.create_tracer() + + def annotate(value): + trace = tracer.begin(value, op_handle=value) + trace.registered(value) + trace.finish("success", completed_blocks=value) + + with ThreadPoolExecutor(max_workers=4) as pool: + list(pool.map(annotate, range(1, 33))) + events = payloads(recording, "source.pin.completed") + assert len({int(p["pin_id"]) for p in events}) == 32 + assert {int(p["pin_request_id"]) for p in events} == set(range(1, 33)) + assert all(p["completed_blocks"] == p["requested_blocks"] for p in events) + assert recording.registered == { + "source.pin.framework", + "source.pin.registered", + "source.pin.waiter", + "source.pin.completed", + "source.pin.detached", + } diff --git a/uv.lock b/uv.lock index b74021b..b34ca0a 100644 --- a/uv.lock +++ b/uv.lock @@ -266,6 +266,14 @@ dependencies = [ { name = "pyzmq" }, ] +[package.optional-dependencies] +profiling = [ + { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, + { name = "numpy", version = "2.5.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "nvtx" }, +] + [package.dev-dependencies] dev = [ { name = "pytest" }, @@ -276,8 +284,11 @@ dev = [ requires-dist = [ { name = "msgspec", specifier = ">=0.21.0,<1" }, { name = "nixl", specifier = "==1.5.0" }, + { name = "numpy", marker = "extra == 'profiling'", specifier = ">=1.26,<3" }, + { name = "nvtx", marker = "extra == 'profiling'", specifier = ">=0.2.16,<0.3" }, { name = "pyzmq", specifier = ">=25.0.0,<28" }, ] +provides-extras = ["profiling"] [package.metadata.requires-dev] dev = [ @@ -893,6 +904,38 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/a8/64/3708a90d1ebe202ffdeb7185f878a3c84d15c2b2c31858da2ce0583e2def/nvidia_nvtx-13.0.85-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:cb7780edb6b14107373c835bf8b72e7a178bac7367e23da7acb108f973f157a6", size = 148878, upload-time = "2025-09-04T08:28:53.627Z" }, ] +[[package]] +name = "nvtx" +version = "0.2.16" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/7d/0b/7e9571ea963fd54bc0ac05184e735df7c4672c1a93cbb50d4fe3e7fa9f99/nvtx-0.2.16.tar.gz", hash = "sha256:252fdb870308be4a2beecb3ed174b1d59099eeb63c883f269944a765ce5f7720", size = 193099, upload-time = "2026-08-12T17:45:47.018Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/23/81/72d4fc2eaf32f851a815f467dcf42aac1595b18c181654d38f4c99a83d9d/nvtx-0.2.16-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:4592470f9dd67a23ee8523850a6b98fed4f220337a21605809c4b9a20daec726", size = 2728969, upload-time = "2026-08-12T17:56:22.208Z" }, + { url = "https://files.pythonhosted.org/packages/bd/df/9c5620c91234689085740b6bf396a678ce495070b41571e7b562eae90975/nvtx-0.2.16-cp310-cp310-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:23f30fcaf68f53d1895282315cb35aed5f605d59aeb33e75e276545ff95c4af6", size = 2754594, upload-time = "2026-08-12T17:56:01.323Z" }, + { url = "https://files.pythonhosted.org/packages/c2/cc/41d48962330d184cf1f0e713e9b83cab8305e09f89f135587be6a3c69400/nvtx-0.2.16-cp310-cp310-win_amd64.whl", hash = "sha256:ae4f45566ea3754f8599241a4537c8bf5fb5a1f0044878c9edca4185a0ef1fcf", size = 406498, upload-time = "2026-08-12T17:55:40.925Z" }, + { url = "https://files.pythonhosted.org/packages/db/c0/82bde35c60663a89efb45507045fc1b18a09f9ea1cbf2c6d91d6e848df66/nvtx-0.2.16-cp310-cp310-win_arm64.whl", hash = "sha256:cda9dc7019247b6ecf5fdce72e5fb6d2b70b82d2bf6e15441f344287e4174780", size = 377934, upload-time = "2026-08-12T17:55:20.392Z" }, + { url = "https://files.pythonhosted.org/packages/e6/37/1a7405737b07c40aa8756b028bf43f2674ba65d34d33cc5bdb32642cb3e7/nvtx-0.2.16-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:18bcb56547640c58ebc38c6674526df43d7f34afd288bdff92bf399fd54e2b04", size = 2855859, upload-time = "2026-08-12T17:54:19.085Z" }, + { url = "https://files.pythonhosted.org/packages/b8/ca/f615ad2a7c4b4f6c10dbcddb267a25878098305d630f1cba5b0020156011/nvtx-0.2.16-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f6d2922f5637b5f4c646599a066a770e7f26e154697dc217ad0aa8b831ea1338", size = 2886295, upload-time = "2026-08-12T17:54:47.46Z" }, + { url = "https://files.pythonhosted.org/packages/62/44/66947e41b79acc868770c30d1621e1b04b8b669e9ebd522340abef11b801/nvtx-0.2.16-cp311-cp311-win_amd64.whl", hash = "sha256:c5dbff0f10565da97b826a357b0835a0fd8008cb646e968ed5224168082ed917", size = 408408, upload-time = "2026-08-12T17:53:57.993Z" }, + { url = "https://files.pythonhosted.org/packages/70/85/62cb6881d71924ad891ca49370258de83055193c1c99361260f2d2802ceb/nvtx-0.2.16-cp311-cp311-win_arm64.whl", hash = "sha256:d4e814ecc617847d8baf32c87dde4d67f5097c4b5128866da2212326d9d2b541", size = 376873, upload-time = "2026-08-12T17:53:33.49Z" }, + { url = "https://files.pythonhosted.org/packages/3d/07/b33bef66e90c8f8b1861cb792cc4d0bcfc3c8640b0608a27c1e00262beb6/nvtx-0.2.16-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:4a94c321fc3fe489e40f82e2d7fe7db7ef312882e85228aaaa97ee3dfadc067c", size = 2792878, upload-time = "2026-08-12T17:53:12.123Z" }, + { url = "https://files.pythonhosted.org/packages/5f/09/adb32c8c6926d14c90eafc023d6a4c16ac22ca279f70c406bb9590b0e708/nvtx-0.2.16-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b6230be4ee35edcfe00411c6f71d48de5d50147146c4efc06b710d2516086031", size = 2849502, upload-time = "2026-08-12T17:52:52.966Z" }, + { url = "https://files.pythonhosted.org/packages/7d/23/c91834e518a1f627b08fa96e92b58083214aa826455379fe546bedb8a9e1/nvtx-0.2.16-cp312-cp312-win_amd64.whl", hash = "sha256:916cc45efa8794d3ad4c1f91e32110eb419f2d7669af6a10f47143e9e0844855", size = 398255, upload-time = "2026-08-12T17:52:28.869Z" }, + { url = "https://files.pythonhosted.org/packages/e5/ee/a3148b75ac4840f2c6bc7047c0a04340a87efb482212233ed5b83a39a26c/nvtx-0.2.16-cp312-cp312-win_arm64.whl", hash = "sha256:f70870ead7b6cf4227e4ae6ceba4e89a4bc3eecffbb62baca2c39a58d808df7e", size = 365146, upload-time = "2026-08-12T17:52:08.687Z" }, + { url = "https://files.pythonhosted.org/packages/61/6f/f25e14e910206552248e427f4182893fee6c320e0fe0f7f1740b5f661d4c/nvtx-0.2.16-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:a5dbb81e4098382bd2209a15dc8a8a53ff3db309e1ad160ff2fd5c18a94c3cfd", size = 2759374, upload-time = "2026-08-12T17:51:46.373Z" }, + { url = "https://files.pythonhosted.org/packages/2b/e0/848bc4f7f8df6da5f4497bd1adaf6e11b297177a6b131b550487ae0101dd/nvtx-0.2.16-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:624bf44efe550528e210ccb9fae2f53512f1f22a2f2629eca08b1c1e020e13ea", size = 2807408, upload-time = "2026-08-12T17:51:20.642Z" }, + { url = "https://files.pythonhosted.org/packages/58/f7/ee312f604f8d2af81521b12be4a54125522f64ea587254ffc2ced28b5493/nvtx-0.2.16-cp313-cp313-win_amd64.whl", hash = "sha256:92a6ba76504967d5686f3c02bb980af22567ad2d58f50c934fc49f0799f5193e", size = 397215, upload-time = "2026-08-12T17:50:37.329Z" }, + { url = "https://files.pythonhosted.org/packages/17/59/0480c91d129c47f5a2354b41119395bcb092e4cf6d372440777a25e4541d/nvtx-0.2.16-cp313-cp313-win_arm64.whl", hash = "sha256:c0aebfba86aaff2775518549fb4e4b8f96d3626091b31e1679e77af73fcbbb00", size = 363644, upload-time = "2026-08-12T17:50:13.259Z" }, + { url = "https://files.pythonhosted.org/packages/53/4c/819e3fdbb405db297a15e06679f6b2b0a9a0420fea14dd562adc8eaa5ec9/nvtx-0.2.16-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:95a0da1f7ba534dcf144f95f8c98ad6e4a639ca239a665b96fd6262790a90f1a", size = 2745011, upload-time = "2026-08-12T17:50:57.64Z" }, + { url = "https://files.pythonhosted.org/packages/04/fb/8d570c9cd23eaacebd29064d4f883f3548db90594c21d0e7fc2ea4970112/nvtx-0.2.16-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f1ce0c2509a6136ceaaebf73bb4b4d20fb741a2852c5c63ae4b6c1e5a5bb8fac", size = 2773878, upload-time = "2026-08-12T17:49:33.528Z" }, + { url = "https://files.pythonhosted.org/packages/0e/44/8a82ff02b52c6c28f69f60936ada7dc3d73496fbf1fbfc1f5dc43af52be7/nvtx-0.2.16-cp314-cp314-win_amd64.whl", hash = "sha256:e74ddf31c16c76ab000fbbe7e5a58cd88ade578a4de13bb012446580e961a4cb", size = 406097, upload-time = "2026-08-12T17:49:53.397Z" }, + { url = "https://files.pythonhosted.org/packages/79/db/c8c21be61328a4f64edafa15ebceaf18689aad134f1cdab6ee9780a79e03/nvtx-0.2.16-cp314-cp314-win_arm64.whl", hash = "sha256:48cb7700eea0d8fd588c95bb3023c5eb46ec0ca7936d4acba843bfce91807880", size = 373991, upload-time = "2026-08-12T17:49:05.301Z" }, + { url = "https://files.pythonhosted.org/packages/cc/a9/b7bda4f8326f1e96b89f2128612e09f3b16e5f1814df4207f0d08808947e/nvtx-0.2.16-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:13bb22db325321a0b383d156d84d22851127b3ffb686b54c2fabe0a7d182d53b", size = 2841353, upload-time = "2026-08-12T17:48:22.055Z" }, + { url = "https://files.pythonhosted.org/packages/46/89/28ee9f1904df73197f2296e1ab39e33e8a9993f56018a9b5b2e5294d290b/nvtx-0.2.16-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:154df7d6d5f47d2831275c6e2ad6fc9e917abe90bbc2513c501943fc4ec1dbff", size = 2742004, upload-time = "2026-08-12T17:48:47.964Z" }, + { url = "https://files.pythonhosted.org/packages/48/90/5ceec603d15acc132575025b19e64a85e91ba810386367f6fe0b39559a91/nvtx-0.2.16-cp314-cp314t-win_amd64.whl", hash = "sha256:9bfa46d2f98888813fc5dc17294cd657a073e407bfd29ce147ddfdbcf05fc4c5", size = 435586, upload-time = "2026-08-12T17:48:01.459Z" }, + { url = "https://files.pythonhosted.org/packages/57/2d/0fccbbc34dce15daed674aafaac761df99d84a1940ff416719069bd0e6d1/nvtx-0.2.16-cp314-cp314t-win_arm64.whl", hash = "sha256:e7c0be5e0784576adb495ac6c99a3a5ed6cf79158903d3d3b94a0f724d15e3bd", size = 396466, upload-time = "2026-08-12T17:46:52.532Z" }, +] + [[package]] name = "packaging" version = "26.3" From 895b2cd332b0bcb13eeaa97e07085080f0166ce0 Mon Sep 17 00:00:00 2001 From: aknvda Date: Mon, 5 Oct 2026 16:10:33 -0700 Subject: [PATCH 2/3] feat: trace remote source NIXL transfer lifecycle Signed-off-by: aknvda --- docs/nvtx-events.md | 45 +++++ src/kvcr/_nvtx.py | 269 ++++++++++++++++++++++++++++++ src/kvcr/progress.py | 128 +++++++++++++- src/kvcr/remote_fw_dram.py | 95 ++++++++++- tests/unit/test_nvtx_transfers.py | 181 ++++++++++++++++++++ 5 files changed, 713 insertions(+), 5 deletions(-) create mode 100644 docs/nvtx-events.md create mode 100644 tests/unit/test_nvtx_transfers.py diff --git a/docs/nvtx-events.md b/docs/nvtx-events.md new file mode 100644 index 0000000..94786da --- /dev/null +++ b/docs/nvtx-events.md @@ -0,0 +1,45 @@ +# Remote delivery NVTX events + +Install the `profiling` extra. `KVCR_NVTX_LEVEL=off|low|medium` selects +library-side detail; the default is low when the optional binding is available. +Annotations do not depend on whether Nsight is attached. See [profiling](profiling.md) +for capture commands and the pin schema. + +## Source transfers + +`nixl.write.submit` is a same-thread synchronous scope. `nixl.write.posted` +reports the native posting result, including rejection and ambiguity. +`nixl.done_observed` records the first observed native DONE, including synchronous +DONE returned by posting. It is distinct from `nixl.write.released` (successful +handle release) and `source.write.completed` (logical source operation result). +A prior error or cancellation can make the logical result fail even after DONE. +The status on the release event reflects the transfer outcome, not a release +failure. Release failures emit `nixl.write.release_retry` at medium detail. + +`nixl.write.error`, `source.write.cancel_requested`, and +`source.write.shutdown_unresolved` preserve the error, timeout/cancellation, and +unresolved ownership boundaries. Elapsed time does not prove DMA quiescence. +`source.write.refused` identifies a stale route before posting. No native numeric +error is inferred from exception text; `native_error_known=0` means unavailable. + +Schema 2 events use `(instance_hi, instance_lo, trace_id)` for a local lifecycle. +Join schema 1 pin/waiter records using `(instance_hi, instance_lo, source_op_id)`. +Across workers use `(target_agent_hi, target_agent_lo, +target_incarnation_hi, target_incarnation_lo, op_handle)`; incarnation is taken +from existing control metadata, without extending the wire protocol. Zero +identity fields mean unavailable. Handles remain signed, including local fills. +Transfer IDs are local to the source instance. Unknown counts/bytes are `-1`. + +`source_tier`/`destination_tier` describe memory ownership: 0 unknown, 1 framework, +2 KVCR-owned, 3 mixed. `source_memory`/`destination_memory` describe a verified +registration: 0 unknown, 1 DRAM, 2 VRAM, 3 FILE. Ownership alone does not identify +the physical memory tier; remote registration facts are left unknown when absent. + +Request identities use a stable 128-bit BLAKE2b digest of the complete UTF-8 +string. `request.context` carries a display prefix as an unsigned byte array, +explicit byte length (up to 256), and a truncation flag. Decode exactly `length` +bytes as UTF-8; a truncated prefix may end inside a character. Join by the digest, +never by the display prefix. Empty strings and unavailable IDs are distinct via +`request_known`. Request strings never become registered event names. Session, +parent-session and framework utilization remain explicitly unknown because the +current bindings do not provide them. diff --git a/src/kvcr/_nvtx.py b/src/kvcr/_nvtx.py index 816cda9..d59bb71 100644 --- a/src/kvcr/_nvtx.py +++ b/src/kvcr/_nvtx.py @@ -11,8 +11,10 @@ import os from dataclasses import dataclass from enum import IntEnum +from hashlib import blake2b from importlib import import_module from itertools import count +from types import MappingProxyType from uuid import uuid4 _LOGGER = logging.getLogger(__name__) @@ -33,6 +35,9 @@ class Status(IntEnum): FAILED = 4 TIMEOUT = 5 CANCELLED = 6 + REJECTED = 7 + AMBIGUOUS = 8 + UNRESOLVED = 9 class Reason(IntEnum): @@ -45,6 +50,16 @@ class Reason(IntEnum): CANCELLED = 6 SHUTDOWN = 7 NO_WAITERS = 8 + SUBMIT_REJECTED = 9 + SUBMIT_AMBIGUOUS = 10 + CREATE_ERROR = 11 + PROGRESS_ERROR = 12 + RELEASE_ERROR = 13 + CONTROL_ERROR = 14 + INVALID_NOTIFICATION = 15 + REMOTE_FAILURE = 16 + ROUTE_CHANGED = 17 + SOURCE_STALLED = 18 def _warn_once(key: str, message: str) -> None: @@ -95,6 +110,7 @@ def __init__(self, domain, numpy, level): self.instance_hi, self.instance_lo = instance >> 64, instance & (2**64 - 1) self._ids = count(1) self.category = domain.get_category_id("framework_pin") + self.lifecycle_category = domain.get_category_id("remote_deliver") for name in _NAMES: domain.get_registered_string(name) self.dtype = numpy.dtype( @@ -114,10 +130,263 @@ def __init__(self, domain, numpy, level): ("fw_dram_utilization_known", "u1"), ] ) + self.event_dtype = numpy.dtype( + [ + ("schema_version", "u2"), + ("instance_hi", "u8"), + ("instance_lo", "u8"), + ("trace_id", "u8"), + ("request_hi", "u8"), + ("request_lo", "u8"), + ("request_known", "u1"), + ("session_known", "u1"), + ("parent_session_known", "u1"), + ("op_handle", "i8"), + ("source_op_id", "i8"), + ("target_agent_hi", "u8"), + ("target_agent_lo", "u8"), + ("target_incarnation_hi", "u8"), + ("target_incarnation_lo", "u8"), + ("route_generation", "u8"), + ("transfer_id", "u8"), + ("requested_blocks", "i8"), + ("requested_bytes", "i8"), + ("selected_blocks", "i8"), + ("completed_blocks", "i8"), + ("selected_bytes", "i8"), + ("completed_bytes", "i8"), + ("status", "u1"), + ("reason", "u1"), + ("local_fill", "u1"), + ("source_tier", "u1"), + ("destination_tier", "u1"), + ("source_memory", "u1"), + ("destination_memory", "u1"), + ("native_error_known", "u1"), + ("native_error", "i8"), + ] + ) + # Nsight 2025.3 exports NumPy Unicode payloads as empty strings. A + # bounded byte array and explicit length also preserve embedded NULs. + self.context_dtype = numpy.dtype( + [ + ("schema_version", "u2"), + ("instance_hi", "u8"), + ("instance_lo", "u8"), + ("trace_id", "u8"), + ("request_hi", "u8"), + ("request_lo", "u8"), + ("request_known", "u1"), + ("length", "u2"), + ("truncated", "u1"), + ("value_utf8", "u1", (256,)), + ] + ) def begin(self, block_count, *, source_op_id=0, op_handle=0): return _PinTrace(self, next(self._ids), block_count, source_op_id, op_handle) + def name_progress_thread(self): + # Python 3.12 does not propagate Thread.name to Linux. Nsight displays + # the OS name; nvtx 0.2.16 has no Python thread-naming API. + try: + import ctypes + + ctypes.CDLL(None).prctl(15, ctypes.c_char_p(b"kvcr-progress"), 0, 0, 0) + except Exception: + pass + + def lifecycle( + self, *, target_agent=None, target_incarnation=None, request_id=None, **fields + ): + try: + hi, lo = identity(target_agent) + inc_hi, inc_lo = identity(target_incarnation) + req_hi, req_lo = identity(request_id) + trace = _LifecycleTrace( + self, + next(self._ids), + { + "target_agent_hi": hi, + "target_agent_lo": lo, + "target_incarnation_hi": inc_hi, + "target_incarnation_lo": inc_lo, + "request_hi": req_hi, + "request_lo": req_lo, + "request_known": request_id is not None, + **fields, + }, + ) + trace.context(request_id) + return trace + except Exception: + return None + + def source(self, op, kvcr): + """Prepare context once; diagnostics cannot interrupt a source write.""" + try: + refs = tuple(ref for block in op.src_descriptors for ref in block) + targets = tuple(ref for block in op.dst_descriptors for ref in block) + return self.lifecycle( + target_agent=op.route[0] or None, + target_incarnation=op.target_incarnation, + op_handle=op.op_handle, + source_op_id=op.op_id[1], + route_generation=op.route[1], + requested_blocks=op.requested_blocks, + selected_blocks=len(op.source_keys), + selected_bytes=kvcr._descriptor_bytes(refs), + source_tier=tier(refs), + destination_tier=tier(targets), + source_memory=memory_kind(refs, kvcr._memory_regions), + ) + except Exception: + return None + + +def identity(value): + """Stable 128-bit label identity; no Python hash randomization or truncation.""" + if value is None: + return 0, 0 + digest = blake2b(value.encode("utf-8"), digest_size=16).digest() + return int.from_bytes(digest[:8], "big"), int.from_bytes(digest[8:], "big") + + +def tier(refs): + """Storage ownership, deliberately independent of DRAM versus VRAM.""" + kinds = {1 if ref.framework else 2 for ref in refs} + return next(iter(kinds)) if len(kinds) == 1 else 3 if kinds else 0 + + +def memory_kind(refs, regions): + kinds = set() + for ref in refs: + table = regions[0 if ref.framework else 1] + region = table.get(ref.label) + if region is None: + region = table.get( + ref.label.partition(":")[0] + (":*" if ref.framework else "") + ) + kinds.add( + {"DRAM": 1, "VRAM": 2, "FILE": 3}.get(getattr(region, "mem_type", None), 0) + ) + return next(iter(kinds)) if len(kinds) == 1 else 0 + + +class _LifecycleTrace: + """Immutable context plus event bookkeeping owned by the operation lifecycle.""" + + def __init__(self, tracer, trace_id, fields): + self.tracer, self.trace_id = tracer, trace_id + self._fields = MappingProxyType(fields) + self._observed = set() + self.failure_reason = Reason.UNKNOWN + + def context(self, request_id): + try: + t = self.tracer + encoded = b"" if request_id is None else request_id.encode("utf-8") + payload = t.numpy.zeros((), dtype=t.context_dtype) + values = dict( + self._fields, + schema_version=2, + instance_hi=t.instance_hi, + instance_lo=t.instance_lo, + trace_id=self.trace_id, + length=min(len(encoded), 256), + truncated=len(encoded) > 256, + ) + for key in t.context_dtype.names: + if key in values: + payload[key] = values[key] + payload["value_utf8"][: min(len(encoded), 256)] = t.numpy.frombuffer( + encoded[:256], dtype="u1" + ) + t.domain.mark( + t.domain.get_event_attributes( + message="request.context", + category=t.lifecycle_category, + payload=payload, + ) + ) + except Exception: + pass + + def complete_source(self, success, blocks): + requested = self._fields.get("requested_blocks", -1) + self.mark( + "source.write.completed", + once=True, + status=(Status.PARTIAL if blocks < requested else Status.SUCCESS) + if success + else Status.FAILED, + reason=Reason.NONE + if success and blocks >= requested + else self.failure_reason, + completed_blocks=blocks if success else 0, + completed_bytes=self._fields.get("selected_bytes", -1) if success else 0, + ) + + def _attributes(self, name, fields): + t = self.tracer + values = { + "schema_version": 2, + "instance_hi": t.instance_hi, + "instance_lo": t.instance_lo, + "trace_id": self.trace_id, + "requested_blocks": -1, + "requested_bytes": -1, + "selected_blocks": -1, + "completed_blocks": -1, + "selected_bytes": -1, + "completed_bytes": -1, + "reason": Reason.NONE, + "status": Status.PENDING, + **self._fields, + **fields, + } + payload = t.numpy.array( + tuple(values.get(key, 0) for key in t.event_dtype.names), + dtype=t.event_dtype, + ) + return t.domain.get_event_attributes( + message=name, category=t.lifecycle_category, payload=payload + ) + + def mark(self, name, *, once=False, detail=False, **fields): + reason = fields.get("reason", Reason.NONE) + if reason not in ( + Reason.NONE, + Reason.UNKNOWN, + Reason.RELEASE_ERROR, + Reason.SHUTDOWN, + ): + self.failure_reason = reason + if detail and self.tracer.level != "medium": + return + if once: + if name in self._observed: + return + self._observed.add(name) + try: + self.tracer.domain.mark(self._attributes(name, fields)) + except Exception: + pass + + def push(self, name, **fields): + try: + self.tracer.domain.push_range(self._attributes(name, fields)) + return True + except Exception: + return False + + def pop(self, pushed): + if pushed: + try: + self.tracer.domain.pop_range() + except Exception: + pass + @dataclass class _PinTrace: diff --git a/src/kvcr/progress.py b/src/kvcr/progress.py index 4f92505..97aa716 100644 --- a/src/kvcr/progress.py +++ b/src/kvcr/progress.py @@ -18,6 +18,7 @@ import numpy as np from nixl import nixl_agent, nixl_agent_config +from . import _nvtx from .types import BlockKey, RegionDescriptor logger = logging.getLogger(__name__) @@ -87,6 +88,7 @@ class _TransferState: outcome: bool | None = None telemetry: Any | None = None next_release_log_at: float = 0.0 + trace: _nvtx._LifecycleTrace | None = None @dataclass @@ -209,6 +211,22 @@ def poll_transfer( if state.outcome is not False: logger.warning("NIXL transfer progress failed", exc_info=True) xfer_state = "ERR" + if state.trace is not None: + if xfer_state == "DONE": + state.trace.mark( + "nixl.done_observed", + once=True, + transfer_id=transfer_id, + status=_nvtx.Status.SUCCESS, + ) + elif xfer_state not in ("PROC", "PEND"): + state.trace.mark( + "nixl.write.error", + once=True, + transfer_id=transfer_id, + status=_nvtx.Status.FAILED, + reason=_nvtx.Reason.PROGRESS_ERROR, + ) # Releasing a pending NIXL/UCX handle can leave DMA running and lose # its completion signal. Remote writes retain it until actual DONE. if require_completion and xfer_state != "DONE": @@ -243,8 +261,46 @@ def submit_transfer( backend: str | None = None, notif_msg: bytes = b"", capture_telemetry: bool = False, + trace: _nvtx._LifecycleTrace | None = None, ) -> tuple[int, bool]: """Submit aligned local and remote descriptors to NIXL.""" + pushed = trace.push("nixl.write.submit") if trace is not None else False + try: + return self._submit_transfer( + operation, + local_descriptors, + remote_descriptors, + remote_side_agent=remote_side_agent, + backend=backend, + notif_msg=notif_msg, + capture_telemetry=capture_telemetry, + trace=trace, + ) + except Exception: + if trace is not None: + trace.mark( + "nixl.write.rejected", + once=True, + status=_nvtx.Status.REJECTED, + reason=_nvtx.Reason.CREATE_ERROR, + ) + raise + finally: + if trace is not None: + trace.pop(pushed) + + def _submit_transfer( + self, + operation, + local_descriptors, + remote_descriptors, + *, + remote_side_agent, + backend, + notif_msg, + capture_telemetry, + trace, + ): if not local_descriptors or not remote_descriptors: raise ValueError("NIXL transfer descriptors must be non-empty") if not isinstance(remote_side_agent, str) or not remote_side_agent: @@ -281,9 +337,12 @@ def submit_transfer( raise RuntimeError("NIXL transfer creation returned None") self._next_transfer_id += 1 transfer_id = self._next_transfer_id - state = _TransferState(handle, remote_side_agent, capture_telemetry) + state = _TransferState( + handle, remote_side_agent, capture_telemetry, trace=trace + ) self._active_transfers[transfer_id] = state submitted = True + post_status, post_reason = _nvtx.Status.PENDING, _nvtx.Reason.NONE try: post_state = agent.transfer(handle) if post_state == "DONE": @@ -291,12 +350,39 @@ def submit_transfer( elif post_state == "ERR": state.outcome = False submitted = False + post_status, post_reason = ( + _nvtx.Status.REJECTED, + _nvtx.Reason.SUBMIT_REJECTED, + ) elif post_state not in ("PROC", "PEND"): submitted = False + post_status, post_reason = ( + _nvtx.Status.AMBIGUOUS, + _nvtx.Reason.SUBMIT_AMBIGUOUS, + ) logger.warning("NIXL transfer returned unexpected state %r", post_state) except Exception: submitted = False + post_status, post_reason = ( + _nvtx.Status.AMBIGUOUS, + _nvtx.Reason.SUBMIT_AMBIGUOUS, + ) logger.warning("NIXL transfer submission was ambiguous", exc_info=True) + if trace is not None: + trace.mark( + "nixl.write.posted", + once=True, + transfer_id=transfer_id, + status=post_status, + reason=post_reason, + ) + if state.outcome is True: + trace.mark( + "nixl.done_observed", + once=True, + transfer_id=transfer_id, + status=_nvtx.Status.SUCCESS, + ) return transfer_id, submitted def cancel_transfer(self, transfer_id: int) -> bool: @@ -305,6 +391,14 @@ def cancel_transfer(self, transfer_id: int) -> bool: if state is None: return True state.outcome = False + if state.trace is not None: + state.trace.mark( + "nixl.write.cancel_requested", + once=True, + transfer_id=transfer_id, + status=_nvtx.Status.CANCELLED, + reason=_nvtx.Reason.CANCELLED, + ) return self._release_transfer(transfer_id, state) def _make_transfer_descriptors(self, descriptors: Sequence[_MemDescriptor]) -> Any: @@ -408,10 +502,26 @@ def _prepared_indices( def _release_transfer(self, transfer_id: int, state: _TransferState) -> bool: release_xfer = getattr(self.nixl_agent, "release_xfer_handle", None) if release_xfer is None: + if state.trace is not None: + state.trace.mark( + "nixl.write.release_retry", + detail=True, + transfer_id=transfer_id, + status=_nvtx.Status.UNRESOLVED, + reason=_nvtx.Reason.RELEASE_ERROR, + ) return False try: released = release_xfer(state.handle) is not False except Exception: + if state.trace is not None: + state.trace.mark( + "nixl.write.release_retry", + detail=True, + transfer_id=transfer_id, + status=_nvtx.Status.UNRESOLVED, + reason=_nvtx.Reason.RELEASE_ERROR, + ) now = time.monotonic() if now >= state.next_release_log_at: logger.warning( @@ -422,8 +532,24 @@ def _release_transfer(self, transfer_id: int, state: _TransferState) -> bool: state.next_release_log_at = now + _RELEASE_LOG_INTERVAL_SECONDS return False if not released: + if state.trace is not None: + state.trace.mark( + "nixl.write.release_retry", + detail=True, + transfer_id=transfer_id, + status=_nvtx.Status.UNRESOLVED, + reason=_nvtx.Reason.RELEASE_ERROR, + ) return False self._active_transfers.pop(transfer_id, None) + if state.trace is not None: + state.trace.mark( + "nixl.write.released", + once=True, + transfer_id=transfer_id, + status=_nvtx.Status.SUCCESS if state.outcome else _nvtx.Status.FAILED, + reason=_nvtx.Reason.NONE if state.outcome else _nvtx.Reason.UNKNOWN, + ) return True def start(self) -> None: diff --git a/src/kvcr/remote_fw_dram.py b/src/kvcr/remote_fw_dram.py index 36039dd..71fe95d 100644 --- a/src/kvcr/remote_fw_dram.py +++ b/src/kvcr/remote_fw_dram.py @@ -275,6 +275,7 @@ class _SourcePinOp(_Op): framework_pins: set[PinHandle] = field(default_factory=set) pending_pin_ids: set[PinRequestId] = field(default_factory=set) framework_acquire_attempted: bool = False + target_incarnation: str | None = None @dataclass @@ -301,24 +302,40 @@ class _SourceWriteOp(_RemoteOp): success: bool = False completed_indices: tuple[int, ...] = () route: tuple[str, int] = ("", 0) + requested_blocks: int = -1 + target_incarnation: str | None = None + trace: _nvtx._LifecycleTrace | None = field(default=None, repr=False, compare=False) + trace_initialized: bool = False def progress( self, progress: _KVCRProgress, _event: object | None ) -> tuple[bool, bool]: backend = self._backend + if not self.trace_initialized: + self.trace_initialized = True + if backend._nvtx is not None: + self.trace = backend._nvtx.source(self, backend._kvcr) + trace = self.trace observed_work = False write_id = (self.route[0], self.op_handle) status = backend._dangling_ops.source_writes[write_id] if self.transfer_id is None: - if ( - not backend._dangling_ops.check_source_progress() - or status.cancel_requested - ): + source_responsive = backend._dangling_ops.check_source_progress() + if not source_responsive or status.cancel_requested: self.state = _SourceWriteState.NOTIFY_FAILURE if ( self.state is _SourceWriteState.NOTIFY_FAILURE or backend._kvcr._clock() >= self.deadline ): + failure_reason = ( + _nvtx.Reason.CANCELLED + if status.cancel_requested + else _nvtx.Reason.SOURCE_STALLED + if not source_responsive + else _nvtx.Reason.UNKNOWN + if self.state is _SourceWriteState.NOTIFY_FAILURE + else _nvtx.Reason.DEADLINE + ) backend._send_write_done( progress, self.remote_agent, self.op_handle, False ) @@ -328,6 +345,15 @@ def progress( "source_write", self.started_at, "failed" ) backend._dangling_ops.finish_source(self) + if trace is not None: + trace.mark( + "source.write.completed", + once=True, + status=_nvtx.Status.FAILED, + reason=failure_reason, + completed_blocks=0, + completed_bytes=0, + ) return True, True if self.state is not _SourceWriteState.READY_TO_WRITE: raise RuntimeError(f"KVCR source operation {self.op_id!r} is not ready") @@ -340,6 +366,13 @@ def progress( # same handle back for a reused name. The target hears a # refusal instead of receiving the dead generation's bytes. self.state = _SourceWriteState.NOTIFY_FAILURE + if trace is not None: + trace.mark( + "source.write.refused", + once=True, + status=_nvtx.Status.REJECTED, + reason=_nvtx.Reason.ROUTE_CHANGED, + ) return False, True if not self.src_descriptors: backend._send_write_done( @@ -351,6 +384,19 @@ def progress( "source_write", self.started_at, "failed" ) backend._dangling_ops.finish_source(self) + if trace is not None: + trace.mark( + "source.write.completed", + once=True, + status=_nvtx.Status.PARTIAL + if self.requested_blocks + else _nvtx.Status.SUCCESS, + reason=_nvtx.Reason.UNKNOWN + if self.requested_blocks + else _nvtx.Reason.NONE, + completed_blocks=0, + completed_bytes=0, + ) return True, True submit_started_at = backend._kvcr._timer() status.submitted = True @@ -367,6 +413,7 @@ def progress( completed_indices=self.completed_indices, ), capture_telemetry=backend._telemetry_enabled, + trace=trace, ) self.transfer_id = transfer_id self.state = ( @@ -402,6 +449,15 @@ def progress( ) self.state = _SourceWriteState.FINISHED backend._dangling_ops.finish_source(self) + if trace is not None: + trace.mark( + "source.write.completed", + once=True, + status=_nvtx.Status.FAILED, + reason=_nvtx.Reason.CREATE_ERROR, + completed_blocks=0, + completed_bytes=0, + ) return True, True transfer_id = self.transfer_id @@ -411,6 +467,17 @@ def progress( status.cancel_requested or backend._kvcr._clock() >= self.deadline ): self.state = _SourceWriteState.CANCEL_PENDING + if trace is not None: + trace.mark( + "source.write.cancel_requested", + once=True, + status=_nvtx.Status.CANCELLED + if status.cancel_requested + else _nvtx.Status.TIMEOUT, + reason=_nvtx.Reason.CANCELLED + if status.cancel_requested + else _nvtx.Reason.DEADLINE, + ) observed_work = True cancelling = self.state is _SourceWriteState.CANCEL_PENDING transfer_result = backend._dangling_ops.poll_source( @@ -451,6 +518,8 @@ def progress( ) self.state = _SourceWriteState.FINISHED backend._dangling_ops.finish_source(self) + if trace is not None: + trace.complete_source(self.success, len(self.source_keys)) return True, True def close(self, progress: _KVCRProgress) -> bool: @@ -459,6 +528,13 @@ def close(self, progress: _KVCRProgress) -> bool: progress.poll_transfer(self.transfer_id, require_completion=True) is None ): + if self.trace is not None: + self.trace.mark( + "source.write.shutdown_unresolved", + once=True, + status=_nvtx.Status.UNRESOLVED, + reason=_nvtx.Reason.SHUTDOWN, + ) return False self.transfer_id = None self._backend._send_write_done( @@ -844,6 +920,8 @@ def _finish_target_pull(self, op: _TargetPullOp) -> None: # ------------------------------------------------------------------------- def initialize_progress(self, _progress: _KVCRProgress) -> None: + if self._nvtx is not None: + self._nvtx.name_progress_thread() initialize_control = getattr(self._control, "initialize", None) if initialize_control is not None: initialize_control() @@ -1201,6 +1279,11 @@ def _handle_start_write( dst_descriptors=dst_descriptors, allow_layout_subset=allow_layout_subset, route=(target_agent, self._route_generation.get(target_agent, 0)), + target_incarnation=( + payload.get("sender_incarnation") + if isinstance(payload.get("sender_incarnation"), str) + else None + ), ) if not self._try_local_source_write(progress, source_pin): self._progress_outbound.append(source_pin) @@ -1257,6 +1340,8 @@ def _try_local_source_write( dst_descriptors=source_pin.dst_descriptors, completed_indices=tuple(range(len(source_pin.ordered_keys))), route=source_pin.route, + requested_blocks=len(source_pin.ordered_keys), + target_incarnation=source_pin.target_incarnation, _backend=self, ) # Already on progress: retain cleanup ownership before starting the write. @@ -1354,6 +1439,8 @@ def _submit_prepared_source_write( source_pin.dst_descriptors[index] for index in completed_indices ), route=source_pin.route, + requested_blocks=len(source_pin.ordered_keys), + target_incarnation=source_pin.target_incarnation, _backend=self, framework_pins=framework_pins, source_keys=completed_keys, diff --git a/tests/unit/test_nvtx_transfers.py b/tests/unit/test_nvtx_transfers.py new file mode 100644 index 0000000..cc5fac4 --- /dev/null +++ b/tests/unit/test_nvtx_transfers.py @@ -0,0 +1,181 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +"""Observe NIXL state and resource release without changing their ordering.""" + +import pytest +import test_nvtx +from test_nvtx import payloads +from test_progress import _mem, _transfer_progress, _TransferAgent + +from kvcr import _nvtx + +recording = test_nvtx.recording + + +@pytest.mark.parametrize("value", [None, "", "req-α-😀\x00x", "😀" * 100]) +def test_context_mapping_preserves_identity_and_bounded_utf8(recording, value): + trace = _nvtx.create_tracer().lifecycle(request_id=value) + trace.mark("op.deliver") + context = payloads(recording, "request.context")[0] + event = payloads(recording, "op.deliver")[0] + assert int(context["request_known"]) == (value is not None) + assert (int(event["request_hi"]), int(event["request_lo"])) == _nvtx.identity(value) + encoded = b"" if value is None else value.encode("utf-8") + assert context["value_utf8"].tobytes()[: int(context["length"])] == encoded[:256] + assert bool(context["truncated"]) == (len(encoded) > 256) + assert int(event["session_known"]) == int(event["parent_session_known"]) == 0 + + +def test_source_deadline_and_unresolved_close_are_not_native_completion(recording): + source, agent, pinning = test_nvtx.make_source(recording) + agent.state = "PROC" + pinning.complete(0) + backend = source._core._remote_fw_dram + test_nvtx._poll_until( + source, lambda _: bool(payloads(recording, "nixl.write.posted")) + ) + source._core._clock = lambda: float("inf") + test_nvtx._poll_until( + source, lambda _: bool(payloads(recording, "source.write.cancel_requested")) + ) + assert not payloads(recording, "nixl.done_observed") + assert not payloads(recording, "source.write.completed") + op = next( + o + for o in source._core._progress._in_flight_ops.values() + if o.op_id[0] == "source" + ) + assert not op.close(source._core._progress) + assert len(payloads(recording, "source.write.shutdown_unresolved")) == 1 + agent.state = "DONE" + test_nvtx._poll_until( + source, lambda _: bool(payloads(recording, "source.write.completed")) + ) + final = payloads(recording, "source.write.completed")[0] + assert int(final["reason"]) == int(_nvtx.Reason.DEADLINE) + assert not backend._pending_pin_ops + + +def test_source_path_connects_pin_to_transfer(recording): + source, agent, pinning = test_nvtx.make_source(recording) + agent.state = "DONE" + pinning.complete(0) + test_nvtx._poll_until(source, lambda _: pinning.unpins == ["pin"]) + pin = payloads(recording, "source.pin.completed")[0] + submitted = payloads(recording, "nixl.write.posted") + assert len(submitted) == 1 + assert int(submitted[0]["source_op_id"]) == int(pin["source_op_id"]) + assert int(submitted[0]["instance_hi"]) == int(pin["instance_hi"]) + assert int(submitted[0]["op_handle"]) == 9 + assert int(submitted[0]["selected_blocks"]) == 1 + assert int(submitted[0]["selected_bytes"]) > 0 + assert len(payloads(recording, "source.write.completed")) == 1 + + +def setup_transfer(): + agent = _TransferAgent() + progress = _transfer_progress(agent) + trace = _nvtx.create_tracer().lifecycle( + op_handle=-9, source_op_id=2, target_agent="target-α", requested_blocks=2 + ) + return agent, progress, trace + + +def submit(progress, trace): + return progress.submit_transfer( + "WRITE", + [_mem(0)], + [_mem(1, owner="remote-agent")], + remote_side_agent="remote-agent", + trace=trace, + ) + + +@pytest.mark.parametrize("post", ["PROC", "DONE"]) +def test_native_done_and_release_are_distinct_once(recording, post): + agent, progress, trace = setup_transfer() + agent.transfer_result = post + transfer, accepted = submit(progress, trace) + assert accepted + assert len(payloads(recording, "nixl.done_observed")) == (post == "DONE") + agent.release_failures = 1 + assert progress.poll_transfer(transfer, require_completion=True) is None + assert transfer in progress._active_transfers + assert len(payloads(recording, "nixl.done_observed")) == 1 + assert not payloads(recording, "nixl.write.released") + assert progress.poll_transfer(transfer, require_completion=True) == (True, None) + assert transfer not in progress._active_transfers + done = payloads(recording, "nixl.done_observed")[0] + released = payloads(recording, "nixl.write.released")[0] + assert int(done["transfer_id"]) == int(released["transfer_id"]) == transfer + assert int(done["trace_id"]) == int(released["trace_id"]) + assert int(released["op_handle"]) == -9 + assert not recording.stack + + +@pytest.mark.parametrize("post", ["ERR", "unexpected", "exception"]) +def test_rejected_or_ambiguous_post_retains_native_handle(recording, post): + agent, progress, trace = setup_transfer() + agent.transfer_result = post + agent.submit_exception = post == "exception" + transfer, accepted = submit(progress, trace) + assert not accepted + posted = payloads(recording, "nixl.write.posted")[0] + assert int(posted["status"]) == int( + _nvtx.Status.REJECTED if post == "ERR" else _nvtx.Status.AMBIGUOUS + ) + agent.state = "PROC" + assert progress.poll_transfer(transfer, require_completion=True) is None + assert transfer in progress._active_transfers + assert not payloads(recording, "nixl.write.released") + agent.state = "DONE" + result = progress.poll_transfer(transfer, require_completion=True) + assert result == (post != "ERR", None) + assert len(payloads(recording, "nixl.done_observed")) == 1 + assert len(payloads(recording, "nixl.write.released")) == 1 + + +def test_poll_error_and_later_native_done_are_both_observed(recording): + agent, progress, trace = setup_transfer() + transfer, _ = submit(progress, trace) + agent.check_exception = True + for _ in range(3): + assert progress.poll_transfer(transfer, require_completion=True) is None + assert len(payloads(recording, "nixl.write.error")) == 1 + assert transfer in progress._active_transfers + agent.check_exception = False + assert progress.poll_transfer(transfer, require_completion=True) == (False, None) + assert len(payloads(recording, "nixl.done_observed")) == 1 + assert not recording.stack + + +def test_creation_failure_closes_scope_and_reports_rejection(recording): + agent, progress, trace = setup_transfer() + agent.make_prepped_xfer = lambda *a, **kw: None + with pytest.raises(RuntimeError, match="creation returned None"): + submit(progress, trace) + assert not progress._active_transfers + assert len(payloads(recording, "nixl.write.rejected")) == 1 + assert not recording.stack + + +@pytest.mark.parametrize("failure", ["attributes", "push", "mark", "pop"]) +def test_annotation_failure_cannot_change_nixl_lifecycle(recording, failure): + agent, progress, trace = setup_transfer() + recording.fail = failure + transfer, accepted = submit(progress, trace) + assert accepted + assert progress.poll_transfer(transfer, require_completion=True) == (True, None) + assert agent.events[-1] == f"release-transfer:{transfer}" + assert not recording.stack + + +@pytest.mark.parametrize("level,count", [("low", 0), ("medium", 1)]) +def test_release_retry_detail_level(recording, monkeypatch, level, count): + monkeypatch.setenv("KVCR_NVTX_LEVEL", level) + agent, progress, trace = setup_transfer() + transfer, _ = submit(progress, trace) + agent.release_failures = 1 + assert progress.poll_transfer(transfer, require_completion=True) is None + assert len(payloads(recording, "nixl.write.release_retry")) == count + assert progress.poll_transfer(transfer, require_completion=True) == (True, None) From 21c867430d0bda8b668eb391ae3f64ff42b6f6f5 Mon Sep 17 00:00:00 2001 From: aknvda Date: Mon, 5 Oct 2026 16:33:05 -0700 Subject: [PATCH 3/3] fix: preserve source route refusal reason in completion traces Signed-off-by: aknvda --- docs/nvtx-events.md | 5 +++++ docs/profiling.md | 11 +++++------ src/kvcr/remote_fw_dram.py | 2 ++ tests/unit/test_nvtx_transfers.py | 25 +++++++++++++++++++++++++ 4 files changed, 37 insertions(+), 6 deletions(-) diff --git a/docs/nvtx-events.md b/docs/nvtx-events.md index 94786da..0584ec4 100644 --- a/docs/nvtx-events.md +++ b/docs/nvtx-events.md @@ -1,3 +1,8 @@ + + # Remote delivery NVTX events Install the `profiling` extra. `KVCR_NVTX_LEVEL=off|low|medium` selects diff --git a/docs/profiling.md b/docs/profiling.md index be18ddb..3501fde 100644 --- a/docs/profiling.md +++ b/docs/profiling.md @@ -10,10 +10,9 @@ remote-memory path. It separates the synchronous `request_pin` callback from the asynchronous wait for its result. A shared pin has one lifetime and separate associations to each waiting source operation. -It does **not** yet trace NIXL submission/completion, target processing, router -hints, or caller polling. Complete request/session correlation and performance -validation remain separate work. Review a decoded pin capture before expanding -the hooks. +Source NIXL submission, native completion, release and cancellation are also +traced; see the [source lifecycle reference](nvtx-events.md). Target processing, +router hints and caller polling remain separate work in this source checkpoint. ## Enable tracing @@ -28,8 +27,8 @@ uv sync --extra profiling | Value | Behavior | | --- | --- | | `off` | No NVTX/NumPy import, pin trace objects, or payload construction. | -| `low` | Pin callback scopes, registration, waiter associations, terminal events. Default when the profiling dependencies are available. | -| `medium` | Low detail plus individual waiter-detachment events. | +| `low` | Pin lifecycles and source NIXL writes. Default when the profiling dependencies are available. | +| `medium` | Low detail plus waiter-detachment and native release-retry events. | Without the optional dependencies, tracing is a no-op. An explicit request for tracing with an unavailable backend warns once. An unsupported level warns once diff --git a/src/kvcr/remote_fw_dram.py b/src/kvcr/remote_fw_dram.py index 71fe95d..9618073 100644 --- a/src/kvcr/remote_fw_dram.py +++ b/src/kvcr/remote_fw_dram.py @@ -336,6 +336,8 @@ def progress( if self.state is _SourceWriteState.NOTIFY_FAILURE else _nvtx.Reason.DEADLINE ) + if failure_reason is _nvtx.Reason.UNKNOWN and trace is not None: + failure_reason = trace.failure_reason backend._send_write_done( progress, self.remote_agent, self.op_handle, False ) diff --git a/tests/unit/test_nvtx_transfers.py b/tests/unit/test_nvtx_transfers.py index cc5fac4..c3ec96b 100644 --- a/tests/unit/test_nvtx_transfers.py +++ b/tests/unit/test_nvtx_transfers.py @@ -72,6 +72,31 @@ def test_source_path_connects_pin_to_transfer(recording): assert len(payloads(recording, "source.write.completed")) == 1 +def test_route_refusal_preserves_cause_at_source_completion(recording, monkeypatch): + from kvcr.remote_fw_dram import _SourceWriteOp + + source, agent, pinning = test_nvtx.make_source(recording) + backend = source._core._remote_fw_dram + submit = source._core._progress.submit + + def replace_route(op): + if isinstance(op, _SourceWriteOp): + backend._route_generation[op.route[0]] = op.route[1] + 1 + submit(op) + + monkeypatch.setattr(source._core._progress, "submit", replace_route) + agent.state = "DONE" + pinning.complete(0) + test_nvtx._poll_until( + source, lambda _: bool(payloads(recording, "source.write.completed")) + ) + assert len(payloads(recording, "source.write.refused")) == 1 + assert not payloads(recording, "nixl.write.posted") + assert int(payloads(recording, "source.write.completed")[0]["reason"]) == int( + _nvtx.Reason.ROUTE_CHANGED + ) + + def setup_transfer(): agent = _TransferAgent() progress = _transfer_progress(agent)