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}")

---

Go deeper with CFSGPT

Get AI-powered deep-dives, save terms, and run advanced simulations — free account.

Create Free Account