Spaces:
Sleeping
Sleeping
| from __future__ import annotations | |
| import random | |
| from collections import deque | |
| from dataclasses import dataclass | |
| from typing import Deque | |
| import numpy as np | |
| class Transition: | |
| state: np.ndarray | |
| action: int | |
| reward: float | |
| next_state: np.ndarray | |
| done: float | |
| class ReplayBuffer: | |
| def __init__(self, capacity: int = 10000, seed: int | None = None) -> None: | |
| if capacity <= 0: | |
| raise ValueError("capacity must be > 0") | |
| self.capacity = capacity | |
| self.buffer: Deque[Transition] = deque(maxlen=capacity) | |
| self._random = random.Random(seed) | |
| def __len__(self) -> int: | |
| return len(self.buffer) | |
| def add(self, state: np.ndarray, action: int, reward: float, next_state: np.ndarray, done: bool) -> None: | |
| self.buffer.append( | |
| Transition( | |
| state=np.asarray(state, dtype=np.float32), | |
| action=int(action), | |
| reward=float(reward), | |
| next_state=np.asarray(next_state, dtype=np.float32), | |
| done=float(done), | |
| ) | |
| ) | |
| def sample(self, batch_size: int) -> list[Transition]: | |
| return self._random.sample(list(self.buffer), batch_size) | |