Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions src/nnsight/intervention/tracing/tracer.py
Original file line number Diff line number Diff line change
Expand Up @@ -318,6 +318,12 @@ def add(self, module_path: str, key: str, value: Any):
lambda x: x.to(device=_device, dtype=_dtype, non_blocking=True),
torch.Tensor,
)
# Synchronize GPU operations to ensure cached tensors are fully
# transferred before subsequent forward passes access them. Without
# this, non_blocking=True allows async transfers that may not complete
# before the model reads the cached value, causing stale data issues.
if _device is not None and _device.type == "cuda":
torch.cuda.synchronize(_device)

if module_path not in self.cache:
self.cache[module_path] = Cache.Entry(**{key: value})
Expand Down
180 changes: 180 additions & 0 deletions tests/repro_cache_stream_ordering.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,180 @@
"""GPU repro harness for cache stream ordering on main.

Purpose:
- Verify whether Cache.add runs on an unexpected CUDA stream.
- Check if cached tensors can be consumed safely without torch.cuda.synchronize.

How to run:
python tests/repro_cache_stream_ordering.py --runs 50

Notes:
- This script is intentionally standalone (not a pytest test) so it can be
shared directly in PR comments as a reproducibility harness.
- It exits non-zero when it finds a real mismatch/failure signal.
"""

from __future__ import annotations

import argparse
import math
import sys
from dataclasses import dataclass
from typing import Any, List, Optional

import torch

import nnsight
from nnsight.intervention.tracing.tracer import Cache


@dataclass
class AddRecord:
module_path: str
key: str
stream_ptr: Optional[int]
tensor_device: Optional[str]


def first_tensor(value: Any) -> Optional[torch.Tensor]:
if isinstance(value, torch.Tensor):
return value
if isinstance(value, (list, tuple)):
for item in value:
t = first_tensor(item)
if t is not None:
return t
return None
if isinstance(value, dict):
for item in value.values():
t = first_tensor(item)
if t is not None:
return t
return None
return None


def run_repro(runs: int, model_id: str, prompt: str) -> int:
if not torch.cuda.is_available():
print("CUDA is not available in this environment. Cannot run GPU stream repro.")
return 2

device = "cuda:0"
torch.set_grad_enabled(False)

print(f"Loading model: {model_id} on {device}")
model = nnsight.LanguageModel(model_id, device_map=device, dispatch=True)

records: List[AddRecord] = []
orig_add = Cache.add

def wrapped_add(self, module_path: str, key: str, value: Any):
stream_ptr: Optional[int]
if torch.cuda.is_available():
stream_ptr = int(torch.cuda.current_stream().cuda_stream)
else:
stream_ptr = None

tensor = first_tensor(value)
tensor_device = str(tensor.device) if tensor is not None else None
records.append(
AddRecord(
module_path=module_path,
key=key,
stream_ptr=stream_ptr,
tensor_device=tensor_device,
)
)

return orig_add(self, module_path, key, value)

Cache.add = wrapped_add

mismatches = 0
cross_stream_mismatches = 0
nan_failures = 0

consumer_stream = torch.cuda.Stream(device=torch.device(device))

try:
for i in range(runs):
with model.trace(prompt) as tracer:
# Force real GPU work in the cache transform by changing dtype.
cache = tracer.cache(
device=torch.device(device),
dtype=torch.float32,
modules=[model.transformer.h[0]],
)

cached_output = cache["model.transformer.h.0"].output

# Same-stream consume (this is the natural path in user code).
same_stream_value = cached_output.sum().item()

# Different-stream consume (stress case). We intentionally do not
# add any explicit stream wait here.
with torch.cuda.stream(consumer_stream):
cross_stream_tensor = cached_output.sum()

consumer_stream.synchronize()
cross_stream_value = cross_stream_tensor.item()

# Fresh run baseline.
with model.trace(prompt):
fresh_output = model.transformer.h[0].output.save()
fresh_value = fresh_output.to(torch.float32).sum().item()

if math.isnan(same_stream_value) or math.isnan(cross_stream_value):
nan_failures += 1

if not math.isclose(same_stream_value, fresh_value, rel_tol=0.0, abs_tol=1e-4):
mismatches += 1

if not math.isclose(cross_stream_value, fresh_value, rel_tol=0.0, abs_tol=1e-4):
cross_stream_mismatches += 1

if (i + 1) % 10 == 0 or i == runs - 1:
print(
f"run {i + 1}/{runs} | same_stream_mismatch={mismatches} "
f"cross_stream_mismatch={cross_stream_mismatches} nan_failures={nan_failures}"
)

finally:
Cache.add = orig_add

stream_ptrs = sorted({r.stream_ptr for r in records if r.stream_ptr is not None})
null_stream_seen = any(ptr == 0 for ptr in stream_ptrs)
cuda_records = [r for r in records if r.tensor_device and "cuda" in r.tensor_device]

print("\n=== Summary ===")
print(f"cache_add_calls={len(records)}")
print(f"cuda_cache_add_calls={len(cuda_records)}")
print(f"unique_cache_add_stream_ptrs={stream_ptrs}")
print(f"null_stream_seen={null_stream_seen}")
print(f"same_stream_mismatches={mismatches}")
print(f"cross_stream_mismatches={cross_stream_mismatches}")
print(f"nan_failures={nan_failures}")

# Fail only when there is an actual numerical mismatch/failure signal.
if mismatches > 0 or cross_stream_mismatches > 0 or nan_failures > 0:
print("RESULT: FAIL (found mismatch/failure signal)")
return 1

print("RESULT: PASS (no mismatch signal without global synchronize)")
return 0


def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Repro harness for cache stream ordering")
parser.add_argument("--runs", type=int, default=50, help="Number of repeated trials")
parser.add_argument("--model", type=str, default="openai-community/gpt2")
parser.add_argument(
"--prompt",
type=str,
default="Madison Square Garden is located in the city of",
)
return parser.parse_args()


if __name__ == "__main__":
args = parse_args()
sys.exit(run_repro(args.runs, args.model, args.prompt))
31 changes: 31 additions & 0 deletions tests/test_lm.py
Original file line number Diff line number Diff line change
Expand Up @@ -1294,6 +1294,37 @@ def test_cache_alias(self, MSG_prompt: str):
cache.model.model.h["second_layer"].output,
)

@torch.no_grad()
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_cache_gpu_synchronization(self, gpt2: nnsight.LanguageModel, MSG_prompt: str):
"""Test that GPU-cached tensors are synchronized and ready for use.

When caching to GPU with non_blocking=True, we must synchronize to ensure
the tensor transfer completes before subsequent forward passes access the cache.
This test verifies the fix for the GPU sync gap where stale data could be read.
"""
# Cache to GPU
with gpt2.trace(MSG_prompt) as tracer:
cache = tracer.cache(device=torch.device("cuda"), modules=[gpt2.transformer.h[0]])

# Retrieve cached tensor — if synchronization failed, this could return stale data
cached_output = cache["model.transformer.h.0"].output

# Verify the cached tensor is actually on GPU and ready
assert cached_output.device.type == "cuda"

# Perform immediate computation on cached tensor to verify it's fully transferred
# and not stuck in an async transfer. This would fail with a CUDA error if sync is missing.
result = cached_output.sum()
assert result.is_cuda

# Re-trace the same input and verify the cached value matches fresh computation
with gpt2.trace(MSG_prompt):
fresh_output = gpt2.transformer.h[0].output.save()

# The cached version should match the fresh computation (same input)
assert torch.allclose(cached_output.cpu(), fresh_output.cpu())


# =============================================================================
# Module Renaming
Expand Down