MCPcopy Create free account
hub / github.com/pytorch/examples / finish_episode

Function finish_episode

reinforcement_learning/reinforce.py:63–79  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

61
62
63def finish_episode():
64 R = 0
65 policy_loss = []
66 returns = deque()
67 for r in policy.rewards[::-1]:
68 R = r + args.gamma * R
69 returns.appendleft(R)
70 returns = torch.tensor(returns)
71 returns = (returns - returns.mean()) / (returns.std() + eps)
72 for log_prob, R in zip(policy.saved_log_probs, returns):
73 policy_loss.append(-log_prob * R)
74 optimizer.zero_grad()
75 policy_loss = torch.cat(policy_loss).sum()
76 policy_loss.backward()
77 optimizer.step()
78 del policy.rewards[:]
79 del policy.saved_log_probs[:]
80
81
82def main():

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected