File size: 3,019 Bytes
e791b16 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 | # Experience replay buffer utilities
import queue
import threading
from dataclasses import dataclass
from typing import List
import torch
@dataclass
class Experience:
prompt_ids: torch.Tensor
generated_ids: torch.Tensor
log_probs: torch.Tensor
reward: float
version: int
class BoundedReplayBuffer:
"""Thread‑safe bounded replay buffer.
- Non‑blocking `push`; if full, discards oldest entry.
- `sample` returns up to `batch_size` experiences, removing them from the buffer.
"""
def __init__(self, max_size: int = 10000):
self.max_size = max_size
self.queue = queue.Queue(maxsize=max_size)
self.lock = threading.Lock()
def push(self, exp: Experience):
with self.lock:
try:
self.queue.put_nowait(exp)
except queue.Full:
# discard oldest and insert new
try:
self.queue.get_nowait()
except queue.Empty:
pass
self.queue.put_nowait(exp)
def sample(self, batch_size: int) -> List[Experience]:
batch: List[Experience] = []
with self.lock:
while len(batch) < batch_size:
try:
batch.append(self.queue.get_nowait())
except queue.Empty:
break
return batch
def size(self) -> int:
return self.queue.qsize()
@dataclass
class VersionedExperience:
prompt_ids: torch.Tensor
generated_ids: torch.Tensor
log_probs: torch.Tensor
reward: float
policy_version: int
generation_step: int
class VersionedReplayBuffer:
"""Replay buffer that evicts experiences older than ``max_staleness`` versions.
- ``max_size``: maximum number of experiences to store.
- ``max_staleness``: maximum allowed age in policy versions.
An experience with ``policy_version < current_version - max_staleness`` is dropped.
- ``current_version``: set externally by the orchestrator before each push.
"""
def __init__(self, max_size: int = 10000, max_staleness: int = 5):
self.max_size = max_size
self.max_staleness = max_staleness
self.buffer: List[VersionedExperience] = []
self.current_version: int = 0
self.lock = threading.Lock()
def push(self, exp: VersionedExperience) -> None:
"""Push an experience, silently dropping it if it is too stale."""
with self.lock:
if exp.policy_version < self.current_version - self.max_staleness:
return # stale — drop
if len(self.buffer) >= self.max_size:
self.buffer.pop(0) # evict oldest
self.buffer.append(exp)
def sample(self, batch_size: int) -> List[VersionedExperience]:
"""Return up to ``batch_size`` experiences (no removal)."""
with self.lock:
return self.buffer[:batch_size]
def __len__(self) -> int:
return len(self.buffer)
|