import torch.optim as optim # Define the optimizer optimizer = optim.Adam(policy_network.parameters(), lr=0.001) # Training loop for epoch in range(num_epochs): optimizer.zero_grad() loss = compute_policy_loss(policy_network, trajectories) loss.backward() optimizer.step()