[tune] Reduce sampling API clutter (#4739)

Adds some sugar for tune sampling API (for commonplace sampling idioms).
This commit is contained in:
Richard Liaw
2019-05-06 17:42:39 -07:00
committed by GitHub
parent 71b2dec3b4
commit 7f50c96adb
8 changed files with 85 additions and 65 deletions
+3 -8
View File
@@ -1,16 +1,11 @@
from ray.tune.suggest.search import SearchAlgorithm
from ray.tune.suggest.basic_variant import BasicVariantGenerator
from ray.tune.suggest.suggestion import SuggestionAlgorithm
from ray.tune.suggest.variant_generator import grid_search, function, \
sample_from
from ray.tune.suggest.variant_generator import grid_search
__all__ = [
"SearchAlgorithm",
"BasicVariantGenerator",
"SuggestionAlgorithm",
"grid_search",
"function",
"sample_from",
"SearchAlgorithm", "BasicVariantGenerator", "SuggestionAlgorithm",
"grid_search"
]
+1 -43
View File
@@ -9,6 +9,7 @@ import random
import types
from ray.tune import TuneError
from ray.tune.sample import sample_from
logger = logging.getLogger(__name__)
@@ -54,49 +55,6 @@ def grid_search(values):
return {"grid_search": values}
class sample_from(object):
"""Specify that tune should sample configuration values from this function.
The use of function arguments in tune configs must be disambiguated by
either wrapped the function in tune.sample_from() or tune.function().
Arguments:
func: An callable function to draw a sample from.
"""
def __init__(self, func):
self.func = func
def __str__(self):
return "tune.sample_from({})".format(str(self.func))
def __repr__(self):
return "tune.sample_from({})".format(repr(self.func))
class function(object):
"""Wraps `func` to make sure it is not expanded during resolution.
The use of function arguments in tune configs must be disambiguated by
either wrapped the function in tune.sample_from() or tune.function().
Arguments:
func: A function literal.
"""
def __init__(self, func):
self.func = func
def __call__(self, *args, **kwargs):
return self.func(*args, **kwargs)
def __str__(self):
return "tune.function({})".format(str(self.func))
def __repr__(self):
return "tune.function({})".format(repr(self.func))
_STANDARD_IMPORTS = {
"random": random,
"np": numpy,