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)
| 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): |
nothing calls this directly
no outgoing calls
no test coverage detected