kernelmachine/gpt3-quality-filter
2
1from typing import Any, Dict, List, Union2import numpy as np3import logging4import os5 6# Create a custom logger7logger = logging.getLogger(__name__)8logger.setLevel(logging.DEBUG)9 10 11 12class RandomSearch:13 14 @staticmethod15 def random_choice(args: List[Any], n: int = 1):16 """17 pick a random element from a set.18 19 Example:20 >> sampler = RandomSearch.random_choice(1,2,3)21 >> sampler()22 223 """24 choices = []25 for arg in args:26 choices.append(arg)27 if n == 1:28 return lambda: np.random.choice(choices, replace=False)29 else:30 return lambda: np.random.choice(choices, n, replace=False)31 32 @staticmethod33 def random_integer(low: Union[int, float], high: Union[int, float]):34 """35 pick a random integer between two bounds36 37 Example:38 >> sampler = RandomSearch.random_integer(1, 10)39 >> sampler()40 941 """42 return lambda: int(np.random.randint(low, high))43 44 @staticmethod45 def random_loguniform(low: Union[float, int], high: Union[float, int]):46 """47 pick a random float between two bounds, using loguniform distribution48 49 Example:50 >> sampler = RandomSearch.random_loguniform(1e-5, 1e-2)51 >> sampler()52 0.000453 """54 return lambda: np.exp(np.random.uniform(np.log(low), np.log(high)))55 56 @staticmethod57 def random_uniform(low: Union[float, int], high: Union[float, int]):58 """59 pick a random float between two bounds, using uniform distribution60 61 Example:62 >> sampler = RandomSearch.random_uniform(0, 1)63 >> sampler()64 0.0165 """66 return lambda: np.random.uniform(low, high)67 68 69class HyperparameterSearch:70 71 def __init__(self, **kwargs):72 self.search_space = {}73 self.lambda_ = lambda: 074 for key, val in kwargs.items():75 self.search_space[key] = val76 77 def parse(self, val: Any):78 79 if isinstance(val, (int, np.int)):80 return int(val)81 elif isinstance(val, (float, np.float)):82 return val83 elif isinstance(val, (np.ndarray, list)):84 return " ".join(val)85 elif val is None:86 return None87 if isinstance(val, str):88 return val89 else:90 val = val()91 if isinstance(val, (int, np.int)):92 return int(val)93 elif isinstance(val, (np.ndarray, list)):94 return " ".join(val)95 else:96 return val97 98 99 def sample(self) -> Dict:100 res = {}101 for key, val in self.search_space.items():102 try:103 res[key] = self.parse(val)104 except (TypeError, ValueError) as error:105 logger.error(f"Could not parse key {key} with value {val}. {error}")106 107 return res108 109 def update_environment(self, sample) -> None:110 for key, val in sample.items():111 os.environ[key] = str(val)112 113 114SEARCH_SPACE = {115 "penalty": RandomSearch.random_choice(["l1", "l2"]),116 "C": RandomSearch.random_uniform(0, 1),117 "solver": "liblinear",118 "multi_class": "auto",119 "tol": RandomSearch.random_loguniform(10e-5, 10e-3),120 "stopwords": RandomSearch.random_choice([0, 1]),121 "weight": RandomSearch.random_choice(["hash"]),122 "ngram_range": RandomSearch.random_choice(["1 2", "2 3", "1 3"]),123 "random_state": RandomSearch.random_integer(0, 100000)124}125 