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

Method __init__

distributed/rpc/batch/reinforce.py:120–142  ·  view source on GitHub ↗
(self, world_size, batch=True)

Source from the content-addressed store, hash-verified

118
119class Agent:
120 def __init__(self, world_size, batch=True):
121 self.ob_rrefs = []
122 self.agent_rref = RRef(self)
123 self.rewards = {}
124 self.policy = Policy(batch).cuda()
125 self.optimizer = optim.Adam(self.policy.parameters(), lr=1e-2)
126 self.running_reward = 0
127
128 for ob_rank in range(1, world_size):
129 ob_info = rpc.get_worker_info(OBSERVER_NAME.format(ob_rank))
130 self.ob_rrefs.append(remote(ob_info, Observer, args=(batch,)))
131 self.rewards[ob_info.id] = []
132
133 self.states = torch.zeros(len(self.ob_rrefs), 1, 4)
134 self.batch = batch
135 # With batching, saved_log_probs contains a list of tensors, where each
136 # tensor contains probs from all observers in one step.
137 # Without batching, saved_log_probs is a dictionary where the key is the
138 # observer id and the value is a list of probs for that observer.
139 self.saved_log_probs = [] if self.batch else {k:[] for k in range(len(self.ob_rrefs))}
140 self.future_actions = torch.futures.Future()
141 self.lock = threading.Lock()
142 self.pending_states = len(self.ob_rrefs)
143
144 @staticmethod
145 @rpc.functions.async_execution

Callers 1

__init__Method · 0.45

Calls 1

PolicyClass · 0.70

Tested by

no test coverage detected