MCPcopy Create free account
hub / github.com/YeWR/EfficientZero / SharedStorage

Class SharedStorage

core/storage.py:34–148  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

32
33@ray.remote
34class 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)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected