Fraser/dream-coder
Program Synthesis Data Generated program synthesis datasets used to train dreamcoder. Currently just supports text & list data.
6730
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