Team Ai
Apppublic

kernelmachine/gpt3-quality-filter

sourceHugging Faceupdated 5y agoView on Hugging Face
2likes
hyperparameters.py125 linesDownload Raw Back to lr
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