Skip to content
Merged
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
67 changes: 47 additions & 20 deletions benchmarks/ad_hoc/bench_dreamer_v3_learner.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
Example::

python benchmarks/ad_hoc/bench_dreamer_v3_learner.py
python benchmarks/ad_hoc/bench_dreamer_v3_learner.py --device cpu --variants eager --replay-device cpu
"""

from __future__ import annotations
Expand Down Expand Up @@ -149,7 +150,8 @@ def _measure(step, synchronize, *, warmup: int, iterations: int):
for _ in range(warmup):
step()
synchronize()
torch.cuda.reset_peak_memory_stats()
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats()
host_latencies = []
started = time.perf_counter()
for _ in range(iterations):
Expand All @@ -160,13 +162,11 @@ def _measure(step, synchronize, *, warmup: int, iterations: int):
return (time.perf_counter() - started) * 1000 / iterations, host_latencies


def _profile_update(step, synchronize) -> dict:
with torch.profiler.profile(
activities=[
torch.profiler.ProfilerActivity.CPU,
torch.profiler.ProfilerActivity.CUDA,
]
) as profile:
def _profile_update(step, synchronize, device: torch.device) -> dict:
activities = [torch.profiler.ProfilerActivity.CPU]
if device.type == "cuda":
activities.append(torch.profiler.ProfilerActivity.CUDA)
with torch.profiler.profile(activities=activities) as profile:
step()
synchronize()

Expand Down Expand Up @@ -200,6 +200,13 @@ def main() -> None:
parser.add_argument("--unroll", type=int, default=8)
parser.add_argument("--warmup", type=int, default=10)
parser.add_argument("--iterations", type=int, default=50)
parser.add_argument(
"--device",
choices=("cpu", "cuda"),
default="cuda",
help="Learner device. CPU runs support the eager and compiled variants "
"without CUDA graphs.",
)
parser.add_argument(
"--replay-device",
choices=("cpu", "cuda"),
Expand All @@ -215,19 +222,27 @@ def main() -> None:
"compiled_train_step",
"compiled_train_step_cuda_graph",
),
default=(
"cuda_graph",
"compiled_train_step_cuda_graph",
),
default=None,
)
args = parser.parse_args()
if not torch.cuda.is_available():
raise RuntimeError("This benchmark requires CUDA.")
if args.variants is None:
args.variants = (
("cuda_graph", "compiled_train_step_cuda_graph")
if args.device == "cuda"
else ("eager",)
)
if args.device == "cuda" or args.replay_device == "cuda":
if not torch.cuda.is_available():
raise RuntimeError("CUDA devices were requested but CUDA is unavailable.")
if args.device != "cuda" and any(
"cuda_graph" in variant for variant in args.variants
):
raise RuntimeError("CUDA graph variants require --device cuda.")

repo_root = Path(__file__).parents[2]
example = _load_example(repo_root)
template_cfg = _load_config(repo_root)
device = torch.device("cuda:0")
device = torch.device("cuda:0" if args.device == "cuda" else "cpu")
torch.set_float32_matmul_precision("high")

real_env = example["make_env"](template_cfg, template_cfg.env.seed)
Expand Down Expand Up @@ -285,7 +300,11 @@ def main() -> None:
replay_step = None
if args.replay_device is None:
step = ft.partial(learner_update.step, None, data.clone())
synchronize = ft.partial(torch.cuda.synchronize, device)
synchronize = (
ft.partial(torch.cuda.synchronize, device)
if device.type == "cuda"
else lambda: None
)
workload = "complete_learner_update"
else:
replay_step = _ReplayLearnerStep(
Expand All @@ -307,15 +326,21 @@ def main() -> None:
warmup=args.warmup,
iterations=args.iterations,
)
peak_memory = torch.cuda.max_memory_allocated(device) / 2**20
profile_metrics = _profile_update(step, synchronize)
peak_memory = (
torch.cuda.max_memory_allocated(device) / 2**20
if device.type == "cuda"
else None
)
profile_metrics = _profile_update(step, synchronize, device)
finally:
if replay_step is not None:
replay_step.close()
result = {
"variant": variant,
"workload": workload,
"device": torch.cuda.get_device_name(device),
"device": (
torch.cuda.get_device_name(device) if device.type == "cuda" else "cpu"
),
"replay_device": args.replay_device,
"torch_version": torch.__version__,
"cuda_version": torch.version.cuda,
Expand All @@ -328,6 +353,7 @@ def main() -> None:
"warmup_updates": args.warmup,
"measured_updates": args.iterations,
"mean_update_ms": mean_ms,
"updates_per_second": 1000 / mean_ms,
"transitions_per_second": args.batch * args.steps * 1000 / mean_ms,
"p50_host_update_ms": statistics.median(host_latencies),
"p95_host_update_ms": statistics.quantiles(
Expand All @@ -341,7 +367,8 @@ def main() -> None:
# variant before measuring the next one's live and peak allocations.
del step, synchronize, learner_update, learner, replay_step
gc.collect()
torch.cuda.empty_cache()
if device.type == "cuda":
torch.cuda.empty_cache()


if __name__ == "__main__":
Expand Down
67 changes: 52 additions & 15 deletions sota-implementations/dreamer_v3/dreamer_v3_replay.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,12 +40,58 @@ def collector_action_budget(
return (vector_records - reset_records) * num_envs


def _last_rows_per_coordinate(coordinates: torch.Tensor) -> torch.Tensor:
"""Return the last row holding each distinct coordinate, in coordinate order.

