Lint Python files with Yapf (#1872)

This commit is contained in:
Philipp Moritz
2018-04-11 10:11:35 -07:00
committed by Robert Nishihara
parent a3ddde398c
commit 74162d1492
97 changed files with 3927 additions and 3139 deletions
+22 -20
View File
@@ -35,10 +35,13 @@ class AsyncHyperBandScheduler(FIFOScheduler):
halving rate, specified by the reduction factor.
"""
def __init__(
self, time_attr='training_iteration',
reward_attr='episode_reward_mean', max_t=100,
grace_period=10, reduction_factor=3, brackets=3):
def __init__(self,
time_attr='training_iteration',
reward_attr='episode_reward_mean',
max_t=100,
grace_period=10,
reduction_factor=3,
brackets=3):
assert max_t > 0, "Max (time_attr) not valid!"
assert max_t >= grace_period, "grace_period must be <= max_t!"
assert grace_period > 0, "grace_period must be positive!"
@@ -51,8 +54,10 @@ class AsyncHyperBandScheduler(FIFOScheduler):
self._trial_info = {} # Stores Trial -> Bracket
# Tracks state for new trial add
self._brackets = [_Bracket(
grace_period, max_t, reduction_factor, s) for s in range(brackets)]
self._brackets = [
_Bracket(grace_period, max_t, reduction_factor, s)
for s in range(brackets)
]
self._counter = 0 # for
self._num_stopped = 0
self._reward_attr = reward_attr
@@ -60,7 +65,7 @@ class AsyncHyperBandScheduler(FIFOScheduler):
def on_trial_add(self, trial_runner, trial):
sizes = np.array([len(b._rungs) for b in self._brackets])
probs = np.e ** (sizes - sizes.max())
probs = np.e**(sizes - sizes.max())
normalized = probs / probs.sum()
idx = np.random.choice(len(self._brackets), p=normalized)
self._trial_info[trial.trial_id] = self._brackets[idx]
@@ -71,28 +76,23 @@ class AsyncHyperBandScheduler(FIFOScheduler):
action = TrialScheduler.STOP
else:
bracket = self._trial_info[trial.trial_id]
action = bracket.on_result(
trial,
getattr(result, self._time_attr),
getattr(result, self._reward_attr))
action = bracket.on_result(trial, getattr(result, self._time_attr),
getattr(result, self._reward_attr))
if action == TrialScheduler.STOP:
self._num_stopped += 1
return action
def on_trial_complete(self, trial_runner, trial, result):
bracket = self._trial_info[trial.trial_id]
bracket.on_result(
trial,
getattr(result, self._time_attr),
getattr(result, self._reward_attr))
bracket.on_result(trial, getattr(result, self._time_attr),
getattr(result, self._reward_attr))
del self._trial_info[trial.trial_id]
def on_trial_remove(self, trial_runner, trial):
del self._trial_info[trial.trial_id]
def debug_string(self):
out = "Using AsyncHyperBand: num_stopped={}".format(
self._num_stopped)
out = "Using AsyncHyperBand: num_stopped={}".format(self._num_stopped)
out += "\n" + "\n".join([b.debug_str() for b in self._brackets])
return out
@@ -111,6 +111,7 @@ class _Bracket():
>>> b.on_result(trial3, 1, 1) # STOP
>>> b.cutoff(b._rungs[0][1]) == 2.0
"""
def __init__(self, min_t, max_t, reduction_factor, s):
self.rf = reduction_factor
MAX_RUNGS = int(np.log(max_t / min_t) / np.log(self.rf) - s + 1)
@@ -140,9 +141,10 @@ class _Bracket():
return action
def debug_str(self):
iters = " | ".join(
["Iter {:.3f}: {}".format(milestone, self.cutoff(recorded))
for milestone, recorded in self._rungs])
iters = " | ".join([
"Iter {:.3f}: {}".format(milestone, self.cutoff(recorded))
for milestone, recorded in self._rungs
])
return "Bracket: " + iters