Skip to content
2 changes: 1 addition & 1 deletion test/rb/test_ensemble.py
Original file line number Diff line number Diff line change
Expand Up @@ -437,7 +437,7 @@ def test_rb_multidim(self, datatype, datadim, rbtype, storage_cls, sampler_cls):
s = rb.sample()
assert str(rb)
if datatype in ("tensordict", "tensorclass"):
assert (s.exclude("index") == 1).all()
assert (s.exclude("index", "index_generation") == 1).all()
assert s.numel() == 4
else:
for leaf in tree_iter(s):
Expand Down
4 changes: 2 additions & 2 deletions test/rb/test_samplers.py
Original file line number Diff line number Diff line change
Expand Up @@ -243,7 +243,7 @@ def test_sampler_without_rep_state_dict(self, backend):
replay_buffer.extend(transition.clone())
for _ in range(n_samples):
s = replay_buffer.sample(batch_size=1)
assert (s.exclude("index") == 1).all()
assert (s.exclude("index", "index_generation") == 1).all()

replay_buffer.extend(torch.zeros_like(transition))

Expand All @@ -257,7 +257,7 @@ def test_sampler_without_rep_state_dict(self, backend):

new_replay_buffer.load_state_dict(state_dict)
s = new_replay_buffer.sample(batch_size=1)
assert (s.exclude("index") == 0).all()
assert (s.exclude("index", "index_generation") == 0).all()

def test_sampler_without_rep_dumps_loads(self, tmpdir):
d0 = tmpdir + "/save0"
Expand Down
2 changes: 1 addition & 1 deletion test/rb/test_storages.py
Original file line number Diff line number Diff line change
Expand Up @@ -306,7 +306,7 @@ def test_storage_state_dict(self, storage_in, storage_out, init_out, backend):

new_replay_buffer.load_state_dict(state_dict)
s = new_replay_buffer.sample()
assert (s.exclude("index") == 1).all()
assert (s.exclude("index", "index_generation") == 1).all()

