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

Method select_action_batch

distributed/rpc/batch/reinforce.py:146–170  ·  view source on GitHub ↗

r""" Batching select_action: In each step, the agent waits for states from all observers, and process them together. This helps to reduce the number of CUDA kernels launched and hence speed up amortized inference speed.

(agent_rref, ob_id, state)

Source from the content-addressed store, hash-verified

144 @staticmethod
145 @rpc.functions.async_execution
146 def select_action_batch(agent_rref, ob_id, state):
147 r"""
148 Batching select_action: In each step, the agent waits for states from
149 all observers, and process them together. This helps to reduce the
150 number of CUDA kernels launched and hence speed up amortized inference
151 speed.
152 """
153 self = agent_rref.local_value()
154 self.states[ob_id].copy_(state)
155 future_action = self.future_actions.then(
156 lambda future_actions: future_actions.wait()[ob_id].item()
157 )
158
159 with self.lock:
160 self.pending_states -= 1
161 if self.pending_states == 0:
162 self.pending_states = len(self.ob_rrefs)
163 probs = self.policy(self.states.cuda())
164 m = Categorical(probs)
165 actions = m.sample()
166 self.saved_log_probs.append(m.log_prob(actions).t()[0])
167 future_actions = self.future_actions
168 self.future_actions = torch.futures.Future()
169 future_actions.set_result(actions.cpu())
170 return future_action
171
172 @staticmethod
173 def select_action(agent_rref, ob_id, state):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected