# Source: content/notes/deep-learning/deep-rl.md
# Independent CPU example; use the curriculum environment.
# See /notes/ml/#example-environment or /notes/deep-learning/#example-environment.

import copy
import torch
from torch import nn
from torch.distributions import Categorical, kl_divergence

torch.manual_seed(52)
torch.set_num_threads(1)
states = torch.randn(512, 4)
actor = nn.Linear(4, 3)
nn.init.zeros_(actor.weight)
nn.init.zeros_(actor.bias)
old_actor = copy.deepcopy(actor).eval()
old_actor.requires_grad_(False)
critic = nn.Linear(4, 1)
best_action = torch.stack((states[:, 0], states[:, 1], -states[:, 0]), 1).argmax(1)
with torch.no_grad():
    old_distribution = Categorical(logits=old_actor(states))
    actions = old_distribution.sample()
    old_log_prob = old_distribution.log_prob(actions)
    rewards = (actions == best_action).float()
    baseline = rewards.mean()
    advantages = rewards - baseline
    advantages = advantages / advantages.std(unbiased=False).clamp_min(1e-6)
    initial_success_probability = old_distribution.probs.gather(1, best_action[:, None]).mean().item()
opt = torch.optim.Adam(list(actor.parameters()) + list(critic.parameters()), lr=0.03)
for _ in range(12):
    distribution = Categorical(logits=actor(states))
    ratio = (distribution.log_prob(actions) - old_log_prob).exp()
    surrogate = torch.minimum(ratio * advantages,
                               ratio.clamp(0.8, 1.2) * advantages)
    policy_loss = -surrogate.mean()
    value_loss = (critic(states).squeeze(-1) - rewards).square().mean()
    loss = policy_loss + 0.5 * value_loss - 0.01 * distribution.entropy().mean()
    opt.zero_grad(set_to_none=True)
    loss.backward()
    nn.utils.clip_grad_norm_(list(actor.parameters()) + list(critic.parameters()), 1.0)
    opt.step()
with torch.no_grad():
    current = Categorical(logits=actor(states))
    success_probability = current.probs.gather(1, best_action[:, None]).mean().item()
    kl = kl_divergence(old_distribution, current).mean().item()
    ratios = (current.log_prob(actions) - old_log_prob).exp()
    clip_fraction = ((ratios < 0.8) | (ratios > 1.2)).float().mean().item()
assert success_probability > initial_success_probability
assert torch.isfinite(ratios).all() and kl >= 0
assert all(p.grad is None for p in old_actor.parameters())
print({"initial_success_probability": initial_success_probability,
       "final_success_probability": success_probability, "old_to_new_kl": kl,
       "clip_fraction": clip_fraction})
