| 32 | |
| 33 | @ray.remote |
| 34 | class SharedStorage(object): |
| 35 | def __init__(self, model, target_model): |
| 36 | """Shared storage for models and others |
| 37 | Parameters |
| 38 | ---------- |
| 39 | model: any |
| 40 | models for self-play (update every checkpoint_interval) |
| 41 | target_model: any |
| 42 | models for reanalyzing (update every target_model_interval) |
| 43 | """ |
| 44 | self.step_counter = 0 |
| 45 | self.test_counter = 0 |
| 46 | self.model = model |
| 47 | self.target_model = target_model |
| 48 | self.ori_reward_log = [] |
| 49 | self.reward_log = [] |
| 50 | self.reward_max_log = [] |
| 51 | self.test_dict_log = {} |
| 52 | self.eps_lengths = [] |
| 53 | self.eps_lengths_max = [] |
| 54 | self.temperature_log = [] |
| 55 | self.visit_entropies_log = [] |
| 56 | self.priority_self_play_log = [] |
| 57 | self.distributions_log = {} |
| 58 | self.start = False |
| 59 | |
| 60 | def set_start_signal(self): |
| 61 | self.start = True |
| 62 | |
| 63 | def get_start_signal(self): |
| 64 | return self.start |
| 65 | |
| 66 | def get_weights(self): |
| 67 | return self.model.get_weights() |
| 68 | |
| 69 | def set_weights(self, weights): |
| 70 | return self.model.set_weights(weights) |
| 71 | |
| 72 | def get_target_weights(self): |
| 73 | return self.target_model.get_weights() |
| 74 | |
| 75 | def set_target_weights(self, weights): |
| 76 | return self.target_model.set_weights(weights) |
| 77 | |
| 78 | def incr_counter(self): |
| 79 | self.step_counter += 1 |
| 80 | |
| 81 | def get_counter(self): |
| 82 | return self.step_counter |
| 83 | |
| 84 | def set_data_worker_logs(self, eps_len, eps_len_max, eps_ori_reward, eps_reward, eps_reward_max, temperature, visit_entropy, priority_self_play, distributions): |
| 85 | self.eps_lengths.append(eps_len) |
| 86 | self.eps_lengths_max.append(eps_len_max) |
| 87 | self.ori_reward_log.append(eps_ori_reward) |
| 88 | self.reward_log.append(eps_reward) |
| 89 | self.reward_max_log.append(eps_reward_max) |
| 90 | self.temperature_log.append(temperature) |
| 91 | self.visit_entropies_log.append(visit_entropy) |
nothing calls this directly
no outgoing calls
no test coverage detected