CoolFace
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
6likes730downloads
recognition.py1529 linesDownload Raw Back to dreamcoder
1from dreamcoder.enumeration import *2from dreamcoder.grammar import *3# luke4 5 6import gc7 8try:9    import torch10    import torch.nn as nn11    import torch.nn.functional as F12    from torch.autograd import Variable13    from torch.nn.utils.rnn import pack_padded_sequence14except:15    eprint("WARNING: Could not import torch. This is only okay when doing pypy compression.")16 17try:18    import numpy as np19except:20    eprint("WARNING: Could not import np. This is only okay when doing pypy compression.")21    22import json23 24 25def variable(x, volatile=False, cuda=False):26    if isinstance(x, list):27        x = np.array(x)28    if isinstance(x, (np.ndarray, np.generic)):29        x = torch.from_numpy(x)30    if cuda:31        x = x.cuda()32    return Variable(x, volatile=volatile)33 34def maybe_cuda(x, use_cuda):35    if use_cuda:36        return x.cuda()37    else:38        return x39 40 41def is_torch_not_a_number(v):42    """checks whether a tortured variable is nan"""43    v = v.data44    if not ((v == v).item()):45        return True46    return False47 48def is_torch_invalid(v):49    """checks whether a torch variable is nan or inf"""50    if is_torch_not_a_number(v):51        return True52    a = v - v53    if is_torch_not_a_number(a):54        return True55    return False56 57def _relu(x): return x.clamp(min=0)58 59class Entropy(nn.Module):60    """Calculates the entropy of logits"""61    def __init__(self):62        super(Entropy, self).__init__()63 64    def forward(self, x):65        b = F.softmax(x, dim=0) * F.log_softmax(x, dim=0)66        b = -1.0 * b.sum()67        return b68 69class GrammarNetwork(nn.Module):70    """Neural network that outputs a grammar"""71    def __init__(self, inputDimensionality, grammar):72        super(GrammarNetwork, self).__init__()73        self.logProductions = nn.Linear(inputDimensionality, len(grammar)+1)74        self.grammar = grammar75        76    def forward(self, x):77        """Takes as input inputDimensionality-dimensional vector and returns Grammar78        Tensor-valued probabilities"""79        logProductions = self.logProductions(x)80        return Grammar(logProductions[-1].view(1), #logVariable81                       [(logProductions[k].view(1), t, program)82                        for k, (_, t, program) in enumerate(self.grammar.productions)],83                       continuationType=self.grammar.continuationType)84 85    def batchedLogLikelihoods(self, xs, summaries):86        """Takes as input BxinputDimensionality vector & B likelihood summaries;87        returns B-dimensional vector containing log likelihood of each summary"""88        use_cuda = xs.device.type == 'cuda'89 90        B = xs.size(0)91        assert len(summaries) == B92        logProductions = self.logProductions(xs)93 94        # uses[b][p] is # uses of primitive p by summary b95        uses = np.zeros((B,len(self.grammar) + 1))96        for b,summary in enumerate(summaries):97            for p, production in enumerate(self.grammar.primitives):98                uses[b,p] = summary.uses.get(production, 0.)99            uses[b,len(self.grammar)] = summary.uses.get(Index(0), 0)100 101        numerator = (logProductions * maybe_cuda(torch.from_numpy(uses).float(),use_cuda)).sum(1)102        numerator += maybe_cuda(torch.tensor([summary.constant for summary in summaries ]).float(), use_cuda)103 104        alternativeSet = {normalizer105                          for s in summaries106                          for normalizer in s.normalizers }107        alternativeSet = list(alternativeSet)108 109        mask = np.zeros((len(alternativeSet), len(self.grammar) + 1))110        for tau in range(len(alternativeSet)):111            for p, production in enumerate(self.grammar.primitives):112                mask[tau,p] = 0. if production in alternativeSet[tau] else NEGATIVEINFINITY113            mask[tau,len(self.grammar)] = 0. if Index(0) in alternativeSet[tau] else NEGATIVEINFINITY114        mask = maybe_cuda(torch.tensor(mask).float(), use_cuda)115 116        # mask: Rx|G|117        # logProductions: Bx|G|118        # Want: mask + logProductions : BxRx|G| = z119        z = mask.repeat(B,1,1) + logProductions.repeat(len(alternativeSet),1,1).transpose(1,0)120        # z: BxR121        z = torch.logsumexp(z, 2) # pytorch 1.0 dependency122 123        # Calculate how many times each normalizer was used124        N = np.zeros((B, len(alternativeSet)))125        for b, summary in enumerate(summaries):126            for tau, alternatives in enumerate(alternativeSet):127                N[b, tau] = summary.normalizers.get(alternatives,0.)128 129        denominator = (maybe_cuda(torch.tensor(N).float(),use_cuda) * z).sum(1)130        return numerator - denominator131 132        133 134class ContextualGrammarNetwork_LowRank(nn.Module):135    def __init__(self, inputDimensionality, grammar, R=16):136        """Low-rank approximation to bigram model. Parameters is linear in number of primitives.137        R: maximum rank"""138        139        super(ContextualGrammarNetwork_LowRank, self).__init__()140 141        self.grammar = grammar142 143        self.R = R # embedding size144 145        # library now just contains a list of indicies which go with each primitive146        self.grammar = grammar147        self.library = {}148        self.n_grammars = 0149        for prim in grammar.primitives:150            numberOfArguments = len(prim.infer().functionArguments())151            idx_list = list(range(self.n_grammars, self.n_grammars+numberOfArguments))152            self.library[prim] = idx_list153            self.n_grammars += numberOfArguments154        155        # We had an extra grammar for when there is no parent and for when the parent is a variable156        self.n_grammars += 2157        self.transitionMatrix = LowRank(inputDimensionality, self.n_grammars, len(grammar) + 1, R)158        159    def grammarFromVector(self, logProductions):160        return Grammar(logProductions[-1].view(1),161                       [(logProductions[k].view(1), t, program)162                        for k, (_, t, program) in enumerate(self.grammar.productions)],163                       continuationType=self.grammar.continuationType)164 165    def forward(self, x):166        assert len(x.size()) == 1, "contextual grammar doesn't currently support batching"167 168        transitionMatrix = self.transitionMatrix(x)169        170        return ContextualGrammar(self.grammarFromVector(transitionMatrix[-1]), self.grammarFromVector(transitionMatrix[-2]),171                {prim: [self.grammarFromVector(transitionMatrix[j]) for j in js]172                 for prim, js in self.library.items()} )173        174    def vectorizedLogLikelihoods(self, x, summaries):175        B = len(summaries)176        G = len(self.grammar) + 1177 178        # Which column of the transition matrix corresponds to which primitive179        primitiveColumn = {p: c180                           for c, (_1,_2,p) in enumerate(self.grammar.productions) }181        primitiveColumn[Index(0)] = G - 1182        # Which row of the transition matrix corresponds to which context183        contextRow = {(parent, index): r184                      for parent, indices in self.library.items()185                      for index, r in enumerate(indices) }186        contextRow[(None,None)] = self.n_grammars - 1187        contextRow[(Index(0),None)] = self.n_grammars - 2188 189        transitionMatrix = self.transitionMatrix(x)190 191        # uses[b][g][p] is # uses of primitive p by summary b for parent g192        uses = np.zeros((B,self.n_grammars,len(self.grammar)+1))193        for b,summary in enumerate(summaries):194            for e, ss in summary.library.items():195                for g,s in zip(self.library[e], ss):196                    assert g < self.n_grammars - 2197                    for p, production in enumerate(self.grammar.primitives):198                        uses[b,g,p] = s.uses.get(production, 0.)199                    uses[b,g,len(self.grammar)] = s.uses.get(Index(0), 0)200                    201            # noParent: this is the last network output202            for p, production in enumerate(self.grammar.primitives):            203                uses[b, self.n_grammars - 1, p] = summary.noParent.uses.get(production, 0.)204            uses[b, self.n_grammars - 1, G - 1] = summary.noParent.uses.get(Index(0), 0.)205 206            # variableParent: this is the penultimate network output207            for p, production in enumerate(self.grammar.primitives):            208                uses[b, self.n_grammars - 2, p] = summary.variableParent.uses.get(production, 0.)209            uses[b, self.n_grammars - 2, G - 1] = summary.variableParent.uses.get(Index(0), 0.)210 211        uses = maybe_cuda(torch.tensor(uses).float(),use_cuda)212        numerator = uses.view(B, -1) @ transitionMatrix.view(-1)213        214        constant = np.zeros(B)215        for b,summary in enumerate(summaries):216            constant[b] += summary.noParent.constant + summary.variableParent.constant217            for ss in summary.library.values():218                for s in ss:219                    constant[b] += s.constant220            221        numerator = numerator + maybe_cuda(torch.tensor(constant).float(),use_cuda)222 223        # Calculate the god-awful denominator224        # Map from (parent, index, {set-of-alternatives}) to [occurrences-in-summary-zero, occurrences-in-summary-one, ...]225        alternativeSet = {}226        for b,summary in enumerate(summaries):227            for normalizer, frequency in summary.noParent.normalizers.items():228                k = (None,None,normalizer)229                alternativeSet[k] = alternativeSet.get(k, np.zeros(B))230                alternativeSet[k][b] += frequency231            for normalizer, frequency in summary.variableParent.normalizers.items():232                k = (Index(0),None,normalizer)233                alternativeSet[k] = alternativeSet.get(k, np.zeros(B))234                alternativeSet[k][b] += frequency235            for parent, ss in summary.library.items():236                for argumentIndex, s in enumerate(ss):237                    for normalizer, frequency in s.normalizers.items():238                        k = (parent, argumentIndex, normalizer)239                        alternativeSet[k] = alternativeSet.get(k, zeros(B))240                        alternativeSet[k][b] += frequency241 242        # Calculate each distinct normalizing constant243        alternativeNormalizer = {}244        for parent, index, alternatives in alternativeSet:245            r = transitionMatrix[contextRow[(parent, index)]]246            entries = r[ [primitiveColumn[alternative] for alternative in alternatives ]]247            alternativeNormalizer[(parent, index, alternatives)] = torch.logsumexp(entries, dim=0)248 249        # Concatenate the normalizers into a vector250        normalizerKeys = list(alternativeSet.keys())251        normalizerVector = torch.cat([ alternativeNormalizer[k] for k in normalizerKeys])252 253        assert False, "This function is still in progress."254        255 256    def batchedLogLikelihoods(self, xs, summaries):257        """Takes as input BxinputDimensionality vector & B likelihood summaries;258        returns B-dimensional vector containing log likelihood of each summary"""259        use_cuda = xs.device.type == 'cuda'260        261        B = xs.shape[0]262        G = len(self.grammar) + 1263        assert len(summaries) == B264 265        # logProductions: Bx n_grammars x G266        logProductions = self.transitionMatrix(xs)267        # uses[b][g][p] is # uses of primitive p by summary b for parent g268        uses = np.zeros((B,self.n_grammars,len(self.grammar)+1))269        for b,summary in enumerate(summaries):270            for e, ss in summary.library.items():271                for g,s in zip(self.library[e], ss):272                    assert g < self.n_grammars - 2273                    for p, production in enumerate(self.grammar.primitives):274                        uses[b,g,p] = s.uses.get(production, 0.)275                    uses[b,g,len(self.grammar)] = s.uses.get(Index(0), 0)276                    277            # noParent: this is the last network output278            for p, production in enumerate(self.grammar.primitives):            279                uses[b, self.n_grammars - 1, p] = summary.noParent.uses.get(production, 0.)280            uses[b, self.n_grammars - 1, G - 1] = summary.noParent.uses.get(Index(0), 0.)281 282            # variableParent: this is the penultimate network output283            for p, production in enumerate(self.grammar.primitives):            284                uses[b, self.n_grammars - 2, p] = summary.variableParent.uses.get(production, 0.)285            uses[b, self.n_grammars - 2, G - 1] = summary.variableParent.uses.get(Index(0), 0.)286            287        numerator = (logProductions*maybe_cuda(torch.tensor(uses).float(),use_cuda)).view(B,-1).sum(1)288 289        constant = np.zeros(B)290        for b,summary in enumerate(summaries):291            constant[b] += summary.noParent.constant + summary.variableParent.constant292            for ss in summary.library.values():293                for s in ss:294                    constant[b] += s.constant295            296        numerator += maybe_cuda(torch.tensor(constant).float(),use_cuda)297        298        if True:299 300            # Calculate the god-awful denominator301            alternativeSet = set()302            for summary in summaries:303                for normalizer in summary.noParent.normalizers: alternativeSet.add(normalizer)304                for normalizer in summary.variableParent.normalizers: alternativeSet.add(normalizer)305                for ss in summary.library.values():306                    for s in ss:307                        for normalizer in s.normalizers: alternativeSet.add(normalizer)308            alternativeSet = list(alternativeSet)309 310            mask = np.zeros((len(alternativeSet), G))311            for tau in range(len(alternativeSet)):312                for p, production in enumerate(self.grammar.primitives):313                    mask[tau,p] = 0. if production in alternativeSet[tau] else NEGATIVEINFINITY314                mask[tau, G - 1] = 0. if Index(0) in alternativeSet[tau] else NEGATIVEINFINITY315            mask = maybe_cuda(torch.tensor(mask).float(), use_cuda)316 317            z = mask.repeat(self.n_grammars,1,1).repeat(B,1,1,1) + \318                logProductions.repeat(len(alternativeSet),1,1,1).transpose(0,1).transpose(1,2)319            z = torch.logsumexp(z, 3) # pytorch 1.0 dependency320 321            N = np.zeros((B, self.n_grammars, len(alternativeSet)))322            for b, summary in enumerate(summaries):323                for e, ss in summary.library.items():324                    for g,s in zip(self.library[e], ss):325                        assert g < self.n_grammars - 2326                        for r, alternatives in enumerate(alternativeSet):                327                            N[b,g,r] = s.normalizers.get(alternatives, 0.)328                # noParent: this is the last network output329                for r, alternatives in enumerate(alternativeSet):330                    N[b,self.n_grammars - 1,r] = summary.noParent.normalizers.get(alternatives, 0.)331                # variableParent: this is the penultimate network output332                for r, alternatives in enumerate(alternativeSet):333                    N[b,self.n_grammars - 2,r] = summary.variableParent.normalizers.get(alternatives, 0.)334            N = maybe_cuda(torch.tensor(N).float(), use_cuda)335            denominator = (N*z).sum(1).sum(1)336        else:337            gs = [ self(xs[b]) for b in range(B) ]338            denominator = torch.cat([ summary.denominator(g) for summary,g in zip(summaries, gs) ])339            340            341 342        343        344        ll = numerator - denominator 345 346        if False: # verifying that batching works correctly347            gs = [ self(xs[b]) for b in range(B) ]348            _l = torch.cat([ summary.logLikelihood(g) for summary,g in zip(summaries, gs) ])349            assert torch.all((ll - _l).abs() < 0.0001)350        return ll351    352class ContextualGrammarNetwork_Mask(nn.Module):353    def __init__(self, inputDimensionality, grammar):354        """Bigram model, but where the bigram transitions are unconditional.355        Individual primitive probabilities are still conditional (predicted by neural network)356        """357        358        super(ContextualGrammarNetwork_Mask, self).__init__()359 360        self.grammar = grammar361 362        # library now just contains a list of indicies which go with each primitive363        self.grammar = grammar364        self.library = {}365        self.n_grammars = 0366        for prim in grammar.primitives:367            numberOfArguments = len(prim.infer().functionArguments())368            idx_list = list(range(self.n_grammars, self.n_grammars+numberOfArguments))369            self.library[prim] = idx_list370            self.n_grammars += numberOfArguments371        372        # We had an extra grammar for when there is no parent and for when the parent is a variable373        self.n_grammars += 2374        self._transitionMatrix = nn.Parameter(nn.init.xavier_uniform(torch.Tensor(self.n_grammars, len(grammar) + 1)))375        self._logProductions = nn.Linear(inputDimensionality, len(grammar)+1)376 377    def transitionMatrix(self, x):378        if len(x.shape) == 1: # not batched379            return self._logProductions(x) + self._transitionMatrix # will broadcast380        elif len(x.shape) == 2: # batched381            return self._logProductions(x).unsqueeze(1).repeat(1,self.n_grammars,1) + \382                self._transitionMatrix.unsqueeze(0).repeat(x.size(0),1,1)383        else:384            assert False, "unknown shape for transition matrix input"385        386    def grammarFromVector(self, logProductions):387        return Grammar(logProductions[-1].view(1),388                       [(logProductions[k].view(1), t, program)389                        for k, (_, t, program) in enumerate(self.grammar.productions)],390                       continuationType=self.grammar.continuationType)391 392    def forward(self, x):393        assert len(x.size()) == 1, "contextual grammar doesn't currently support batching"394 395        transitionMatrix = self.transitionMatrix(x)396        397        return ContextualGrammar(self.grammarFromVector(transitionMatrix[-1]), self.grammarFromVector(transitionMatrix[-2]),398                {prim: [self.grammarFromVector(transitionMatrix[j]) for j in js]399                 for prim, js in self.library.items()} )400        401    def batchedLogLikelihoods(self, xs, summaries):402        """Takes as input BxinputDimensionality vector & B likelihood summaries;403        returns B-dimensional vector containing log likelihood of each summary"""404        use_cuda = xs.device.type == 'cuda'405        406        B = xs.shape[0]407        G = len(self.grammar) + 1408        assert len(summaries) == B409 410        # logProductions: Bx n_grammars x G411        logProductions = self.transitionMatrix(xs)412        # uses[b][g][p] is # uses of primitive p by summary b for parent g413        uses = np.zeros((B,self.n_grammars,len(self.grammar)+1))414        for b,summary in enumerate(summaries):415            for e, ss in summary.library.items():416                for g,s in zip(self.library[e], ss):417                    assert g < self.n_grammars - 2418                    for p, production in enumerate(self.grammar.primitives):419                        uses[b,g,p] = s.uses.get(production, 0.)420                    uses[b,g,len(self.grammar)] = s.uses.get(Index(0), 0)421                    422            # noParent: this is the last network output423            for p, production in enumerate(self.grammar.primitives):            424                uses[b, self.n_grammars - 1, p] = summary.noParent.uses.get(production, 0.)425            uses[b, self.n_grammars - 1, G - 1] = summary.noParent.uses.get(Index(0), 0.)426 427            # variableParent: this is the penultimate network output428            for p, production in enumerate(self.grammar.primitives):            429                uses[b, self.n_grammars - 2, p] = summary.variableParent.uses.get(production, 0.)430            uses[b, self.n_grammars - 2, G - 1] = summary.variableParent.uses.get(Index(0), 0.)431            432        numerator = (logProductions*maybe_cuda(torch.tensor(uses).float(),use_cuda)).view(B,-1).sum(1)433 434        constant = np.zeros(B)435        for b,summary in enumerate(summaries):436            constant[b] += summary.noParent.constant + summary.variableParent.constant437            for ss in summary.library.values():438                for s in ss:439                    constant[b] += s.constant440            441        numerator += maybe_cuda(torch.tensor(constant).float(),use_cuda)442        443        if True:444 445            # Calculate the god-awful denominator446            alternativeSet = set()447            for summary in summaries:448                for normalizer in summary.noParent.normalizers: alternativeSet.add(normalizer)449                for normalizer in summary.variableParent.normalizers: alternativeSet.add(normalizer)450                for ss in summary.library.values():451                    for s in ss:452                        for normalizer in s.normalizers: alternativeSet.add(normalizer)453            alternativeSet = list(alternativeSet)454 455            mask = np.zeros((len(alternativeSet), G))456            for tau in range(len(alternativeSet)):457                for p, production in enumerate(self.grammar.primitives):458                    mask[tau,p] = 0. if production in alternativeSet[tau] else NEGATIVEINFINITY459                mask[tau, G - 1] = 0. if Index(0) in alternativeSet[tau] else NEGATIVEINFINITY460            mask = maybe_cuda(torch.tensor(mask).float(), use_cuda)461 462            z = mask.repeat(self.n_grammars,1,1).repeat(B,1,1,1) + \463                logProductions.repeat(len(alternativeSet),1,1,1).transpose(0,1).transpose(1,2)464            z = torch.logsumexp(z, 3) # pytorch 1.0 dependency465 466            N = np.zeros((B, self.n_grammars, len(alternativeSet)))467            for b, summary in enumerate(summaries):468                for e, ss in summary.library.items():469                    for g,s in zip(self.library[e], ss):470                        assert g < self.n_grammars - 2471                        for r, alternatives in enumerate(alternativeSet):                472                            N[b,g,r] = s.normalizers.get(alternatives, 0.)473                # noParent: this is the last network output474                for r, alternatives in enumerate(alternativeSet):475                    N[b,self.n_grammars - 1,r] = summary.noParent.normalizers.get(alternatives, 0.)476                # variableParent: this is the penultimate network output477                for r, alternatives in enumerate(alternativeSet):478                    N[b,self.n_grammars - 2,r] = summary.variableParent.normalizers.get(alternatives, 0.)479            N = maybe_cuda(torch.tensor(N).float(), use_cuda)480            denominator = (N*z).sum(1).sum(1)481        else:482            gs = [ self(xs[b]) for b in range(B) ]483            denominator = torch.cat([ summary.denominator(g) for summary,g in zip(summaries, gs) ])484            485            486 487        488        489        ll = numerator - denominator490 491        if False: # verifying that batching works correctly492            gs = [ self(xs[b]) for b in range(B) ]493            _l = torch.cat([ summary.logLikelihood(g) for summary,g in zip(summaries, gs) ])494            assert torch.all((ll - _l).abs() < 0.0001)495        return ll496        497                498 499class ContextualGrammarNetwork(nn.Module):500    """Like GrammarNetwork but ~contextual~"""501    def __init__(self, inputDimensionality, grammar):502        super(ContextualGrammarNetwork, self).__init__()503        504        # library now just contains a list of indicies which go with each primitive505        self.grammar = grammar506        self.library = {}507        self.n_grammars = 0508        for prim in grammar.primitives:509            numberOfArguments = len(prim.infer().functionArguments())510            idx_list = list(range(self.n_grammars, self.n_grammars+numberOfArguments))511            self.library[prim] = idx_list512            self.n_grammars += numberOfArguments513        514        # We had an extra grammar for when there is no parent and for when the parent is a variable515        self.n_grammars += 2516        self.network = nn.Linear(inputDimensionality, (self.n_grammars)*(len(grammar) + 1))517 518 519    def grammarFromVector(self, logProductions):520        return Grammar(logProductions[-1].view(1),521                       [(logProductions[k].view(1), t, program)522                        for k, (_, t, program) in enumerate(self.grammar.productions)],523                       continuationType=self.grammar.continuationType)524 525    def forward(self, x):526        assert len(x.size()) == 1, "contextual grammar doesn't currently support batching"527 528        allVars = self.network(x).view(self.n_grammars, -1)529        return ContextualGrammar(self.grammarFromVector(allVars[-1]), self.grammarFromVector(allVars[-2]),530                {prim: [self.grammarFromVector(allVars[j]) for j in js]531                 for prim, js in self.library.items()} )532 533    def batchedLogLikelihoods(self, xs, summaries):534        use_cuda = xs.device.type == 'cuda'535        """Takes as input BxinputDimensionality vector & B likelihood summaries;536        returns B-dimensional vector containing log likelihood of each summary"""537 538        B = xs.shape[0]539        G = len(self.grammar) + 1540        assert len(summaries) == B541 542        # logProductions: Bx n_grammars x G543        logProductions = self.network(xs).view(B, self.n_grammars, G)544        # uses[b][g][p] is # uses of primitive p by summary b for parent g545        uses = np.zeros((B,self.n_grammars,len(self.grammar)+1))546        for b,summary in enumerate(summaries):547            for e, ss in summary.library.items():548                for g,s in zip(self.library[e], ss):549                    assert g < self.n_grammars - 2550                    for p, production in enumerate(self.grammar.primitives):551                        uses[b,g,p] = s.uses.get(production, 0.)552                    uses[b,g,len(self.grammar)] = s.uses.get(Index(0), 0)553                    554            # noParent: this is the last network output555            for p, production in enumerate(self.grammar.primitives):            556                uses[b, self.n_grammars - 1, p] = summary.noParent.uses.get(production, 0.)557            uses[b, self.n_grammars - 1, G - 1] = summary.noParent.uses.get(Index(0), 0.)558 559            # variableParent: this is the penultimate network output560            for p, production in enumerate(self.grammar.primitives):            561                uses[b, self.n_grammars - 2, p] = summary.variableParent.uses.get(production, 0.)562            uses[b, self.n_grammars - 2, G - 1] = summary.variableParent.uses.get(Index(0), 0.)563            564        numerator = (logProductions*maybe_cuda(torch.tensor(uses).float(),use_cuda)).view(B,-1).sum(1)565 566        constant = np.zeros(B)567        for b,summary in enumerate(summaries):568            constant[b] += summary.noParent.constant + summary.variableParent.constant569            for ss in summary.library.values():570                for s in ss:571                    constant[b] += s.constant572            573        numerator += maybe_cuda(torch.tensor(constant).float(),use_cuda)574 575        # Calculate the god-awful denominator576        alternativeSet = set()577        for summary in summaries:578            for normalizer in summary.noParent.normalizers: alternativeSet.add(normalizer)579            for normalizer in summary.variableParent.normalizers: alternativeSet.add(normalizer)580            for ss in summary.library.values():581                for s in ss:582                    for normalizer in s.normalizers: alternativeSet.add(normalizer)583        alternativeSet = list(alternativeSet)584 585        mask = np.zeros((len(alternativeSet), G))586        for tau in range(len(alternativeSet)):587            for p, production in enumerate(self.grammar.primitives):588                mask[tau,p] = 0. if production in alternativeSet[tau] else NEGATIVEINFINITY589            mask[tau, G - 1] = 0. if Index(0) in alternativeSet[tau] else NEGATIVEINFINITY590        mask = maybe_cuda(torch.tensor(mask).float(), use_cuda)591 592        z = mask.repeat(self.n_grammars,1,1).repeat(B,1,1,1) + \593            logProductions.repeat(len(alternativeSet),1,1,1).transpose(0,1).transpose(1,2)594        z = torch.logsumexp(z, 3) # pytorch 1.0 dependency595 596        N = np.zeros((B, self.n_grammars, len(alternativeSet)))597        for b, summary in enumerate(summaries):598            for e, ss in summary.library.items():599                for g,s in zip(self.library[e], ss):600                    assert g < self.n_grammars - 2601                    for r, alternatives in enumerate(alternativeSet):                602                        N[b,g,r] = s.normalizers.get(alternatives, 0.)603            # noParent: this is the last network output604            for r, alternatives in enumerate(alternativeSet):605                N[b,self.n_grammars - 1,r] = summary.noParent.normalizers.get(alternatives, 0.)606            # variableParent: this is the penultimate network output607            for r, alternatives in enumerate(alternativeSet):608                N[b,self.n_grammars - 2,r] = summary.variableParent.normalizers.get(alternatives, 0.)609        N = maybe_cuda(torch.tensor(N).float(), use_cuda)610        611 612        613        denominator = (N*z).sum(1).sum(1)614        ll = numerator - denominator615 616        if False: # verifying that batching works correctly617            gs = [ self(xs[b]) for b in range(B) ]618            _l = torch.cat([ summary.logLikelihood(g) for summary,g in zip(summaries, gs) ])619            assert torch.all((ll - _l).abs() < 0.0001)620 621        return ll622        623 624class RecognitionModel(nn.Module):625    def __init__(self,featureExtractor,grammar,hidden=[64],activation="tanh",626                 rank=None,contextual=False,mask=False,627                 cuda=False,628                 previousRecognitionModel=None,629                 id=0):630        super(RecognitionModel, self).__init__()631        self.id = id632        self.trained=False633        self.use_cuda = cuda634 635        self.featureExtractor = featureExtractor636        # Sanity check - make sure that all of the parameters of the637        # feature extractor were added to our parameters as well638        if hasattr(featureExtractor, 'parameters'):639            for parameter in featureExtractor.parameters():640                assert any(myParameter is parameter for myParameter in self.parameters())641 642        # Build the multilayer perceptron that is sandwiched between the feature extractor and the grammar643        if activation == "sigmoid":644            activation = nn.Sigmoid645        elif activation == "relu":646            activation = nn.ReLU647        elif activation == "tanh":648            activation = nn.Tanh649        else:650            raise Exception('Unknown activation function ' + str(activation))651        self._MLP = nn.Sequential(*[ layer652                                     for j in range(len(hidden))653                                     for layer in [654                                             nn.Linear(([featureExtractor.outputDimensionality] + hidden)[j],655                                                       hidden[j]),656                                             activation()]])657 658        self.entropy = Entropy()659 660        if len(hidden) > 0:661            self.outputDimensionality = self._MLP[-2].out_features662            assert self.outputDimensionality == hidden[-1]663        else:664            self.outputDimensionality = self.featureExtractor.outputDimensionality665 666        self.contextual = contextual667        if self.contextual:668            if mask:669                self.grammarBuilder = ContextualGrammarNetwork_Mask(self.outputDimensionality, grammar)670            else:671                self.grammarBuilder = ContextualGrammarNetwork_LowRank(self.outputDimensionality, grammar, rank)672        else:673            self.grammarBuilder = GrammarNetwork(self.outputDimensionality, grammar)674        675        self.grammar = ContextualGrammar.fromGrammar(grammar) if contextual else grammar676        self.generativeModel = grammar677        678        self._auxiliaryPrediction = nn.Linear(self.featureExtractor.outputDimensionality, 679                                              len(self.grammar.primitives))680        self._auxiliaryLoss = nn.BCEWithLogitsLoss()681 682        if cuda: self.cuda()683 684        if previousRecognitionModel:685            self._MLP.load_state_dict(previousRecognitionModel._MLP.state_dict())686            self.featureExtractor.load_state_dict(previousRecognitionModel.featureExtractor.state_dict())687            688    def auxiliaryLoss(self, frontier, features):689        # Compute a vector of uses690        ls = frontier.bestPosterior.program691        def uses(summary):692            if hasattr(summary, 'uses'): 693                return torch.tensor([ float(int(p in summary.uses))694                                      for p in self.generativeModel.primitives ])695            assert hasattr(summary, 'noParent')696            u = uses(summary.noParent) + uses(summary.variableParent)697            for ss in summary.library.values():698                for s in ss:699                    u += uses(s)700            return u701        u = uses(ls)702        u[u > 1.] = 1.703        if self.use_cuda: u = u.cuda()704        al = self._auxiliaryLoss(self._auxiliaryPrediction(features), u)705        return al706            707    def taskEmbeddings(self, tasks):708        return {task: self.featureExtractor.featuresOfTask(task).data.cpu().numpy()709                for task in tasks}710 711    def forward(self, features):712        """returns either a Grammar or a ContextualGrammar713        Takes as input the output of featureExtractor.featuresOfTask"""714        features = self._MLP(features)715        return self.grammarBuilder(features)716 717    def auxiliaryPrimitiveEmbeddings(self):718        """Returns the actual outputDimensionality weight vectors for each of the primitives."""719        auxiliaryWeights = self._auxiliaryPrediction.weight.data.cpu().numpy()720        primitivesDict =  {self.grammar.primitives[i] : auxiliaryWeights[i, :] for i in range(len(self.grammar.primitives))}721        return primitivesDict722 723    def grammarOfTask(self, task):724        features = self.featureExtractor.featuresOfTask(task)725        if features is None: return None726        return self(features)727 728    def grammarLogProductionsOfTask(self, task):729        """Returns the grammar logits from non-contextual models."""730 731        features = self.featureExtractor.featuresOfTask(task)732        if features is None: return None733 734        if hasattr(self, 'hiddenLayers'):735            # Backward compatability with old checkpoints.736            for layer in self.hiddenLayers:737                features = self.activation(layer(features))738            # return features739            return self.noParent[1](features)740        else:741            features = self._MLP(features)742 743        if self.contextual:744            if hasattr(self.grammarBuilder, 'variableParent'):745                return self.grammarBuilder.variableParent.logProductions(features)746            elif hasattr(self.grammarBuilder, 'network'):747                return self.grammarBuilder.network(features).view(-1)748            elif hasattr(self.grammarBuilder, 'transitionMatrix'):749                return self.grammarBuilder.transitionMatrix(features).view(-1)750            else:751                assert False752        else:753            return self.grammarBuilder.logProductions(features)754 755    def grammarFeatureLogProductionsOfTask(self, task):756        return torch.tensor(self.grammarOfTask(task).untorch().featureVector())757 758    def grammarLogProductionDistanceToTask(self, task, tasks):759        """Returns the cosine similarity of all other tasks to a given task."""760        taskLogits = self.grammarLogProductionsOfTask(task).unsqueeze(0) # Change to [1, D]761        assert taskLogits is not None, 'Grammar log productions are not defined for this task.'762        otherTasks = [t for t in tasks if t is not task] # [nTasks -1 , D]763 764        # Build matrix of all other tasks.765        otherLogits = torch.stack([self.grammarLogProductionsOfTask(t) for t in otherTasks])766        cos = nn.CosineSimilarity(dim=1, eps=1e-6)767        cosMatrix = cos(taskLogits, otherLogits)768        return cosMatrix.data.cpu().numpy()769 770    def grammarEntropyOfTask(self, task):771        """Returns the entropy of the grammar distribution from non-contextual models for a task."""772        grammarLogProductionsOfTask = self.grammarLogProductionsOfTask(task)773 774        if grammarLogProductionsOfTask is None: return None775 776        if hasattr(self, 'entropy'):777            return self.entropy(grammarLogProductionsOfTask)778        else:779            e = Entropy()780            return e(grammarLogProductionsOfTask)781 782    def taskAuxiliaryLossLayer(self, tasks):783        return {task: self._auxiliaryPrediction(self.featureExtractor.featuresOfTask(task)).view(-1).data.cpu().numpy()784                for task in tasks}785                786    def taskGrammarFeatureLogProductions(self, tasks):787        return {task: self.grammarFeatureLogProductionsOfTask(task).data.cpu().numpy()788                for task in tasks}789 790    def taskGrammarLogProductions(self, tasks):791        return {task: self.grammarLogProductionsOfTask(task).data.cpu().numpy()792                for task in tasks}793 794    def taskGrammarStartProductions(self, tasks):795        return {task: np.array([l for l,_1,_2 in g.productions ])796                for task in tasks797                for g in [self.grammarOfTask(task).untorch().noParent] }798 799    def taskHiddenStates(self, tasks):800        return {task: self._MLP(self.featureExtractor.featuresOfTask(task)).view(-1).data.cpu().numpy()801                for task in tasks}802 803    def taskGrammarEntropies(self, tasks):804        return {task: self.grammarEntropyOfTask(task).data.cpu().numpy()805                for task in tasks}806 807    def frontierKL(self, frontier, auxiliary=False, vectorized=True):808        features = self.featureExtractor.featuresOfTask(frontier.task)809        if features is None:810            return None, None811        # Monte Carlo estimate: draw a sample from the frontier812        entry = frontier.sample()813 814        al = self.auxiliaryLoss(frontier, features if auxiliary else features.detach())815 816        if not vectorized:817            g = self(features)818            return - entry.program.logLikelihood(g), al819        else:820            features = self._MLP(features).unsqueeze(0)821            822            ll = self.grammarBuilder.batchedLogLikelihoods(features, [entry.program]).view(-1)823            return -ll, al824            825 826    def frontierBiasOptimal(self, frontier, auxiliary=False, vectorized=True):827        if not vectorized:828            features = self.featureExtractor.featuresOfTask(frontier.task)829            if features is None: return None, None830            al = self.auxiliaryLoss(frontier, features if auxiliary else features.detach())831            g = self(features)832            summaries = [entry.program for entry in frontier]833            likelihoods = torch.cat([entry.program.logLikelihood(g) + entry.logLikelihood834                                     for entry in frontier ])835            best = likelihoods.max()836            return -best, al837            838        batchSize = len(frontier.entries)839        features = self.featureExtractor.featuresOfTask(frontier.task)840        if features is None: return None, None841        al = self.auxiliaryLoss(frontier, features if auxiliary else features.detach())842        features = self._MLP(features)843        features = features.expand(batchSize, features.size(-1))  # TODO844        lls = self.grammarBuilder.batchedLogLikelihoods(features, [entry.program for entry in frontier])845        actual_ll = torch.Tensor([ entry.logLikelihood for entry in frontier])846        lls = lls + (actual_ll.cuda() if self.use_cuda else actual_ll)847        ml = -lls.max() #Beware that inputs to max change output type848        return ml, al849 850    def replaceProgramsWithLikelihoodSummaries(self, frontier):851        return Frontier(852            [FrontierEntry(853                program=self.grammar.closedLikelihoodSummary(frontier.task.request, e.program),854                logLikelihood=e.logLikelihood,855                logPrior=e.logPrior) for e in frontier],856            task=frontier.task)857 858    def train(self, frontiers, _=None, steps=None, lr=0.001, topK=5, CPUs=1,859              timeout=None, evaluationTimeout=0.001,860              helmholtzFrontiers=[], helmholtzRatio=0., helmholtzBatch=500,861              biasOptimal=None, defaultRequest=None, auxLoss=False, vectorized=True):862        """863        helmholtzRatio: What fraction of the training data should be forward samples from the generative model?864        helmholtzFrontiers: Frontiers from programs enumerated from generative model (optional)865        If helmholtzFrontiers is not provided then we will sample programs during training866        """867        assert (steps is not None) or (timeout is not None), \868            "Cannot train recognition model without either a bound on the number of gradient steps or bound on the training time"869        if steps is None: steps = 9999999870        if biasOptimal is None: biasOptimal = len(helmholtzFrontiers) > 0871        872        requests = [frontier.task.request for frontier in frontiers]873        if len(requests) == 0 and helmholtzRatio > 0 and len(helmholtzFrontiers) == 0:874            assert defaultRequest is not None, "You are trying to random Helmholtz training, but don't have any frontiers. Therefore we would not know the type of the program to sample. Try specifying defaultRequest=..."875            requests = [defaultRequest]876        frontiers = [frontier.topK(topK).normalize()877                     for frontier in frontiers if not frontier.empty]878        if len(frontiers) == 0:879            eprint("You didn't give me any nonempty replay frontiers to learn from. Going to learn from 100% Helmholtz samples")880            helmholtzRatio = 1.881 882        # Should we sample programs or use the enumerated programs?883        randomHelmholtz = len(helmholtzFrontiers) == 0884        885        class HelmholtzEntry:886            def __init__(self, frontier, owner):887                self.request = frontier.task.request888                self.task = None889                self.programs = [e.program for e in frontier]890                self.frontier = Thunk(lambda: owner.replaceProgramsWithLikelihoodSummaries(frontier))891                self.owner = owner892 893            def clear(self): self.task = None894 895            def calculateTask(self):896                assert self.task is None897                p = random.choice(self.programs)898                return self.owner.featureExtractor.taskOfProgram(p, self.request)899 900            def makeFrontier(self):901                assert self.task is not None902                f = Frontier(self.frontier.force().entries,903                             task=self.task)904                return f905        906            907            908 909        # Should we recompute tasks on the fly from Helmholtz?  This910        # should be done if the task is stochastic, or if there are911        # different kinds of inputs on which it could be run. For912        # example, lists and strings need this; towers and graphics do913        # not. There is no harm in recomputed the tasks, it just914        # wastes time.915        if not hasattr(self.featureExtractor, 'recomputeTasks'):916            self.featureExtractor.recomputeTasks = True917        helmholtzFrontiers = [HelmholtzEntry(f, self)918                              for f in helmholtzFrontiers]919        random.shuffle(helmholtzFrontiers)920        921        helmholtzIndex = [0]922        def getHelmholtz():923            if randomHelmholtz:924                if helmholtzIndex[0] >= len(helmholtzFrontiers):925                    updateHelmholtzTasks()926                    helmholtzIndex[0] = 0927                    return getHelmholtz()928                helmholtzIndex[0] += 1929                return helmholtzFrontiers[helmholtzIndex[0] - 1].makeFrontier()930 931            f = helmholtzFrontiers[helmholtzIndex[0]]932            if f.task is None:933                with timing("Evaluated another batch of Helmholtz tasks"):934                    updateHelmholtzTasks()935                return getHelmholtz()936 937            helmholtzIndex[0] += 1938            if helmholtzIndex[0] >= len(helmholtzFrontiers):939                helmholtzIndex[0] = 0940                random.shuffle(helmholtzFrontiers)941                if self.featureExtractor.recomputeTasks:942                    for fp in helmholtzFrontiers:943                        fp.clear()944                    return getHelmholtz() # because we just cleared everything945            assert f.task is not None946            return f.makeFrontier()947            948        def updateHelmholtzTasks():949            updateCPUs = CPUs if hasattr(self.featureExtractor, 'parallelTaskOfProgram') and self.featureExtractor.parallelTaskOfProgram else 1950            if updateCPUs > 1: eprint("Updating Helmholtz tasks with",updateCPUs,"CPUs",951                                      "while using",getThisMemoryUsage(),"memory")952            953            if randomHelmholtz:954                newFrontiers = self.sampleManyHelmholtz(requests, helmholtzBatch, CPUs)955                newEntries = []956                for f in newFrontiers:957                    e = HelmholtzEntry(f,self)958                    e.task = f.task959                    newEntries.append(e)960                helmholtzFrontiers.clear()961                helmholtzFrontiers.extend(newEntries)962                return 963 964            # Save some memory by freeing up the tasks as we go through them965            if self.featureExtractor.recomputeTasks:966                for hi in range(max(0, helmholtzIndex[0] - helmholtzBatch,967                                    min(helmholtzIndex[0], len(helmholtzFrontiers)))):968                    helmholtzFrontiers[hi].clear()969 970            if hasattr(self.featureExtractor, 'tasksOfPrograms'):971                eprint("batching task calculation")972                newTasks = self.featureExtractor.tasksOfPrograms(973                    [random.choice(hf.programs)974                     for hf in helmholtzFrontiers[helmholtzIndex[0]:helmholtzIndex[0] + helmholtzBatch] ],975                    [hf.request976                     for hf in helmholtzFrontiers[helmholtzIndex[0]:helmholtzIndex[0] + helmholtzBatch] ])977            else:978                newTasks = [hf.calculateTask() 979                            for hf in helmholtzFrontiers[helmholtzIndex[0]:helmholtzIndex[0] + helmholtzBatch]]980 981                """982                # catwong: Disabled for ensemble training.983                newTasks = \984                           parallelMap(updateCPUs,985                                       lambda f: f.calculateTask(),986                                       helmholtzFrontiers[helmholtzIndex[0]:helmholtzIndex[0] + helmholtzBatch],987                                       seedRandom=True)988                """989            badIndices = []990            endingIndex = min(helmholtzIndex[0] + helmholtzBatch, len(helmholtzFrontiers))991            for i in range(helmholtzIndex[0], endingIndex):992                helmholtzFrontiers[i].task = newTasks[i - helmholtzIndex[0]]993                if helmholtzFrontiers[i].task is None: badIndices.append(i)994            # Permanently kill anything which failed to give a task995            for i in reversed(badIndices):996                assert helmholtzFrontiers[i].task is None997                del helmholtzFrontiers[i]998 999 1000        # We replace each program in the frontier with its likelihoodSummary1001        # This is because calculating likelihood summaries requires juggling types1002        # And type stuff is expensive!1003        frontiers = [self.replaceProgramsWithLikelihoodSummaries(f).normalize()1004                     for f in frontiers]1005 1006        eprint("(ID=%d): Training a recognition model from %d frontiers, %d%% Helmholtz, feature extractor %s." % (1007            self.id, len(frontiers), int(helmholtzRatio * 100), self.featureExtractor.__class__.__name__))1008        eprint("(ID=%d): Got %d Helmholtz frontiers - random Helmholtz training? : %s"%(1009            self.id, len(helmholtzFrontiers), len(helmholtzFrontiers) == 0))1010        eprint("(ID=%d): Contextual? %s" % (self.id, str(self.contextual)))1011        eprint("(ID=%d): Bias optimal? %s" % (self.id, str(biasOptimal)))1012        eprint(f"(ID={self.id}): Aux loss? {auxLoss} (n.b. we train a 'auxiliary' classifier anyway - this controls if gradients propagate back to the future extractor)")1013 1014        # The number of Helmholtz samples that we generate at once1015        # Should only affect performance and shouldn't affect anything else1016        helmholtzSamples = []1017 1018        optimizer = torch.optim.Adam(self.parameters(), lr=lr, eps=1e-3, amsgrad=True)1019        start = time.time()1020        losses, descriptionLengths, realLosses, dreamLosses, realMDL, dreamMDL = [], [], [], [], [], []1021        classificationLosses = []1022        totalGradientSteps = 01023        epochs = 99999991024        for i in range(1, epochs + 1):1025            if timeout and time.time() - start > timeout:1026                break1027 1028            if totalGradientSteps > steps:1029                break1030 1031            if helmholtzRatio < 1.:1032                permutedFrontiers = list(frontiers)1033                random.shuffle(permutedFrontiers)1034            else:1035                permutedFrontiers = [None]1036 1037            finishedSteps = False1038            for frontier in permutedFrontiers:1039                # Randomly decide whether to sample from the generative model1040                dreaming = random.random() < helmholtzRatio1041                if dreaming: frontier = getHelmholtz()1042                self.zero_grad()1043                loss, classificationLoss = \1044                        self.frontierBiasOptimal(frontier, auxiliary=auxLoss, vectorized=vectorized) if biasOptimal \1045                        else self.frontierKL(frontier, auxiliary=auxLoss, vectorized=vectorized)1046                if loss is None:1047                    if not dreaming:1048                        eprint("ERROR: Could not extract features during experience replay.")1049                        eprint("Task is:",frontier.task)1050                        eprint("Aborting - we need to be able to extract features of every actual task.")1051                        assert False1052                    else:1053                        continue1054                if is_torch_invalid(loss):1055                    eprint("Invalid real-data loss!")1056                else:1057                    (loss + classificationLoss).backward()1058                    classificationLosses.append(classificationLoss.data.item())1059                    optimizer.step()1060                    totalGradientSteps += 11061                    losses.append(loss.data.item())1062                    descriptionLengths.append(min(-e.logPrior for e in frontier))1063                    if dreaming:1064                        dreamLosses.append(losses[-1])1065                        dreamMDL.append(descriptionLengths[-1])1066                    else:1067                        realLosses.append(losses[-1])1068                        realMDL.append(descriptionLengths[-1])1069                    if totalGradientSteps > steps:1070                        break # Stop iterating, then print epoch and loss, then break to finish.1071                        1072            if (i == 1 or i % 10 == 0) and losses:1073                eprint("(ID=%d): " % self.id, "Epoch", i, "Loss", mean(losses))1074                if realLosses and dreamLosses:1075                    eprint("(ID=%d): " % self.id, "\t\t(real loss): ", mean(realLosses), "\t(dream loss):", mean(dreamLosses))1076                eprint("(ID=%d): " % self.id, "\tvs MDL (w/o neural net)", mean(descriptionLengths))1077                if realMDL and dreamMDL:1078                    eprint("\t\t(real MDL): ", mean(realMDL), "\t(dream MDL):", mean(dreamMDL))1079                eprint("(ID=%d): " % self.id, "\t%d cumulative gradient steps. %f steps/sec"%(totalGradientSteps,1080                                                                       totalGradientSteps/(time.time() - start)))1081                eprint("(ID=%d): " % self.id, "\t%d-way auxiliary classification loss"%len(self.grammar.primitives),sum(classificationLosses)/len(classificationLosses))1082                losses, descriptionLengths, realLosses, dreamLosses, realMDL, dreamMDL = [], [], [], [], [], []1083                classificationLosses = []1084                gc.collect()1085        1086        eprint("(ID=%d): " % self.id, " Trained recognition model in",time.time() - start,"seconds")1087        self.trained=True1088        return self1089 1090    def sampleHelmholtz(self, requests, statusUpdate=None, seed=None):1091        if seed is not None:1092            random.seed(seed)1093        request = random.choice(requests)1094 1095        program = self.generativeModel.sample(request, maximumDepth=6, maxAttempts=100)1096        if program is None:1097            return None1098        task = self.featureExtractor.taskOfProgram(program, request)1099 1100        if statusUpdate is not None:1101            flushEverything()1102        if task is None:1103            return None1104 1105        if hasattr(self.featureExtractor, 'lexicon'):1106            if self.featureExtractor.tokenize(task.examples) is None:1107                return None1108        1109        ll = self.generativeModel.logLikelihood(request, program)1110        frontier = Frontier([FrontierEntry(program=program,1111                                           logLikelihood=0., logPrior=ll)],1112                            task=task)1113        return frontier1114 1115    def sampleManyHelmholtz(self, requests, N, CPUs):1116        eprint("Sampling %d programs from the prior on %d CPUs..." % (N, CPUs))1117        flushEverything()1118        frequency = N / 501119        startingSeed = random.random()1120 1121        # Sequentially for ensemble training.1122        samples = [self.sampleHelmholtz(requests,1123                                           statusUpdate='.' if n % frequency == 0 else None,1124                                           seed=startingSeed + n) for n in range(N)]1125 1126        # (cathywong) Disabled for ensemble training. 1127        # samples = parallelMap(1128        #     1,1129        #     lambda n: self.sampleHelmholtz(requests,1130        #                                    statusUpdate='.' if n % frequency == 0 else None,1131        #                                    seed=startingSeed + n),1132        #     range(N))1133        eprint()1134        flushEverything()1135        samples = [z for z in samples if z is not None]1136        eprint()1137        eprint("Got %d/%d valid samples." % (len(samples), N))1138        flushEverything()1139 1140        return samples1141 1142    def enumerateFrontiers(self,1143                           tasks,1144                           enumerationTimeout=None,1145                           testing=False,1146                           solver=None,1147                           CPUs=1,1148                           frontierSize=None,1149                           maximumFrontier=None,1150                           evaluationTimeout=None):1151        with timing("Evaluated recognition model"):1152            grammars = {task: self.grammarOfTask(task)1153                        for task in tasks}1154            #untorch seperately to make sure you filter out None grammars1155            grammars = {task: grammar.untorch() for task, grammar in grammars.items() if grammar is not None}1156 1157        return multicoreEnumeration(grammars, tasks,1158                                    testing=testing,1159                                    solver=solver,1160                                    enumerationTimeout=enumerationTimeout,1161                                    CPUs=CPUs, maximumFrontier=maximumFrontier,1162                                    evaluationTimeout=evaluationTimeout)1163 1164 1165class RecurrentFeatureExtractor(nn.Module):1166    def __init__(self, _=None,1167                 tasks=None,1168                 cuda=False,1169                 # what are the symbols that can occur in the inputs and1170                 # outputs1171                 lexicon=None,1172                 # how many hidden units1173                 H=32,1174                 # Should the recurrent units be bidirectional?1175                 bidirectional=False,1176                 # What should be the timeout for trying to construct Helmholtz tasks?1177                 helmholtzTimeout=0.25,1178                 # What should be the timeout for running a Helmholtz program?1179                 helmholtzEvaluationTimeout=0.01):1180        super(RecurrentFeatureExtractor, self).__init__()1181 1182        assert tasks is not None, "You must provide a list of all of the tasks, both those that have been hit and those that have not been hit. Input examples are sampled from these tasks."1183 1184        # maps from a requesting type to all of the inputs that we ever saw with that request1185        self.requestToInputs = {1186            tp: [list(map(fst, t.examples)) for t in tasks if t.request == tp ]1187            for tp in {t.request for t in tasks}1188        }1189 1190        inputTypes = {t1191                      for task in tasks1192                      for t in task.request.functionArguments()}1193        # maps from a type to all of the inputs that we ever saw having that type1194        self.argumentsWithType = {1195            tp: [ x1196                  for t in tasks1197                  for xs,_ in t.examples1198                  for tpp, x in zip(t.request.functionArguments(), xs)1199                  if tpp == tp]1200            for tp in inputTypes

Showing the first 1,200 of 1529 lines. Download the file for the rest.