Team Ai
Datasetpublic

Fraser/dream-coder

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

sourceHugging Facemitupdated 4y agoView on Hugging Face
6likes690downloads
type.py379 linesDownload Raw Back to dreamcoder
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