Team Ai
Datasetpublic

Fraser/dream-coder

Program Synthesis Data Generated program synthesis datasets used to train dreamcoder. Currently just supports text & list data.

sourceHugging Facemitupdated 4y agoView on Hugging Face
6likes690downloads
likelihoodModel.py408 linesDownload Raw Back to dreamcoder
1from dreamcoder.task import Task, EvaluationTimeout2import gc3from dreamcoder.utilities import *4from collections import Counter5import math6 7from dreamcoder.domains.regex.groundtruthRegexes import gt_dict8 9gt_dict = {"Data column no. "+str(num): r_str for num, r_str in gt_dict.items()}10 11class AllOrNothingLikelihoodModel:12    def __init__(self, timeout=None):13        self.timeout = timeout14 15    def score(self, program, task):16        logLikelihood = task.logLikelihood(program, self.timeout)17        return valid(logLikelihood), logLikelihood18 19 20class EuclideanLikelihoodModel:21    """Likelihood is based on Euclidean distance between features"""22 23    def __init__(self, featureExtractor, successCutoff=0.9):24        self.extract = featureExtractor25        self.successCutoff = successCutoff26 27    def score(self, program, task):28        taskFeat = self.extract.featuresOfTask(task)29        progFeat = self.extract.featuresOfProgram(program, task.request)30        assert len(taskFeat) == len(progFeat)31        distance = sum((x1 - x2)**2 for x1, x2 in zip(taskFeat, progFeat))32        logLikelihood = float(-distance)  # FIXME: this is really naive33        return exp(logLikelihood) > self.successCutoff, logLikelihood34 35def longest_common_substr(arr):36    #array of examples 37 38# Python 3 program to find the stem39# of given list of words40# function to find the stem (longest41# common substring) from the string array42    # Determine size of the array43    n = len(arr)44 45    # Take first word from array46    # as reference47    s = arr[0]48    l = len(s)49    res = ""50    for i in range(l) :51        for j in range( i + 1, l + 1) :52            # generating all possible substrings53            # of our reference string arr[0] i.e s54            stem = s[i:j]55            k = 156            for k in range(1, n):57 58                # Check if the generated stem is59                # common to to all words60                if stem not in arr[k]:61                    break62 63                # If current substring is present in64                # all strings and its length is greater65                # than current result66            if (k + 1 == n and len(res) < len(stem)): res = stem67    return res 68 69def add_string_constants(tasks):70    for task in tasks:71        task.str_const = longest_common_substr([example[1] for example in task.examples])72    return tasks73 74def get_gt_ll(name, examples):75    #gets groundtruth from dict76    import pregex as pre77    r_str = gt_dict[name]78    preg = pre.create(r_str)79 80    if type(examples[0]) == list:81        examples = [ "".join(example) for example in examples]82 83    s = sum( preg.match(example) for example in examples)84    if s == float("-inf"):85        print("bad for ", name)86        print('preg:', preg)87        print('preg sample:', [preg.sample() for i in range(3)])88        print("exs", examples)89        #assert False 90    return s91 92 93def add_cutoff_values(tasks, ll_cutoff):94    from dreamcoder.domains.regex.makeRegexTasks import makeNewTasks95    if ll_cutoff is None or ll_cutoff == "None":96        for task in tasks:97            task.ll_cutoff = None98        return tasks99    if ll_cutoff == "gt":100        from dreamcoder.domains.regex.makeRegexTasks import regexHeldOutExamples101        for task in tasks:102            task.ll_cutoff = None103            task.gt = get_gt_ll(task.name, [example[1] for example in task.examples])104            task.gt_test = get_gt_ll(task.name,105                                     [example[1] for example in regexHeldOutExamples(task) ])106        return tasks107    elif ll_cutoff == "plus":108        for task in tasks:109            task.ll_cutoff = regex_plus_bound([example[1] for example in task.examples])110        return tasks111    elif ll_cutoff == "bigram":112        eprint("WARNING: using entire corpus to make bigram model")113        #this means i do it twice, which is eh whatever114        model = make_corpus_bigram(show_tasks(makeNewTasks()))115        for task in tasks:116            task.ll_cutoff = bigram_corpus_score([example[1] for example in task.examples], model)117        return tasks118    elif ll_cutoff =="unigram":119        eprint("WARNING: using entire corpus to make unigram model")120        #this means i do it twice, which is eh whatever121        model = make_corpus_unigram(show_tasks(makeNewTasks()))122        for task in tasks:123            task.ll_cutoff = unigram_corpus_score([example[1] for example in task.examples], model)124        return tasks125    elif ll_cutoff =="mix":126        eprint("WARNING: using entire corpus to make bigram model")127        eprint("WARNING: using entire corpus to make unigram model")128        #this means i do it twice, which is eh whatever129        unigram = make_corpus_unigram(show_tasks(makeNewTasks()))130        bigram = make_corpus_bigram(show_tasks(makeNewTasks()))131        for task in tasks:132            uniscore = unigram_corpus_score([example[1] for example in task.examples], unigram)133            biscore = bigram_corpus_score([example[1] for example in task.examples], bigram)134            task.ll_cutoff = math.log(0.75*math.exp(biscore) + 0.25*math.exp(uniscore))135        return tasks136    else:137        eprint("not implemented")138        eprint("cutoff val:")139        eprint(ll_cutoff)140        assert False141 142def show_tasks(dataset):143    task_list = []144    for task in dataset:145        task_list.append([example[1] for example in task.examples])146    return task_list147 148def regex_plus_bound(X):149    from pregex import pregex150    c = Counter(X)151    regexes = [152        pregex.create(".+"),153        pregex.create("\d+"),154        pregex.create("\w+"),155        pregex.create("\s+"),156        pregex.create("\\u+"),157        pregex.create("\l+")]158    regex_scores = []159    for r in regexes:160        regex_scores.append(sum(c[x] * r.match(x) for x in c)/float(sum([len(x) for x in X])) )161    return max(regex_scores)162 163 164def make_corpus_unigram(C):165    str_list = [example + '\n' for task in C for example in task]166    c = Counter(char for example in str_list for char in example )167    n = sum(c.values())168 169    logp = {x:math.log(c[x]/n) for x in c}170    return logp171 172def unigram_corpus_score(X, logp):173    task_ll = 0174    for x in X:175        x = x + '\n'176        task_ll += sum( logp.get(c, float('-inf')) for c in x)/len(x)177 178    ll = task_ll/len(X)179    return ll180 181def unigram_task_score(X):182    """183    Given a list of strings, X, calculate the maximum log-likelihood per character for a unigram model over characters (including STOP symbol)184    """185    c = Counter(x for s in X for x in s)186    c.update("end" for s in X)187    n = sum(c.values())188    logp = {x:math.log(c[x]/n) for x in c}189    return sum(c[x]*logp[x] for x in c)/n190 191def make_corpus_bigram(C):192    #using newline as "end"193    #C is a list of tasks194 195    #make one big list of strings196    str_list = [example + '\n' for task in C for example in task]197 198    #make list of 199    head_count = Counter(element[0] for element in str_list)200    head_n = sum(head_count.values())201    head_logp = {x:math.log(head_count[x]/head_n) for x in head_count}202 203    body_count = Counter(element[i:i+2] for element in str_list for i in range(len(element)-1))204    body_bigram_n = sum(body_count.values())205    #body_count/body_bigram_n gives the joint of a bigram206    body_character_n = Counter(char for element in str_list for char in element)207    body_unigram_n = sum(body_character_n.values())208 209    body_logp = {x:math.log(body_count[x] / body_bigram_n / body_character_n[x[0]] * body_unigram_n) for x in body_count}210 211    return {**head_logp, **body_logp}212 213def bigram_corpus_score(X, logp):214    #assume you have a logp dict215    task_ll = 0216    for x in X:217        bigram_list = [x[0]] + [x[i:i+2] for i in range(len(x)-1)] + [x[-1] + '\n']218        bigram_list = [ ''.join(b) if isinstance(b,list) else b219                        for b in bigram_list ]220 221        string_ll = sum(logp.get(bigram, float('-inf')) for bigram in bigram_list) #/(len(x) + 1)222 223        task_ll += string_ll224 225    ll = task_ll #/len(X)226    return ll227 228 229class ProbabilisticLikelihoodModel:230 231    def __init__(self, timeout):232        self.timeout = timeout233        # i need timeout234 235    def score(self, program, task):236        # need a try, catch here for problems, and for timeouts237        # can copy task.py for the timeout structure238        try:239            def timeoutCallBack(_1, _2): raise EvaluationTimeout()240            signal.signal(signal.SIGVTALRM, timeoutCallBack)241            signal.setitimer(signal.ITIMER_VIRTUAL, self.timeout)242            try:243                string_pregex = program.evaluate([])244                # if 'left_paren' in program.show(False):245                #eprint("string_pregex:", string_pregex)246                #eprint("string_pregex:", string_pregex)247                preg = string_pregex  # pregex.create(string_pregex)248            except IndexError:249                # free variable250                return False, NEGATIVEINFINITY251            except Exception as e:252                eprint("Exception during evaluation:", e)253                if "Attempt to evaluate fragment variable" in e:254                    eprint("program (bc fragment error)", program)255                return False, NEGATIVEINFINITY256 257        #tries and catches258 259        # include prior somehow260        # right now, just summing up log likelihoods. IDK if this is correct.261        # also not using prior at all.262 263            cum_ll = 0264 265            example_list = [example[1] for example in task.examples]266            c_example_list = Counter(example_list)267 268            for c_example in c_example_list:269                #might want a try, except around the following line:270 271                try:272                    #eprint("about to match", program)273                    #print("preg:", preg)274                    ll = preg.match(c_example)275                    #eprint("completed match", ll, program)276                except ValueError as e:277                    eprint("ValueError:", e)278                    ll = float('-inf')279                280                #eprint("pregex:", string_pregex)281                #eprint("example[1]", example[1])282 283                if ll == float('-inf'):284                    return False, NEGATIVEINFINITY285                else:286                    #ll_per_char = ll/float(len(example[1]))287                    #cum_ll_per_char += ll_per_char288 289                    cum_ll += c_example_list[c_example] * ll290            291            #normalized_cum_ll_per_char = cum_ll_per_char/float(len(task.examples))292            #avg_char_num = sum([len(example[1]) for example in task.examples])/float(len(task.examples))293            294            #cutoff_ll = regex_plus_bound(example_list)   295 296            normalized_cum_ll = cum_ll/ float(sum([len(example) for example in example_list]))297 298 299 300            #TODO: change the way normalized_cum_ll is calculated 301            #TODO: refactor to pass in bigram_model, and others302            #TODO: refactor to do 95% certainty thing josh wants303            success = normalized_cum_ll > task.ll_cutoff304 305 306 307            #eprint("cutoff_ll:", cutoff_ll, ", norm_cum_ll:", normalized_cum_ll)	308 309            return success, normalized_cum_ll310 311        except EvaluationTimeout:312            eprint("Timed out while evaluating", program)313            return False, NEGATIVEINFINITY314        finally:315            signal.signal(signal.SIGVTALRM, lambda *_: None)316            signal.setitimer(signal.ITIMER_VIRTUAL, 0)317 318 319try:320    import torch321    import torch.nn as nn322    import torch.nn.functional as F323    from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence324    from torch.autograd import Variable325 326    class FeatureDiscriminatorLikelihoodModel(nn.Module):327        def __init__(self, tasks, featureExtractor,328                     successCutoff=0.6, H=8, trainingSuccessRatio=0.5):329            super(FeatureDiscriminatorLikelihoodModel, self).__init__()330            self.extract = featureExtractor331            self.successCutoff = successCutoff332            self.trainingSuccessRatio = trainingSuccessRatio333 334            self.W = nn.Linear(featureExtractor.outputDimensionality, H)335            self.output = nn.Linear(H, 1)336 337            # training on initialization338            self.train(tasks)339 340        def forward(self, examples):341            """342            Examples is a list of feature sets corresponding to a particular example.343            Output in [0,1] whether all examples correspond to the same program344            """345            assert all(346                len(x) == self.extract.outputDimensionality for x in examples)347            examples = [F.tanh(self.W(ex)) for ex in examples]348            maxed, _ = torch.max(torch.stack(examples), dim=0)349            return F.sigmoid(self.output(maxed))350 351        def train(self, tasks, steps=400):352            # list of list of features for each example in each task353            optimizer = torch.optim.Adam(self.parameters())354            with timing("Trained discriminator"):355                losses = []356                for i in range(steps):357                    self.zero_grad()358                    if random.random() <= self.trainingSuccessRatio:359                        # success360                        t = random.choice(tasks)361                        features = [self.extract.featuresOfTask(362                            Task(t.name, t.request, [ex], t.features))363                            for ex in t.examples]364                        loss = (self(features) - 1.0)**2365                    else:366                        # fail367                        t1, t2 = random.sample(tasks, 2)368                        features1 = [self.extract.featuresOfTask(369                            Task(t1.name, t1.request, [ex], t1.features))370                            for ex in t1.examples[:len(t1.examples) / 2]]371                        features2 = [self.extract.featuresOfTask(372                            Task(t2.name, t2.request, [ex], t2.features))373                            for ex in t2.examples[len(t2.examples) / 2:]]374                        features = features1 + features2375                        loss = self(features)**2376 377                    loss.backward()378                    optimizer.step()379                    losses.append(loss.data[0])380                    if not i % 50:381                        eprint(382                            "Discriminator Epoch",383                            i,384                            "Loss",385                            sum(losses) /386                            len(losses))387                        gc.collect()388 389        def score(self, program, task):390            taskFeatures = self.extract.featuresOfTask(task)391            progFeatures = self.extract.featuresOfProgram(392                program, task.request)393            likelihood = self([taskFeatures] + [progFeatures])394            likelihood = float(likelihood)395            return likelihood > self.successCutoff, log(likelihood)396except ImportError:397    pass398 399 400if __name__=="__main__":401 402    arr = ['MAM.OSBS.2014.06', 'MAM.OSBS.2013.07', 'MAM.OSBS.2013.09', 'MAM.OSBS.2014.05', 'MAM.OSBS.2014.11']403    stems = longest_common_substr(arr)404    print(stems)405 406 407 408