Packs the non-negative integer columns into one key so a single stable sort
orders the rows lexicographically with ties in their original order; the
last row of every run of equal keys then wins. When the packed key would
overflow ``int64``, the columns are sorted one at a time instead.
"""
coordinates = coordinates.long()
n_rows, n_columns = coordinates.shape
radices = coordinates.amax(0) + 1
if bool(radices.double().log2().sum() < 62):
strides = torch.ones_like(radices)
strides[:-1] = radices[1:].flip(0).cumprod(0).flip(0)
key = (coordinates * strides).sum(-1)
order = key.argsort(stable=True)
ordered = key[order].unsqueeze(-1)
else:
order = torch.arange(n_rows, device=coordinates.device)
for column in range(n_columns - 1, -1, -1):
order = order[coordinates[order, column].argsort(stable=True)]
ordered = coordinates[order]
last = torch.ones(n_rows, dtype=torch.bool, device=coordinates.device)
last[:-1] = (ordered[:-1] != ordered[1:]).any(-1)
return order[last]


def _index_on(index: torch.Tensor, device: torch.device) -> torch.Tensor:
"""Move a host index to ``device`` without synchronizing the current stream.

A pageable host-to-device copy waits for every kernel already enqueued,
which after a learner step is the whole step; staging through pinned memory
keeps the copy asynchronous.
"""
if index.device == device:
return index
if device.type == "cuda" and index.device.type == "cpu":
return index.pin_memory().to(device, non_blocking=True)
return index.to(device)


def replay_context_update(
sample: TensorDictBase,
state: torch.Tensor,
belief: torch.Tensor,
) -> tuple[TensorDictBase, torch.Tensor, Mapping[NestedKey, torch.Tensor]]:
"""Build a deduplicated update for the context rows after a sampled step."""
"""Build a deduplicated update for the context rows after a sampled step.

