Speed up first-open memory sampling

This commit is contained in:
2026-05-09 23:57:51 +09:00
parent c570fc3ea5
commit 8217cecd19
2 changed files with 65 additions and 0 deletions
@@ -19,6 +19,8 @@ class ReservoirMemory:
def __init__(self, capacity: int | None = None) -> None: def __init__(self, capacity: int | None = None) -> None:
self.capacity = capacity self.capacity = capacity
self._samples: list[TrainingSample] = [] self._samples: list[TrainingSample] = []
self._first_open_indices: list[int] = []
self._first_open_positions: dict[int, int] = {}
self.seen = 0 self.seen = 0
def __len__(self) -> int: def __len__(self) -> int:
@@ -36,11 +38,14 @@ class ReservoirMemory:
) )
if self.capacity is None or len(self._samples) < self.capacity: if self.capacity is None or len(self._samples) < self.capacity:
self._samples.append(sample) self._samples.append(sample)
self._track_index(len(self._samples) - 1, sample)
return return
rng = rng or np.random.default_rng() rng = rng or np.random.default_rng()
index = int(rng.integers(0, self.seen)) index = int(rng.integers(0, self.seen))
if index < self.capacity: if index < self.capacity:
self._untrack_index(index)
self._samples[index] = sample self._samples[index] = sample
self._track_index(index, sample)
def extend(self, samples: list[TrainingSample], rng: np.random.Generator | None = None) -> None: def extend(self, samples: list[TrainingSample], rng: np.random.Generator | None = None) -> None:
self.add_many(samples, rng) self.add_many(samples, rng)
@@ -65,6 +70,16 @@ class ReservoirMemory:
first_open_only: bool = False, first_open_only: bool = False,
) -> list[TrainingSample]: ) -> list[TrainingSample]:
candidates = self._samples candidates = self._samples
if first_open_only and player is None:
if not self._first_open_indices:
raise ValueError("cannot sample from empty memory")
size = min(int(batch_size), len(self._first_open_indices))
indices = rng.choice(
len(self._first_open_indices),
size=size,
replace=len(self._first_open_indices) < size,
)
return [self._samples[self._first_open_indices[int(index)]] for index in indices]
if player is not None: if player is not None:
candidates = [sample for sample in candidates if sample.player == player] candidates = [sample for sample in candidates if sample.player == player]
if first_open_only: if first_open_only:
@@ -76,9 +91,27 @@ class ReservoirMemory:
return [candidates[int(index)] for index in indices] return [candidates[int(index)] for index in indices]
def count(self, *, player: int | None = None, first_open_only: bool = False) -> int: def count(self, *, player: int | None = None, first_open_only: bool = False) -> int:
if first_open_only and player is None:
return len(self._first_open_indices)
candidates = self._samples candidates = self._samples
if player is not None: if player is not None:
candidates = [sample for sample in candidates if sample.player == player] candidates = [sample for sample in candidates if sample.player == player]
if first_open_only: if first_open_only:
candidates = [sample for sample in candidates if sample.is_first_open] candidates = [sample for sample in candidates if sample.is_first_open]
return len(candidates) return len(candidates)
def _track_index(self, index: int, sample: TrainingSample) -> None:
if not sample.is_first_open:
return
self._first_open_positions[index] = len(self._first_open_indices)
self._first_open_indices.append(index)
def _untrack_index(self, index: int) -> None:
position = self._first_open_positions.pop(index, None)
if position is None:
return
last_index = self._first_open_indices.pop()
if position == len(self._first_open_indices):
return
self._first_open_indices[position] = last_index
self._first_open_positions[last_index] = position
@@ -996,6 +996,38 @@ def test_reservoir_memory_filters_first_open_batches() -> None:
assert memory.count(first_open_only=True) == 3 assert memory.count(first_open_only=True) == 3
def test_reservoir_memory_updates_first_open_index_on_replacement() -> None:
memory = ReservoirMemory(capacity=1)
replacement_rng = np.random.default_rng(1)
sample_rng = np.random.default_rng(37)
memory.add(
TrainingSample(
info_state=np.asarray([1], dtype=np.float32),
target=np.asarray([1], dtype=np.float32),
legal_mask=np.asarray([True]),
iteration=1,
player=0,
is_first_open=True,
),
replacement_rng,
)
memory.add(
TrainingSample(
info_state=np.asarray([2], dtype=np.float32),
target=np.asarray([2], dtype=np.float32),
legal_mask=np.asarray([True]),
iteration=2,
player=0,
is_first_open=False,
),
replacement_rng,
)
assert memory.count(first_open_only=True) == 0
with pytest.raises(ValueError, match="cannot sample from empty memory"):
memory.sample(1, sample_rng, first_open_only=True)
def test_deep_cfr_trainer_can_oversample_first_open_advantage_batches(tmp_path) -> None: def test_deep_cfr_trainer_can_oversample_first_open_advantage_batches(tmp_path) -> None:
trainer = DeepCFRTrainer( trainer = DeepCFRTrainer(
_deep_cfr_config( _deep_cfr_config(