@pytest.mark.skipif(
TORCH_VERSION < version.parse("2.5.0"), reason="requires Torch >= 2.5.0"
Expand Down
174 changes: 174 additions & 0 deletions test/rb/test_writers.py
Original file line number Diff line number Diff line change
Expand Up @@ -403,6 +403,180 @@ def test_roundrobin_dumps_loads_write_count(self, tmp_path):
assert writer2._write_count == 23


class TestWriterGeneration:
def test_default_writer_tracks_generations(self):
rb = ReplayBuffer(storage=LazyTensorStorage(10))
assert rb._writer.tracks_generations is True
index = rb.extend(torch.arange(10))
gen = rb._writer.generations_of(index)
assert gen.dtype == torch.int64
assert gen.shape == index.shape
assert (gen == 0).all()

def test_non_tracking_writer_reports_minus_one(self):
writer = TensorDictMaxValueWriter(rank_key="key")
assert writer.tracks_generations is False
gen = writer.generations_of(torch.arange(4))
torch.testing.assert_close(gen, torch.full((4,), -1, dtype=torch.int64))

def test_generation_increments_on_reuse(self):
size = 4
rb = ReplayBuffer(storage=LazyTensorStorage(size))
rb.extend(torch.arange(size))
torch.testing.assert_close(
rb._writer.generations_of(torch.arange(size)),
torch.zeros(size, dtype=torch.int64),
)
rb.extend(torch.arange(size, size + 3))
torch.testing.assert_close(
rb._writer.generations_of(torch.arange(size)), torch.tensor([1, 1, 1, 0])
)

def test_generation_wraparound(self):
size = 5
rb = ReplayBuffer(storage=LazyTensorStorage(size))
rb.extend(torch.arange(2 * size))
torch.testing.assert_close(
rb._writer.generations_of(torch.arange(size)),
torch.full((size,), 1, dtype=torch.int64),
)

def test_generation_add(self):
size = 3
rb = ReplayBuffer(storage=LazyTensorStorage(size))
for i in range(size + 1):
rb.add(torch.tensor(i))
torch.testing.assert_close(
rb._writer.generations_of(torch.arange(size)), torch.tensor([1, 0, 0])
)

def test_generations_of_unwritten_reports_minus_one(self):
rb = ReplayBuffer(storage=LazyTensorStorage(4))
rb.extend(torch.arange(2))
gen = rb._writer.generations_of(torch.arange(4))
torch.testing.assert_close(gen, torch.tensor([0, 0, -1, -1]))

def test_generation_tensordict_writer(self):
size = 4
rb = TensorDictReplayBuffer(storage=LazyTensorStorage(size))
rb.extend(TensorDict({"a": torch.arange(2 * size)}, [2 * size]))
torch.testing.assert_close(
rb._writer.generations_of(torch.arange(size)),
torch.full((size,), 1, dtype=torch.int64),
)

def test_generation_write_at(self):
storage = LazyTensorStorage(4)
writer = RoundRobinWriter()
writer.register_storage(storage)
writer.extend(torch.arange(4))
writer.write_at(torch.tensor([0, 1]), torch.tensor([10, 11]))
torch.testing.assert_close(
writer.generations_of(torch.arange(4)), torch.tensor([1, 1, 0, 0])
)

def test_empty_is_monotonic(self):
rb = ReplayBuffer(storage=LazyTensorStorage(10))
index = rb.extend(torch.arange(10))
before = rb._writer.generations_of(index)
rb.empty()
rb.extend(torch.arange(10))
after = rb._writer.generations_of(index)
assert (after > before).all()

def test_empty_invalidates_handles_immediately(self):
rb = ReplayBuffer(storage=LazyTensorStorage(10))
index = rb.extend(torch.arange(10))
gen = rb._writer.generations_of(index)
rb.empty()
assert (rb._writer.generations_of(index) != gen).all()

def test_generation_state_dict_roundtrip(self):
size = 4
rb = ReplayBuffer(storage=LazyTensorStorage(size))
rb.extend(torch.arange(size + 1))
sd = rb.state_dict()
rb2 = ReplayBuffer(storage=LazyTensorStorage(size))
rb2.load_state_dict(sd)
torch.testing.assert_close(
rb2._writer.generations_of(torch.arange(size)),
rb._writer.generations_of(torch.arange(size)),
)

def test_legacy_state_dict_without_generation_loads(self):
rb = ReplayBuffer(storage=LazyTensorStorage(10))
rb.extend(torch.arange(5))
sd = rb.state_dict()
del sd["_writer"]["_generation"]
rb2 = ReplayBuffer(storage=LazyTensorStorage(10))
rb2.load_state_dict(sd)
assert rb2._writer._cursor == 5

def test_generation_dumps_loads(self, tmp_path):
writer = RoundRobinWriter()
writer._cursor = 2
writer._write_count = 9
writer._generation = torch.tensor([3, 2, 2, 1])
writer.dumps(tmp_path)
writer2 = RoundRobinWriter()
writer2.loads(tmp_path)
assert writer2._cursor == 2
assert writer2._write_count == 9
torch.testing.assert_close(
writer2.generations_of(torch.arange(4)), torch.tensor([3, 2, 2, 1])
)

def test_sample_returns_generation(self):
size = 8
rb = ReplayBuffer(storage=LazyTensorStorage(size))
rb.extend(torch.arange(size))
_, info = rb.sample(4, return_info=True)
assert "index_generation" in info
gen = torch.as_tensor(info["index_generation"])
idx = torch.as_tensor(info["index"])
assert gen.shape == idx.shape
torch.testing.assert_close(gen, rb._writer.generations_of(idx))

def test_non_tracking_sample_has_no_generation(self):
rb = TensorDictReplayBuffer(
storage=LazyTensorStorage(10),
writer=TensorDictMaxValueWriter(rank_key="key"),
)
rb.extend(TensorDict({"key": torch.arange(10), "a": torch.arange(10)}, [10]))
_, info = rb.sample(4, return_info=True)
assert "index_generation" not in info

def test_tensordict_sample_has_generation_key(self):
size = 8
rb = TensorDictReplayBuffer(storage=LazyTensorStorage(size))
rb.extend(TensorDict({"a": torch.arange(size)}, [size]))
sample = rb.sample(4)
assert "index_generation" in sample.keys()
assert sample["index_generation"].shape[0] == 4

def test_wraparound_race_detectable(self):
size = 8
rb = ReplayBuffer(storage=LazyTensorStorage(size))
rb.extend(torch.arange(size))
_, info = rb.sample(4, return_info=True)
sampled_index = torch.as_tensor(info["index"])
sampled_generation = torch.as_tensor(info["index_generation"])
rb.extend(torch.arange(size, 2 * size))
current = rb._writer.generations_of(sampled_index)
assert (current != sampled_generation).all()

def test_partial_reuse_detectable(self):
size = 8
rb = ReplayBuffer(storage=LazyTensorStorage(size))
rb.extend(torch.arange(size))
_, info = rb.sample(size, return_info=True)
idx = torch.as_tensor(info["index"])
gen = torch.as_tensor(info["index_generation"])
rb.extend(torch.arange(size, size + 3))
stale = rb._writer.generations_of(idx) != gen
torch.testing.assert_close(stale, idx < 3)


if __name__ == "__main__":
args, unknown = argparse.ArgumentParser().parse_known_args()
pytest.main([__file__, "--capture", "no", "--exitfirst"] + unknown)
8 changes: 8 additions & 0 deletions torchrl/data/replay_buffers/replay_buffers.py
Original file line number Diff line number Diff line change
Expand Up @@ -1516,6 +1516,8 @@ def _sample(self, batch_size: int) -> tuple[Any, dict]:
with self._replay_lock if not is_comp else nc, self._write_lock if not is_comp else nc:
index, info = self._sampler.sample(self._storage, batch_size)
info["index"] = index
if self._writer.tracks_generations:
info["index_generation"] = self._writer.generations_of(index)
data = self._storage.get(_storage_index(index, self._storage))
if not isinstance(index, INT_CLASSES):
data = self._collate_fn(data)
Expand Down Expand Up @@ -2135,6 +2137,8 @@ def _sample(self, batch_size: int) -> tuple[Any, dict]:
):
index, info = self.prioritized_sampler.sample(self._storage, batch_size)
info["index"] = index
if self._writer.tracks_generations:
info["index_generation"] = self._writer.generations_of(index)
data = self._storage.get(_storage_index(index, self._storage))
if not isinstance(index, INT_CLASSES):
data = self._collate_fn(data)
Expand Down Expand Up @@ -2564,6 +2568,8 @@ def _sample(self, batch_size: int) -> tuple[Any, dict]:
with self._replay_lock if not is_comp else nc, self._write_lock if not is_comp else nc:
index, info = self._sampler.sample(self._storage, batch_size)
info["index"] = index
if self._writer.tracks_generations:
info["index_generation"] = self._writer.generations_of(index)
data = self._storage.get(_storage_index(index, self._storage))
if not isinstance(index, INT_CLASSES):
data = self._collate_fn(data)
Expand Down Expand Up @@ -2929,6 +2935,8 @@ def _sample(self, batch_size: int) -> tuple[Any, dict]:
):
index, info = self.prioritized_sampler.sample(self._storage, batch_size)
info["index"] = index
if self._writer.tracks_generations:
info["index_generation"] = self._writer.generations_of(index)
data = self._storage.get(_storage_index(index, self._storage))
if not isinstance(index, INT_CLASSES):
data = self._collate_fn(data)
Expand Down
Loading
Loading