Offline Reinforcement Learning
# Offline Reinforcement Learning
## Introduction & Motivation
Offline RL: learn from fixed batch of trajectories without environment interaction. Addresses safety-critical applications. Avoids distribution shift via conservative value estimation.
Motivation: Enable learning from static datasets in safety-critical domains.
Applications: Robotics, autonomous systems, healthcare.
---
## Core Concepts & Theory
### Batch RL
Learn from collected experience only.
### Conservative Q-Learning
Prevent overestimation from OOD actions.
### Behavior Cloning Regularization
Stay close to data distribution.
### Uncertainty Estimation
Identify low-confidence predictions.
---
## Mathematical Formulation
Conservative Q-Learning:
$$Q(s,a) \leftarrow r + \gamma \min_a Q'(s,a) - \alpha \mathbb{E}_{a \sim \pi(s)}[Q(s,a) - Q(s,a_{ ext{data}})]$$
CQL Loss:
$$L = (Q(s,a_{ ext{data}}) - r - \gamma Q'(s',a'))^2 + \alpha \log \sum_a e^{Q(s,a)}$$
---
## Advanced Theory & Extensions
### Batch Normalization
Normalize Q-values.
### Reward Pessimism
Conservative value estimates.
### Uncertainty-Aware
Ensemble-based confidence.
---
## Computational Considerations
Batch Size: All data in memory.
Iterations: Multiple passes over data.
Total: O(B·T·D) where B is batch size.
---
## Practical Implementation Strategies
### Data Filtering
Remove low-quality trajectories.
### Q-Value Stabilization
Gradient clipping, target networks.
### Policy Constraints
Stay in data distribution.
---
## Benchmark Datasets & Evaluation
D4RL: Offline RL benchmark suite.
Robotics Tasks: Continuous control.
Offline Eval: No environment interaction.
---
## Key Challenges & Limitations
### Distribution Shift
Learning from limited data.
### Value Overestimation
Q-learning overoptimism.
### Data Quality
Sensitive to demonstration quality.
---
## Hyperparameter Tuning
α (penalty): 0.5-1.0.
Q-update steps: 1-10 per batch.
Learning rate: 1e-4 to 1e-3.
---
## Real-World Applications & Case Studies
Medical Treatment: Personalized medicine.
Robotics: Offline skill learning.
Autonomous Systems: Safe offline learning.
---
## Integration with Other Methods
Offline RL + uncertainty estimation; + ensemble learning; + value normalization.
---
## Summary & Key Takeaways
Offline RL learns safely from static datasets.
Principles:
1. Batch: Fixed experience only.
2. Conservative: Prevent overestimation.
3. Distribution: Stay in-distribution.
4. Uncertainty: Identify unreliable estimates.
5. Safety: Avoid catastrophic OOD actions.
---
## Appendix: Practical Labs
### Lab 1: Conservative Q-Learning
import numpy as np
def conservative_q_learning_update(Q, state, action, reward, next_state,
data_actions, gamma=0.99, alpha=0.5):
"""CQL update with penalty for OOD actions"""
# Standard Q-learning
q_next = np.max(Q[next_state])
q_target = reward + gamma * q_next
# CQL penalty: penalize high Q-values for non-data actions
penalty = 0
for a in range(len(data_actions)):
if a not in data_actions:
penalty += Q[state, a]
# Conservative update
cql_target = q_target - alpha * penalty / len(data_actions)
# TD update
td_error = cql_target - Q[state, action]
Q[state, action] += 0.01 * td_error
return Q
Q = np.random.randn(5, 4)
state, action = 0, 1
reward, next_state = 1.0, 2
data_actions = [0, 1]
Q = conservative_q_learning_update(Q, state, action, reward, next_state, data_actions)
print(f"✓ CQL update: Q[0,1]={Q[0,1]:.3f}")### Lab 2: Offline Dataset Analysis
import numpy as np
def analyze_offline_dataset(trajectories):
"""Analyze characteristics of offline dataset"""
total_transitions = sum(len(traj[0]) for traj in trajectories)
# Compute statistics
all_rewards = []
action_counts = {}
for states, actions, rewards in trajectories:
all_rewards.extend(rewards)
for action in actions:
action_counts[action] = action_counts.get(action, 0) + 1
# Summary statistics
stats = {
'num_trajectories': len(trajectories),
'total_transitions': total_transitions,
'avg_trajectory_length': total_transitions / len(trajectories),
'mean_reward': np.mean(all_rewards),
'std_reward': np.std(all_rewards),
'action_distribution': action_counts,
}
return stats
# Synthetic dataset
trajs = [
(np.random.randn(5, 10), np.random.randint(0, 4, 5), np.random.randn(5))
for _ in range(10)
]
stats = analyze_offline_dataset(trajs)
print(f"✓ Dataset: {stats['num_trajectories']} trajectories, {stats['total_transitions']} transitions")
print(f"✓ Reward: mean={stats['mean_reward']:.2f}, std={stats['std_reward']:.2f}")### Lab 3: Batch Normalization for Q-Values
import numpy as np
def normalize_q_values(Q_values, buffer_mean=None, buffer_std=None, update_stats=False):
"""Normalize Q-values based on buffer statistics"""
if buffer_mean is None:
buffer_mean = np.mean(Q_values)
if buffer_std is None:
buffer_std = np.std(Q_values)
# Normalize
q_normalized = (Q_values - buffer_mean) / (buffer_std + 1e-8)
return q_normalized, buffer_mean, buffer_std
Q = np.random.randn(100) * 5 + 10
Q_norm, mean, std = normalize_q_values(Q)
assert abs(np.mean(Q_norm)) < 0.01
assert abs(np.std(Q_norm) - 1.0) < 0.01
print(f"✓ Q-normalization: mean={np.mean(Q_norm):.3f}, std={np.std(Q_norm):.3f}")### Lab 4: Offline RL Agent
import numpy as np
class OfflineRLAgent:
def __init__(self, state_dim=10, action_dim=4):
self.state_dim = state_dim
self.action_dim = action_dim
# Q-network
self.Q = np.random.randn(state_dim, action_dim) * 0.01
# Offline buffer
self.buffer = []
def add_to_buffer(self, state, action, reward, next_state, done):
"""Add transition to replay buffer"""
self.buffer.append((state, action, reward, next_state, done))
def conservative_q_update(self, alpha=0.5, learning_rate=0.01):
"""Offline Q-learning with CQL"""
if len(self.buffer) == 0:
return 0
# Sample from buffer
idx = np.random.randint(len(self.buffer))
state, action, reward, next_state, done = self.buffer[idx]
# Q-target with conservative penalty
q_next = np.max(self.Q[next_state]) if not done else 0
q_target = reward + 0.99 * q_next
# CQL penalty (penalize high Q for all actions)
q_penalty = alpha * np.mean(self.Q[state])
# Conservative update
cql_target = q_target - q_penalty
td_error = cql_target - self.Q[state, action]
self.Q[state, action] += learning_rate * td_error
return abs(td_error)
def select_action(self, state):
"""Select greedy action from Q"""
action = np.argmax(self.Q[state])
return action
agent = OfflineRLAgent()
# Build dataset
for _ in range(50):
s = np.random.randint(0, 10)
a = np.random.randint(0, 4)
r = float(a == 0)
s_next = np.random.randint(0, 10)
agent.add_to_buffer(s, a, r, s_next, False)
# Train
for _ in range(100):
agent.conservative_q_update()
action = agent.select_action(np.random.randint(0, 10))
print(f"✓ Offline RL agent: policy learned, Q shape={agent.Q.shape}")---