Nothing here synchronizes with the device holding ``state`` and
``belief``; the returned patch stays on that device and may still be in
flight, so pass it to an asynchronous replay update.
"""
if sample.ndim != 2:
raise RuntimeError(
"Expected a replay sample with shape [batch, time], got "
Expand All @@ -69,23 +115,14 @@ def replay_context_update(

# Sampled slices may overlap. Keep the last value for each destination so
# indexed writes have deterministic semantics on every device.
order = torch.arange(coordinates.shape[0], device=coordinates.device)
for dimension in range(coordinates.shape[1] - 1, -1, -1):
order = order[coordinates[order, dimension].argsort(stable=True)]
ordered_coordinates = coordinates[order]
keep_ordered = torch.ones(
ordered_coordinates.shape[0], dtype=torch.bool, device=coordinates.device
)
keep_ordered[:-1] = (ordered_coordinates[:-1] != ordered_coordinates[1:]).any(-1)
keep = order[keep_ordered]

state = state.detach().float().reshape(-1, state.shape[-1])
belief = belief.detach().float().reshape(-1, belief.shape[-1])
keep = _last_rows_per_coordinate(coordinates)
state = state.detach().reshape(-1, state.shape[-1])
belief = belief.detach().reshape(-1, belief.shape[-1])
return (
handles[keep],
generations[keep],
{
"state": state[keep.to(state.device)],
"belief": belief[keep.to(belief.device)],
"state": state[_index_on(keep, state.device)].float(),
"belief": belief[_index_on(keep, belief.device)].float(),
},
)
109 changes: 109 additions & 0 deletions test/objectives/test_dreamer_v3.py
Original file line number Diff line number Diff line change
Expand Up @@ -2537,6 +2537,115 @@ def test_dreamer_v3_async_replay_sequences_cross_episode_ends(monkeypatch, onlin
rb.shutdown()


def _reference_replay_context_update(sample, state, belief):
"""Per-column stable sorts, the implementation the packed key replaced."""
handles = sample.get("index")[:, 1:].reshape(-1)
generations = sample.get("index_generation")[:, 1:].reshape(-1)
buffer_ids = handles.get("buffer_ids")
local_indices = handles.get("index")
coordinates = torch.cat(
(buffer_ids.reshape(-1, 1), local_indices.reshape(buffer_ids.numel(), -1)),
-1,
)
order = torch.arange(coordinates.shape[0])
for dimension in range(coordinates.shape[1] - 1, -1, -1):
order = order[coordinates[order, dimension].argsort(stable=True)]
ordered_coordinates = coordinates[order]
keep_ordered = torch.ones(ordered_coordinates.shape[0], dtype=torch.bool)
keep_ordered[:-1] = (ordered_coordinates[:-1] != ordered_coordinates[1:]).any(-1)
keep = order[keep_ordered]
state = state.detach().float().reshape(-1, state.shape[-1])
belief = belief.detach().float().reshape(-1, belief.shape[-1])
return (
handles[keep],
generations[keep],
{"state": state[keep], "belief": belief[keep]},
)


def _overlapping_replay_sample(batch, time, *, local_ndim, scale=1):
"""Slice windows that overlap within members, as sampled sequences do."""
generator = torch.Generator().manual_seed(0)
capacity = 12
starts = torch.randint(0, capacity - time + 1, (batch,), generator=generator)
local = starts.unsqueeze(1) + torch.arange(time)
if local_ndim == 2:
lane = torch.randint(0, 3, (batch, 1), generator=generator).expand(batch, time)
local = torch.stack((lane, local), -1)
buffer_ids = torch.randint(0, 4, (batch, 1), generator=generator).expand(
batch, time
)
# Two identical windows make whole rows collide, not only their tails.
buffer_ids = buffer_ids.clone()
buffer_ids[1] = buffer_ids[0]
local = local.clone()
local[1] = local[0]
return TensorDict(
{
"index": TensorDict(
{"buffer_ids": buffer_ids * scale, "index": local * scale},
[batch, time],
),
"index_generation": torch.randint(0, 3, (batch, time), generator=generator),
},
[batch, time],
)


@pytest.mark.parametrize("local_ndim", [1, 2])
@pytest.mark.parametrize(
"scale", [1, 2**33], ids=["packed_key", "per_column_fallback"]
)
@pytest.mark.parametrize(
"device",
[
"cpu",
pytest.param(
"cuda",
marks=[
pytest.mark.gpu,
pytest.mark.skipif(
not torch.cuda.is_available(), reason="requires CUDA"
),
],
),
],
)
def test_dreamer_v3_replay_context_update_matches_reference(
monkeypatch, local_ndim, scale, device
):
repo_root = Path(__file__).parents[2]
example_dir = repo_root / "sota-implementations/dreamer_v3"
monkeypatch.syspath_prepend(str(example_dir))
replay = runpy.run_path(
example_dir / "dreamer_v3_replay.py", run_name="dreamer_v3_replay_test"
)
batch, time = 16, 5
sample = _overlapping_replay_sample(batch, time, local_ndim=local_ndim, scale=scale)
state = torch.randn(batch, time - 1, 6)
belief = torch.randn(batch, time - 1, 7)

index, generation, patch = replay["replay_context_update"](
sample, state.to(device), belief.to(device)
)
ref_index, ref_generation, ref_patch = _reference_replay_context_update(
sample, state, belief
)

coordinates = torch.cat(
(index["buffer_ids"].reshape(-1, 1), index["index"].reshape(len(index), -1)),
-1,
)
assert coordinates.shape[0] < batch * (time - 1)
assert torch.unique(coordinates, dim=0).shape[0] == coordinates.shape[0]
assert torch.equal(index["buffer_ids"], ref_index["buffer_ids"])
assert torch.equal(index["index"], ref_index["index"])
assert torch.equal(generation, ref_generation)
for key in ("state", "belief"):
assert patch[key].device.type == device
assert torch.equal(patch[key].cpu(), ref_patch[key])


@pytest.mark.skipif(not _has_omegaconf, reason="requires omegaconf")
def test_dreamer_v3_replay_capacity_validation(monkeypatch):
from omegaconf import OmegaConf
Expand Down
Loading
Loading