# Agent SAC Discret
class ReplayBuffer:
def __init__(self, capacity=50000):
self.buffer = deque(maxlen=capacity)
def push(self, state, action, reward, next_state, done):
self.buffer.append((state, action, reward, next_state, done))
def sample(self, batch_size):
indices = np.random.choice(len(self.buffer), batch_size, replace=False)
batch = [self.buffer[i] for i in indices]
states, actions, rewards, next_states, dones = zip(*batch)
return (np.array(states, dtype=np.float32),
np.array(actions, dtype=np.int64),
np.array(rewards, dtype=np.float32),
np.array(next_states, dtype=np.float32),
np.array(dones, dtype=np.float32))
def __len__(self):
return len(self.buffer)
class SACDiscreteAgent:
"""SAC pour actions discretes avec double Q et auto-temperature."""
def __init__(self, state_dim, action_dim, hidden_dim=128, lr=3e-4,
gamma=0.99, tau=0.005, alpha_lr=3e-4, target_entropy_ratio=0.5):
self.gamma = gamma
self.tau = tau
self.action_dim = action_dim
# Q-networks (double)
self.q1 = nn.Sequential(
nn.Linear(state_dim, hidden_dim), nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim), nn.ReLU(),
nn.Linear(hidden_dim, action_dim)
).to(DEVICE)
self.q2 = nn.Sequential(
nn.Linear(state_dim, hidden_dim), nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim), nn.ReLU(),
nn.Linear(hidden_dim, action_dim)
).to(DEVICE)
# Target networks
self.q1_target = nn.Sequential(
nn.Linear(state_dim, hidden_dim), nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim), nn.ReLU(),
nn.Linear(hidden_dim, action_dim)
).to(DEVICE)
self.q2_target = nn.Sequential(
nn.Linear(state_dim, hidden_dim), nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim), nn.ReLU(),
nn.Linear(hidden_dim, action_dim)
).to(DEVICE)
self.q1_target.load_state_dict(self.q1.state_dict())
self.q2_target.load_state_dict(self.q2.state_dict())
# Policy network
self.policy = nn.Sequential(
nn.Linear(state_dim, hidden_dim), nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim), nn.ReLU(),
).to(DEVICE)
self.policy_head = nn.Linear(hidden_dim, action_dim).to(DEVICE)
# Auto-temperature
target_entropy = -np.log(1.0 / action_dim) * target_entropy_ratio
self.log_alpha = torch.zeros(1, requires_grad=True, device=DEVICE)
self.target_entropy = target_entropy
# Optimizers
self.q_optimizer = torch.optim.AdamW(
list(self.q1.parameters()) + list(self.q2.parameters()),
lr=lr, weight_decay=1e-4
)
self.policy_optimizer = torch.optim.AdamW(
list(self.policy.parameters()) + list(self.policy_head.parameters()),
lr=lr, weight_decay=1e-4
)
self.alpha_optimizer = torch.optim.AdamW([self.log_alpha], lr=alpha_lr)
self.replay_buffer = ReplayBuffer(capacity=50000)
@property
def alpha(self):
return self.log_alpha.exp().item()
def get_action(self, state, explore=True):
state_t = torch.tensor(state, dtype=torch.float32).unsqueeze(0).to(DEVICE)
with torch.no_grad():
features = self.policy(state_t)
logits = self.policy_head(features)
probs = F.softmax(logits, dim=-1)
if explore:
dist = Categorical(probs=probs)
action = dist.sample()
else:
action = probs.argmax(dim=-1)
return action.item()
def update(self, batch_size=64):
if len(self.replay_buffer) < batch_size:
return {"q_loss": 0, "policy_loss": 0, "alpha": self.alpha}
states, actions, rewards, next_states, dones = self.replay_buffer.sample(batch_size)
states_t = torch.tensor(states).to(DEVICE)
actions_t = torch.tensor(actions).to(DEVICE)
rewards_t = torch.tensor(rewards).to(DEVICE)
next_states_t = torch.tensor(next_states).to(DEVICE)
dones_t = torch.tensor(dones).to(DEVICE)
with torch.no_grad():
next_features = self.policy(next_states_t)
next_logits = self.policy_head(next_features)
next_probs = F.softmax(next_logits, dim=-1)
next_log_probs = F.log_softmax(next_logits, dim=-1)
next_q1 = self.q1_target(next_states_t)
next_q2 = self.q2_target(next_states_t)
next_q = torch.min(next_q1, next_q2)
next_v = (next_probs * (next_q - self.alpha * next_log_probs)).sum(dim=-1)
target_q = rewards_t + self.gamma * (1 - dones_t) * next_v
q1_values = self.q1(states_t).gather(1, actions_t.unsqueeze(-1)).squeeze(-1)
q2_values = self.q2(states_t).gather(1, actions_t.unsqueeze(-1)).squeeze(-1)
q_loss = F.mse_loss(q1_values, target_q) + F.mse_loss(q2_values, target_q)
self.q_optimizer.zero_grad()
q_loss.backward()
self.q_optimizer.step()
features = self.policy(states_t)
logits = self.policy_head(features)
probs = F.softmax(logits, dim=-1)
log_probs = F.log_softmax(logits, dim=-1)
with torch.no_grad():
q1_val = self.q1(states_t)
q2_val = self.q2(states_t)
q_val = torch.min(q1_val, q2_val)
policy_loss = (probs * (self.alpha * log_probs - q_val)).sum(dim=-1).mean()
self.policy_optimizer.zero_grad()
policy_loss.backward()
self.policy_optimizer.step()
alpha_loss = -(self.log_alpha * (log_probs + self.target_entropy).detach()).mean()
self.alpha_optimizer.zero_grad()
alpha_loss.backward()
self.alpha_optimizer.step()
for param, target_param in zip(self.q1.parameters(), self.q1_target.parameters()):
target_param.data.copy_(self.tau * param.data + (1 - self.tau) * target_param.data)
for param, target_param in zip(self.q2.parameters(), self.q2_target.parameters()):
target_param.data.copy_(self.tau * param.data + (1 - self.tau) * target_param.data)
return {"q_loss": q_loss.item(), "policy_loss": policy_loss.item(), "alpha": self.alpha}
sac_agent = SACDiscreteAgent(STATE_DIM, ACTION_DIM, hidden_dim=128)
n_params_sac = sum(p.numel() for p in list(sac_agent.q1.parameters()) + list(sac_agent.q2.parameters()) + list(sac_agent.policy.parameters()) + list(sac_agent.policy_head.parameters()))
print(f"SAC Discret: {n_params_sac:,} parametres, alpha={sac_agent.alpha:.3f}")