Fraser/dream-coder
Program Synthesis Data Generated program synthesis datasets used to train dreamcoder. Currently just supports text & list data.
6690
1class UnificationFailure(Exception):2 pass3 4 5class Occurs(UnificationFailure):6 pass7 8 9class Type(object):10 def __str__(self): return self.show(True)11 12 def __repr__(self): return str(self)13 14 @staticmethod15 def fromjson(j):16 if "index" in j: return TypeVariable(j["index"])17 if "constructor" in j: return TypeConstructor(j["constructor"],18 [ Type.fromjson(a) for a in j["arguments"] ])19 assert False20 21 22class TypeConstructor(Type):23 def __init__(self, name, arguments):24 self.name = name25 self.arguments = arguments26 self.isPolymorphic = any(a.isPolymorphic for a in arguments)27 28 def makeDummyMonomorphic(self, mapping=None):29 mapping = mapping if mapping is not None else {}30 return TypeConstructor(self.name,31 [ a.makeDummyMonomorphic(mapping) for a in self.arguments ])32 33 def __eq__(self, other):34 return isinstance(other, TypeConstructor) and \35 self.name == other.name and \36 all(x == y for x, y in zip(self.arguments, other.arguments))37 38 def __hash__(self): return hash((self.name,) + tuple(self.arguments))39 40 def __ne__(self, other):41 return not (self == other)42 43 def show(self, isReturn):44 if self.name == ARROW:45 if isReturn:46 return "%s %s %s" % (self.arguments[0].show(47 False), ARROW, self.arguments[1].show(True))48 else:49 return "(%s %s %s)" % (self.arguments[0].show(50 False), ARROW, self.arguments[1].show(True))51 elif self.arguments == []:52 return self.name53 else:54 return "%s(%s)" % (self.name, ", ".join(x.show(True)55 for x in self.arguments))56 57 def json(self):58 return {"constructor": self.name,59 "arguments": [a.json() for a in self.arguments]}60 61 62 def isArrow(self): return self.name == ARROW63 64 def functionArguments(self):65 if self.name == ARROW:66 xs = self.arguments[1].functionArguments()67 return [self.arguments[0]] + xs68 return []69 70 def returns(self):71 if self.name == ARROW:72 return self.arguments[1].returns()73 else:74 return self75 76 def apply(self, context):77 if not self.isPolymorphic:78 return self79 return TypeConstructor(self.name,80 [x.apply(context) for x in self.arguments])81 82 def applyMutable(self, context):83 if not self.isPolymorphic:84 return self85 return TypeConstructor(self.name,86 [x.applyMutable(context) for x in self.arguments])87 88 def occurs(self, v):89 if not self.isPolymorphic:90 return False91 return any(x.occurs(v) for x in self.arguments)92 93 def negateVariables(self):94 return TypeConstructor(self.name,95 [a.negateVariables() for a in self.arguments])96 97 def instantiate(self, context, bindings=None):98 if not self.isPolymorphic:99 return context, self100 if bindings is None:101 bindings = {}102 newArguments = []103 for x in self.arguments:104 (context, x) = x.instantiate(context, bindings)105 newArguments.append(x)106 return (context, TypeConstructor(self.name, newArguments))107 108 def instantiateMutable(self, context, bindings=None):109 if not self.isPolymorphic:110 return self111 if bindings is None:112 bindings = {}113 newArguments = []114 return TypeConstructor(self.name, [x.instantiateMutable(context, bindings)115 for x in self.arguments ])116 117 118 def canonical(self, bindings=None):119 if not self.isPolymorphic:120 return self121 if bindings is None:122 bindings = {}123 return TypeConstructor(self.name,124 [x.canonical(bindings) for x in self.arguments])125 126 127class TypeVariable(Type):128 def __init__(self, j):129 assert isinstance(j, int)130 self.v = j131 self.isPolymorphic = True132 133 def makeDummyMonomorphic(self, mapping=None):134 mapping = mapping if mapping is not None else {}135 if self.v not in mapping:136 mapping[self.v] = TypeConstructor(f"dummy_type_{len(mapping)}", [])137 return mapping[self.v]138 139 140 def __eq__(self, other):141 return isinstance(other, TypeVariable) and self.v == other.v142 143 def __ne__(self, other): return not (self.v == other.v)144 145 def __hash__(self): return self.v146 147 def show(self, _): return "t%d" % self.v148 149 def json(self):150 return {"index": self.v}151 152 def returns(self): return self153 154 def isArrow(self): return False155 156 def functionArguments(self): return []157 158 def apply(self, context):159 for v, t in context.substitution:160 if v == self.v:161 return t.apply(context)162 return self163 164 def applyMutable(self, context):165 s = context.substitution[self.v]166 if s is None: return self167 new = s.applyMutable(context)168 context.substitution[self.v] = new169 return new170 171 def occurs(self, v): return v == self.v172 173 def instantiate(self, context, bindings=None):174 if bindings is None:175 bindings = {}176 if self.v in bindings:177 return (context, bindings[self.v])178 new = TypeVariable(context.nextVariable)179 bindings[self.v] = new180 context = Context(context.nextVariable + 1, context.substitution)181 return (context, new)182 183 def instantiateMutable(self, context, bindings=None):184 if bindings is None: bindings = {}185 if self.v in bindings: return bindings[self.v]186 new = context.makeVariable()187 bindings[self.v] = new188 return new189 190 def canonical(self, bindings=None):191 if bindings is None:192 bindings = {}193 if self.v in bindings:194 return bindings[self.v]195 new = TypeVariable(len(bindings))196 bindings[self.v] = new197 return new198 199 def negateVariables(self):200 return TypeVariable(-1 - self.v)201 202 203class Context(object):204 def __init__(self, nextVariable=0, substitution=[]):205 self.nextVariable = nextVariable206 self.substitution = substitution207 208 def extend(self, j, t):209 return Context(self.nextVariable, [(j, t)] + self.substitution)210 211 def makeVariable(self):212 return (Context(self.nextVariable + 1, self.substitution),213 TypeVariable(self.nextVariable))214 215 def unify(self, t1, t2):216 t1 = t1.apply(self)217 t2 = t2.apply(self)218 if t1 == t2:219 return self220 # t1&t2 are not equal221 if not t1.isPolymorphic and not t2.isPolymorphic:222 raise UnificationFailure(t1, t2)223 224 if isinstance(t1, TypeVariable):225 if t2.occurs(t1.v):226 raise Occurs()227 return self.extend(t1.v, t2)228 if isinstance(t2, TypeVariable):229 if t1.occurs(t2.v):230 raise Occurs()231 return self.extend(t2.v, t1)232 if t1.name != t2.name:233 raise UnificationFailure(t1, t2)234 k = self235 for x, y in zip(t2.arguments, t1.arguments):236 k = k.unify(x, y)237 return k238 239 def __str__(self):240 return "Context(next = %d, {%s})" % (self.nextVariable, ", ".join(241 "t%d ||> %s" % (k, v.apply(self)) for k, v in self.substitution))242 243 def __repr__(self): return str(self)244 245class MutableContext(object):246 def __init__(self):247 self.substitution = []248 249 def extend(self,i,t):250 assert self.substitution[i] is None251 self.substitution[i] = t252 253 def makeVariable(self):254 self.substitution.append(None)255 return TypeVariable(len(self.substitution) - 1)256 257 def unify(self, t1, t2):258 t1 = t1.applyMutable(self)259 t2 = t2.applyMutable(self)260 261 if t1 == t2: return262 263 # t1&t2 are not equal264 if not t1.isPolymorphic and not t2.isPolymorphic:265 raise UnificationFailure(t1, t2)266 267 if isinstance(t1, TypeVariable):268 if t2.occurs(t1.v):269 raise Occurs()270 self.extend(t1.v, t2)271 return 272 if isinstance(t2, TypeVariable):273 if t1.occurs(t2.v):274 raise Occurs()275 self.extend(t2.v, t1)276 return 277 if t1.name != t2.name:278 raise UnificationFailure(t1, t2)279 280 for x, y in zip(t2.arguments, t1.arguments):281 self.unify(x, y)282 283 284Context.EMPTY = Context(0, [])285 286 287def canonicalTypes(ts):288 bindings = {}289 return [t.canonical(bindings) for t in ts]290 291 292def instantiateTypes(context, ts):293 bindings = {}294 newTypes = []295 for t in ts:296 context, t = t.instantiate(context, bindings)297 newTypes.append(t)298 return context, newTypes299 300 301def baseType(n): return TypeConstructor(n, [])302 303 304tint = baseType("int")305treal = baseType("real")306tbool = baseType("bool")307tboolean = tbool # alias308tcharacter = baseType("char")309 310 311def tlist(t): return TypeConstructor("list", [t])312 313 314def tpair(a, b): return TypeConstructor("pair", [a, b])315 316 317def tmaybe(t): return TypeConstructor("maybe", [t])318 319 320tstr = tlist(tcharacter)321t0 = TypeVariable(0)322t1 = TypeVariable(1)323t2 = TypeVariable(2)324 325# regex types326tpregex = baseType("pregex")327 328ARROW = "->"329 330 331def arrow(*arguments):332 if len(arguments) == 1:333 return arguments[0]334 return TypeConstructor(ARROW, [arguments[0], arrow(*arguments[1:])])335 336 337def inferArg(tp, tcaller):338 ctx, tp = tp.instantiate(Context.EMPTY)339 ctx, tcaller = tcaller.instantiate(ctx)340 ctx, targ = ctx.makeVariable()341 ctx = ctx.unify(tcaller, arrow(targ, tp))342 return targ.apply(ctx)343 344 345def guess_type(xs):346 """347 Return a TypeConstructor corresponding to x's python type.348 Raises an exception if the type cannot be guessed.349 """350 if all(isinstance(x, bool) for x in xs):351 return tbool352 elif all(isinstance(x, int) for x in xs):353 return tint354 elif all(isinstance(x, str) for x in xs):355 return tstr356 elif all(isinstance(x, list) for x in xs):357 return tlist(guess_type([y for ys in xs for y in ys]))358 else:359 raise ValueError("cannot guess type from {}".format(xs))360 361 362def guess_arrow_type(examples):363 a = len(examples[0][0])364 input_types = []365 for n in range(a):366 input_types.append(guess_type([xs[n] for xs, _ in examples]))367 output_type = guess_type([y for _, y in examples])368 return arrow(*(input_types + [output_type]))369 370def canUnify(t1, t2):371 k = MutableContext()372 t1 = t1.instantiateMutable(k)373 t2 = t2.instantiateMutable(k)374 try:375 k.unify(t1, t2)376 return True377 except UnificationFailure: return False378 379 