Samsara RL
A vectorized NumPy implementation of foundational reinforcement learning algorithms, following David Silver's RL lecture series, Sutton and Barto book and other papers cited when referenced. Built for clarity, learning as well as speed.
Applications of RL include robotic manipulation, LLM fine-tuning, financial portfolio management, and control systems.
Algorithms
| Algorithm | Category | How It Works | Good For | Limitations |
|---|---|---|---|---|
| Policy Iteration | Planning | Alternates policy evaluation (Bellman expectation) and greedy improvement until convergence | Small MDPs with known dynamics | Requires full model (transition + reward matrices) |
| Value Iteration | Planning | Applies Bellman optimality equation directly; extracts policy at convergence | Small MDPs with known dynamics | Requires full model; slower per-iteration than policy iteration for large state spaces |
| Monte Carlo | Tabular Prediction | Estimates Q(s, a) from sampled episode returns using constant-alpha updates | Episodic tasks; unbiased value estimates | High variance; must wait until episode end |
| TD(λ) | Tabular Prediction | One-step bootstrapping with eligibility traces for online Q updates | Online learning; continuous tasks | Biased estimates from bootstrapping |
| SARSA | Tabular Control | On-policy TD control; bootstraps from action actually taken under ε-greedy | Safe exploration; risk-sensitive tasks | Learns ε-greedy value, not optimal value |
| Q-Learning | Tabular Control | Off-policy TD control; bootstraps from max Q(S', a) regardless of action taken | Learning optimal policy while exploring | Maximization bias; can overestimate Q values |
| Linear Semi-Gradient TD(λ) | Value Approximation | TD(λ) with linear function approximation and eligibility traces | Large/continuous state spaces with known features | Linear capacity; requires manual feature engineering |
| DQN | Value Approximation | Neural network Q-function with replay buffer and target network | High-dimensional continuous state spaces | Maximization bias; training instability |
| Double DQN | Value Approximation | DQN with decoupled action selection and evaluation to reduce overestimation | Same as DQN with more stable Q estimates | Still sensitive to hyperparameters |
| Monte Carlo Policy Gradient | Policy Gradient | Accumulates policy gradients over a batch of episodes using discounted returns, then updates the policy network | Episodic tasks; continuous or large state spaces | High variance; must wait until episode end; sensitive to baseline choice |
| REINFORCE | Policy Gradient | Special case of MC Policy Gradient with batch_size=1; updates after every episode | Simple episodic tasks; learning/prototyping | Highest variance; no gradient averaging across episodes |
Table of Contents
- Installation
- Quick Start
- Planning
- Model-Free Prediction
- Model-Free Control
- Function Approximation
- Deep Q-Network
- Policy Gradient
Installation
pip install samsara-rl
Quick Start
from samsara_rl.mdp.grid_world.grid_world_mdp import GridWorldMDP
from samsara_rl.planning.policy_iteration import PolicyIteration
mdp = GridWorldMDP()
pi = PolicyIteration(mdp)
policy = pi.find_optimal_policy()
Planning
Planning algorithms assume full knowledge of environment dynamics (transition probabilities and reward function). While not "true RL" — agents never have access to dynamics in practice — planning provides the theoretical foundation all RL algorithms build on.
MDP Structure
MDPs are represented as NumPy arrays. The included GridWorldMDP implements the 4x4 grid world from David Silver's Lecture 3.
| Attribute | Shape | Description |
|---|---|---|
state_action_transition_matrix |
(S, A, S') |
T(s, a, s') — transition probabilities |
reward_matrix |
(S, A, S') |
R(s, a, s') — reward for each transition |
Policy Iteration
Alternates between evaluating the current policy using the Bellman expectation equation and improving it greedily until the policy stops changing.
PolicyIteration(mdp, bellman_tolerance)
| Argument | Type | Default | Description |
|---|---|---|---|
mdp |
MDP | MDP instance to solve | |
bellman_tolerance |
float |
0.01 |
Convergence threshold for policy evaluation |
find_optimal_policy(max_iter)
| Argument | Type | Default | Description |
|---|---|---|---|
max_iter |
int |
99 |
Maximum number of policy iteration steps |
Value Iteration
Applies the Bellman optimality equation directly each iteration. Equivalent to policy iteration with k=1 evaluation steps per improvement. Policy is extracted once at convergence.
ValueIteration(mdp, bellman_tolerance)
| Argument | Type | Default | Description |
|---|---|---|---|
mdp |
MDP | MDP instance to solve | |
bellman_tolerance |
float |
0.01 |
Convergence threshold for value iteration |
Examples
from samsara_rl.mdp.grid_world.grid_world_mdp import GridWorldMDP
from samsara_rl.planning.policy_iteration import PolicyIteration
from samsara_rl.planning.value_iteration import ValueIteration
mdp = GridWorldMDP()
policy = PolicyIteration(mdp, bellman_tolerance=0.001).find_optimal_policy(max_iter=50)
policy = ValueIteration(mdp, bellman_tolerance=0.001).find_optimal_policy()
Model-Free Prediction
Model-free methods learn value functions directly from experience (sampled episodes) without access to the MDP's transition or reward dynamics.
Monte Carlo
Every-visit Monte Carlo prediction estimates Q(s, a) from sampled returns. After each episode, the return G_t (discounted cumulative reward from time step t onward) is computed for every visited state-action pair, and the Q-table is updated using constant-alpha learning:
Q(s, a) <- Q(s, a) + α (G_t - Q(s, a))
If the same (s, a) pair appears multiple times in an episode, each occurrence triggers an update. Returns are computed in a fully vectorized pass using a cumulative-sum trick that avoids the standard reverse loop over time steps.
MonteCarloPolicyEvaluation(mdp, policy, alpha, gamma)
| Argument | Type | Default | Description |
|---|---|---|---|
mdp |
MDP | MDP instance to sample episodes from | |
policy |
array |
Stochastic policy of shape (S, A) |
|
alpha |
float |
0.01 |
Learning rate for incremental Q updates |
gamma |
float |
1 |
Discount factor |
evaluate(max_iter)
| Argument | Type | Default | Description |
|---|---|---|---|
max_iter |
int |
10000 |
Number of episodes to sample from |
TD(λ)
TD(λ) learns Q(s, a) online using one-step bootstrapping with eligibility traces. After each step, the TD error is computed against the expected Q-value of the next state under the current policy, and all previously visited state-action pairs are updated proportionally to their eligibility:
δ = R + γ E_π[Q(S', ·)] - Q(S, A)
Q(s, a) <- Q(s, a) + α δ e(s, a)
Eligibility traces use the replacing variant — on each visit to (s, a), the trace is set to 1 rather than incremented. All traces decay by γλ at each time step. Traces are reset to zero between episodes.
TDPolicyEvaluation(mdp, policy, alpha, gamma, _lambda)
| Argument | Type | Default | Description |
|---|---|---|---|
mdp |
MDP | MDP instance to sample episodes from | |
policy |
array |
Stochastic policy of shape (S, A) |
|
alpha |
float |
0.01 |
Learning rate for incremental Q updates |
gamma |
float |
0.9 |
Discount factor |
_lambda |
float |
0.4 |
Trace decay parameter (0 = TD(0), 1 = TD(1)) |
evaluate(max_iter)
| Argument | Type | Default | Description |
|---|---|---|---|
max_iter |
int |
1000 |
Number of episodes to sample from |
Examples
from samsara_rl.mdp.grid_world.grid_world_mdp import GridWorldMDP
from samsara_rl.prediction.monte_carlo import MonteCarloPolicyEvaluation
from samsara_rl.prediction.td import TDPolicyEvaluation
from samsara_rl.utils.policy.policy_utils import init_uniform_random
mdp = GridWorldMDP()
policy = init_uniform_random(mdp)
mc = MonteCarloPolicyEvaluation(mdp, policy, alpha=0.01, gamma=0.9)
mc.evaluate(max_iter=10000)
td = TDPolicyEvaluation(mdp, policy, alpha=0.01, gamma=0.9, _lambda=0.4)
td.evaluate(max_iter=10000)
# V(s) for the random policy (expected value over actions)
v_mc = mc.q.mean(axis=1).reshape(4, 4)
v_td = td.q.mean(axis=1).reshape(4, 4)
Model-Free Control
Control algorithms learn an optimal policy by interleaving evaluation and improvement on every step. Both SARSA and Q-Learning build on the TD(λ) engine, using ε-greedy exploration to balance exploitation with discovery of new state-action pairs.
SARSA
On-policy TD control. Bootstraps from a sampled next action A' drawn from the current policy — the name comes from the quintuple (S, A, R, S', A'). Because the bootstrap target reflects the exploratory policy, SARSA's Q values account for the cost of occasional random actions.
δ = R + γ Q(S', A') - Q(S, A)
Sarsa(mdp, alpha, gamma)
| Argument | Type | Default | Description |
|---|---|---|---|
mdp |
MDP | MDP instance to sample episodes from | |
alpha |
float |
0.01 |
Learning rate for incremental Q updates |
gamma |
float |
0.9 |
Discount factor |
evaluate(max_iter)
| Argument | Type | Default | Description |
|---|---|---|---|
max_iter |
int |
5000 |
Number of episodes to run |
Q-Learning
Off-policy TD control. Bootstraps from the greedy action max_a Q(S', a) regardless of the action actually taken. This means Q-Learning converges to the optimal Q* even while following an exploratory ε-greedy policy.
δ = R + γ max_a Q(S', a) - Q(S, A)
QLearning(mdp, alpha, gamma)
| Argument | Type | Default | Description |
|---|---|---|---|
mdp |
MDP | MDP instance to sample episodes from | |
alpha |
float |
0.01 |
Learning rate for incremental Q updates |
gamma |
float |
0.9 |
Discount factor |
evaluate(max_iter)
| Argument | Type | Default | Description |
|---|---|---|---|
max_iter |
int |
5000 |
Number of episodes to run |
Examples
from samsara_rl.mdp.grid_world.grid_world_gym import GridWorldMDP
from samsara_rl.control.tabular.sarsa import Sarsa
from samsara_rl.control.tabular.q_learning import QLearning
mdp = GridWorldMDP()
sarsa = Sarsa(mdp, alpha=0.01, gamma=0.9)
sarsa.evaluate(max_iter=5000)
ql = QLearning(mdp, alpha=0.01, gamma=0.9)
ql.evaluate(max_iter=5000)
# Optimal value per state (best action)
v_sarsa = sarsa.q.max(axis=1).reshape(4, 4)
v_ql = ql.q.max(axis=1).reshape(4, 4)
Function Approximation
Tabular methods store one value per state-action pair — this breaks down when the state space is large or continuous (e.g. CartPole's 4D observation vector). Function approximation replaces the Q-table with a parameterized function Q(s, a; w) that generalizes across states.
Linear Function Approximation
LinearFunction implements Q(s) = X(s)^T W, where X is a user-provided feature extraction function and W is a learned weight matrix. It exposes a PyTorch-style interface: forward pass via __call__, gradient computation via backward(), and parameter access via params.
For discrete environments, a one-hot encoding X(s) gives the linear approximator the same representational power as a tabular method — useful as a sanity check before moving to richer feature representations.
LinearFunction(feature_count, action_count, X, use_bias)
| Argument | Type | Default | Description |
|---|---|---|---|
feature_count |
int |
Number of input features (output dimension of X) | |
action_count |
int |
Number of discrete actions | |
X |
Callable |
identity |
Feature extraction function: state → feature vector |
use_bias |
bool |
False |
Whether to include a bias term per action |
Semi-Gradient TD(λ) Control
TemporalDifferenceGradient implements semi-gradient TD(λ) control with eligibility traces. On each step, the TD error is computed and used to update the function approximator's parameters in the direction of the gradient, scaled by eligibility traces that assign credit to recently visited state-action pairs.
The TD target function is configurable — SARSA and Q-Learning are implemented as thin subclasses that fix the target.
TemporalDifferenceGradient(mdp, alpha, gamma, q, _lambda)
| Argument | Type | Default | Description |
|---|---|---|---|
mdp |
gym.Env |
Gymnasium-compatible environment | |
alpha |
float |
0.001 |
Learning rate |
gamma |
float |
1 |
Discount factor |
q |
LinearFunction |
Function approximator | |
_lambda |
float |
0.2 |
Eligibility trace decay (0 = TD(0), 1 = MC) |
SARSA (Function Approximation)
On-policy control. Bootstraps from Q(S', A') where A' is the action actually taken under the current ε-greedy policy.
SarsaGradient(**kwargs) — accepts the same arguments as TemporalDifferenceGradient.
Q-Learning (Function Approximation)
Off-policy control. Bootstraps from max_a Q(S', a), learning the optimal policy regardless of exploration behavior.
QLearningGradient(**kwargs) — accepts the same arguments as TemporalDifferenceGradient.
Examples
import numpy as np
from samsara_rl.mdp.grid_world.grid_world_gym import GridWorldMDP
from samsara_rl.control.function_approximation.functions.linear import LinearFunction
from samsara_rl.control.function_approximation.sarsa import SarsaGradient
mdp = GridWorldMDP()
# One-hot encoding gives tabular-equivalent capacity
def one_hot(s):
arr = np.zeros(16)
arr[int(s)] = 1
return arr
q_fn = LinearFunction(16, 4, one_hot)
# SARSA with function approximation
sarsa = SarsaGradient(mdp=mdp, gamma=0.999, q=q_fn, alpha=0.01)
sarsa.evaluate(max_iter=20000)
# Learned value per state
v = np.array([q_fn.W.value[s].max() for s in range(16)]).reshape(4, 4)
Deep Q-Network
Deep Q-Networks (DQN) replace the linear function approximator with a neural network, enabling learning in high-dimensional continuous state spaces. Two key innovations stabilize training:
- Experience Replay — transitions are stored in a replay buffer and sampled in random mini-batches, breaking temporal correlations in the training data.
- Target Network — a frozen copy of the Q-network provides stable TD targets. It is periodically updated to match the online network.
The target computation strategy is configurable. Standard DQN uses the target network for both action selection and evaluation. Double DQN decouples these — the online network selects the action, and the target network evaluates it — reducing the maximization bias that causes Q-value overestimation and training instability.
QNetwork(mdp, alpha, gamma, q, target, target_update_freq, epsilon, epsilon_decay, loss_fn, batch_size)
| Argument | Type | Default | Description |
|---|---|---|---|
mdp |
gym.Env |
Gymnasium-compatible environment | |
alpha |
float |
0.001 |
Learning rate for Adam optimizer |
gamma |
float |
1 |
Discount factor |
q |
nn.Module |
Neural network mapping states to Q values | |
target |
object |
DQNTarget |
Target computation strategy (DQNTarget or DoubleDQNTarget) |
target_update_freq |
int |
1000 |
Training steps between target network swaps |
epsilon |
float |
1 |
Initial exploration rate for ε-greedy |
epsilon_decay |
float |
0.999 |
Multiplicative decay applied to epsilon each episode |
loss_fn |
Callable |
MSELoss |
Loss function (e.g. MSELoss, HuberLoss) |
batch_size |
int |
128 |
Number of transitions per training mini-batch |
Examples
import gymnasium as gym
from samsara_rl.control.function_approximation.batch.deep_q_network.q_network import QNetwork
from samsara_rl.control.function_approximation.batch.deep_q_network.targets.double_d_target import DoubleDQNTarget
from samsara_rl.control.function_approximation.functions.neural_networks.fully_connected import FullyConnected
from samsara_rl.mdp.cart_pole.scaled_cart_pole import ScaledCartPole
env = ScaledCartPole(gym.make("CartPole-v1"))
network = FullyConnected(4, 32, 2)
# Standard DQN
agent = QNetwork(
mdp=env, gamma=0.99, q=network, alpha=0.0001,
log_dir="logs/dqn", experiment_name="cartpole_dqn",
)
agent.evaluate(max_iter=3000)
# Double DQN — swap the target strategy
agent = QNetwork(
mdp=env, gamma=0.99, q=network, alpha=0.0001,
target=DoubleDQNTarget(),
log_dir="logs/double_dqn", experiment_name="cartpole_double_dqn",
)
agent.evaluate(max_iter=3000)
For a full walkthrough with TensorBoard logging, decision surface visualization, and hyperparameter tuning, see the Deep Q-Network tutorial notebook.
Policy Gradient
Policy gradient methods learn a parameterized policy directly, rather than deriving it from a value function. The policy network outputs action probabilities via softmax, and gradient ascent maximizes the expected return. Unlike value-based methods (Q-Learning, DQN), policy gradients can naturally represent stochastic policies and scale to continuous action spaces.
Monte Carlo Policy Gradient
Accumulates policy gradients over a batch of episodes before performing an optimizer step. For each episode, discounted returns are computed for every time step, and the policy gradient loss weights the log-probability of each taken action by its return. An optional running average baseline reduces variance by centering returns around their expected value.
MonteCarloPolicyGradient(mdp, alpha, gamma, policy_network, optimizer, batch_size, use_advantage)
| Argument | Type | Default | Description |
|---|---|---|---|
mdp |
gym.Env |
Gymnasium-compatible environment | |
alpha |
float |
Learning rate | |
gamma |
float |
Discount factor | |
policy_network |
nn.Module |
Neural network that maps states to action logits | |
optimizer |
Optimizer |
Adam |
Optimizer for the policy network |
batch_size |
int |
1 |
Number of episodes to accumulate gradients over before stepping |
use_advantage |
bool |
True |
Subtract a running average baseline to reduce variance |
REINFORCE
REINFORCE (Williams, 1992) is a special case of Monte Carlo Policy Gradient where the optimizer steps after every episode (batch_size=1).
Reinforce(mdp, alpha, gamma, policy_network, optimizer, use_advantage)
Accepts the same arguments as MonteCarloPolicyGradient, without batch_size.
Examples
import gymnasium as gym
from samsara_rl.control.function_approximation.batch.monte_carlo_policy_gradient.monte_carlo_policy_gradient import MonteCarloPolicyGradient
from samsara_rl.control.function_approximation.online.reinforce.reinforce import Reinforce
from samsara_rl.control.function_approximation.functions.neural_networks.fully_connected import FullyConnected
from samsara_rl.mdp.cart_pole.scaled_cart_pole import ScaledCartPole
env = ScaledCartPole(gym.make("CartPole-v1"))
network = FullyConnected(4, 32, 2)
# Monte Carlo Policy Gradient with batch of 32 episodes
agent = MonteCarloPolicyGradient(
mdp=env, gamma=0.99, alpha=0.002, policy_network=network, batch_size=32,
log_dir="logs/mc_pg", experiment_name="cartpole_mcpg",
)
agent.evaluate(max_iter=3000)
# REINFORCE — updates every episode
agent = Reinforce(
mdp=env, gamma=0.99, alpha=0.002, policy_network=network,
log_dir="logs/reinforce", experiment_name="cartpole_reinforce",
)
agent.evaluate(max_iter=3000)
For a full walkthrough, see the REINFORCE tutorial notebook.
Metadata
Release files for samsara-rl 0.0.9
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| samsara_rl-0.0.9.tar.gz | 944.2 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| samsara_rl-0.0.9-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 985.2 kB
Release files / samsara_rl-0.0.9.tar.gz
| Download URL | samsara_rl-0.0.9.tar.gz |
|---|---|
| Size | 944.2 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
5c87fa2a75f78bc0f4f8b1a90a282fac847fa998d8ead5c3678d2d9d4539c782
|
|
BLAKE2b-256 checksum How to use checksums |
cc1655a28a4f51d182f7518feec9cb5f9b60fa4532efcdffc128b2bb7bb5f51a
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
uv/0.10.12 {"installer":{"name":"uv","version":"0.10.12","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"Ubuntu","version":"24.04","id":"noble","libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":true}
|
Release files / samsara_rl-0.0.9-py3-none-any.whl
| Download URL | samsara_rl-0.0.9-py3-none-any.whl |
|---|---|
| Size | 41.0 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
c2fbc9793e5f43d904a2df83ca613ebeaf630f5a245ddfa9309f57907c45c527
|
|
BLAKE2b-256 checksum How to use checksums |
60f02a15dbaaa2040ddbd80ef59dec1259ebe886ae9d0e2b912e28dde850f4b5
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
uv/0.10.12 {"installer":{"name":"uv","version":"0.10.12","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"Ubuntu","version":"24.04","id":"noble","libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":true}
|