packages feed

futhark-0.28.1: rts/python/server.py

# Start of server.py

import sys
import time
import shlex  # For string splitting


class Server:
    def __init__(self, ctx):
        self._ctx = ctx
        self._vars = {}

    class Failure(BaseException):
        def __init__(self, msg):
            self.msg = msg

    def _get_arg(self, args, i):
        if i < len(args):
            return args[i]
        else:
            raise self.Failure("Insufficient command args")

    def _get_entry_point(self, entry):
        if entry in self._ctx.entry_points:
            return self._ctx.entry_points[entry]
        else:
            raise self.Failure("Unknown entry point: %s" % entry)

    def _check_var(self, vname):
        if not vname in self._vars:
            raise self.Failure("Unknown variable: %s" % vname)

    def _check_new_var(self, vname):
        if vname in self._vars:
            raise self.Failure("Variable already exists: %s" % vname)

    def _get_var(self, vname):
        self._check_var(vname)
        return self._vars[vname]

    def _cmd_inputs(self, args):
        entry = self._get_arg(args, 0)
        for t in self._get_entry_point(entry)["inputs"]:
            print(t)

    def _cmd_output(self, args):
        entry = self._get_arg(args, 0)
        print(self._get_entry_point(entry)["output"])

    def _cmd_dummy(self, args):
        pass

    def _cmd_free(self, args):
        for vname in args:
            self._check_var(vname)
            del self._vars[vname]

    def _cmd_rename(self, args):
        oldname = self._get_arg(args, 0)
        newname = self._get_arg(args, 1)
        self._check_var(oldname)
        self._check_new_var(newname)
        self._vars[newname] = self._vars[oldname]
        del self._vars[oldname]

    def _cmd_call(self, args):
        entry = self._get_entry_point(self._get_arg(args, 0))
        entry_fname = entry["name"]
        num_ins = len(entry["inputs"])
        exp_len = 2 + num_ins

        if len(args) != exp_len:
            raise self.Failure("Invalid argument count, expected %d" % exp_len)

        out_vname = args[1]

        self._check_new_var(out_vname)

        in_vnames = args[2:]
        ins = [self._get_var(in_vname) for in_vname in in_vnames]

        try:
            runtime, vals = getattr(self._ctx, entry_fname)(*ins)
        except Exception as e:
            raise self.Failure(str(e))

        print("runtime: %d" % runtime)

        self._vars[out_vname] = vals

    def _store_val(self, f, value):
        # In case we are using the PyOpenCL backend, we first
        # need to convert OpenCL arrays to ordinary NumPy
        # arrays.  We do this in a nasty way.
        if isinstance(value, opaque):
            for component in value.data:
                self._store_val(f, component)
        elif (
            isinstance(value, np.number)
            or isinstance(value, bool)
            or isinstance(value, np.bool_)
            or isinstance(value, np.ndarray)
        ):
            # Ordinary NumPy value.
            f.write(construct_binary_value(value))
        else:
            # Assuming PyOpenCL array.
            f.write(construct_binary_value(value.get()))

    def _cmd_store(self, args):
        fname = self._get_arg(args, 0)

        with open(fname, "wb") as f:
            for i in range(1, len(args)):
                self._store_val(f, self._get_var(args[i]))

    def _restore_val(self, reader, typename):
        if typename in self._ctx.opaque_types:
            vs = []
            for t in self._ctx.opaque_types[typename]["payload"]:
                vs += [read_value(t, reader)]
            return self._opaque(typename, vs)
        else:
            return read_value(typename, reader)

    def _cmd_restore(self, args):
        if len(args) % 2 == 0:
            raise self.Failure("Invalid argument count")

        fname = args[0]
        args = args[1:]

        with open(fname, "rb") as f:
            reader = ReaderInput(f)
            while args != []:
                vname = args[0]
                typename = args[1]
                args = args[2:]

                if vname in self._vars:
                    raise self.Failure("Variable already exists: %s" % vname)

                try:
                    self._vars[vname] = self._restore_val(reader, typename)
                except ValueError:
                    raise self.Failure(
                        "Failed to restore variable %s.\n"
                        "Possibly malformed data in %s.\n" % (vname, fname)
                    )

            skip_spaces(reader)
            if reader.get_char() != b"":
                raise self.Failure("Expected EOF after reading values")

    def _cmd_types(self, args):
        for k in self._ctx.opaques.keys():
            print(k)

    # Types and values.
    #
    # An opaque value is represented by its payload: a flat list of primitive
    # values and arrays, as described by 'opaque_types'. A transparent value of
    # type t is represented by a single NumPy value (or a PyOpenCL array).

    def _opaque(self, tname, payload):
        return opaque(tname, self._ctx.opaques, *payload)

    def _opaque_type(self, tname):
        if tname in self._ctx.opaque_types:
            return self._ctx.opaque_types[tname]
        else:
            raise self.Failure("Unknown opaque type: %s" % tname)

    # Split a transparent type name into its rank and primitive type.
    def _split_type(self, tname):
        rank = 0
        while tname.startswith("[]"):
            rank += 1
            tname = tname[2:]
        if tname not in FUTHARK_PRIMTYPES:
            raise self.Failure("Unknown type: %s" % ("[]" * rank + tname))
        return rank, tname

    def _kind(self, tname):
        if tname in self._ctx.opaque_types:
            return self._ctx.opaque_types[tname]["kind"]
        elif self._split_type(tname)[0] == 0:
            return "primitive"
        else:
            return "array"

    def _array_type(self, tname):
        if self._kind(tname) != "array":
            raise self.Failure("Not an array type")
        if tname in self._ctx.opaque_types:
            t = self._ctx.opaque_types[tname]
            return t["rank"], t["elemtype"]
        else:
            return self._split_type(tname)

    def _record_fields(self, tname):
        if self._kind(tname) != "record":
            raise self.Failure("Not a record type")
        return self._opaque_type(tname)["fields"]

    def _sum_variants(self, tname):
        if self._kind(tname) != "sum":
            raise self.Failure("Not a sum type")
        return self._opaque_type(tname)["variants"]

    def _type_of(self, v):
        if isinstance(v, opaque):
            return v.desc
        dtype = v.dtype if hasattr(v, "dtype") else np.dtype(type(v))
        return "[]" * np.ndim(v) + numpy_type_to_type_name(dtype)

    def _get_typed_var(self, vname, tname):
        v = self._get_var(vname)
        if self._type_of(v) != tname:
            raise self.Failure(
                "Variable %s has type %s, expected %s"
                % (vname, self._type_of(v), tname)
            )
        return v

    def _payload_size(self, tname):
        if tname in self._ctx.opaque_types:
            return len(self._ctx.opaque_types[tname]["payload"])
        else:
            return 1

    def _payload(self, v):
        if isinstance(v, opaque):
            return list(v.data)
        else:
            return [v]

    def _from_payload(self, tname, payload):
        if tname in self._ctx.opaque_types:
            return self._opaque(tname, payload)
        else:
            return payload[0]

    # Split a payload into the parts corresponding to the given types.
    def _split_payload(self, tnames, payload):
        parts = []
        for t in tnames:
            n = self._payload_size(t)
            parts += [payload[:n]]
            payload = payload[n:]
        return parts

    # An empty value of the given transparent type, used for the unused
    # parts of the payload of a sum type.
    def _blank(self, tname):
        rank, pt = self._split_type(tname)
        dtype = FUTHARK_PRIMTYPES[pt]["numpy_type"]
        if rank == 0:
            return dtype(0)
        else:
            return np.zeros((0,) * rank, dtype=dtype)

    # Convert a PyOpenCL array to a NumPy array.
    def _to_host(self, x):
        if isinstance(x, np.ndarray) or not hasattr(x, "get"):
            return x
        else:
            return x.get()

    def _shape(self, v, rank):
        return tuple(int(d) for d in np.shape(self._payload(v)[0])[:rank])

    def _parse_ints(self, args):
        try:
            return [int(x) for x in args]
        except ValueError:
            raise self.Failure("Invalid integer in: %s" % " ".join(args))

    # Index a transparent array, producing a fresh value.
    def _index_transparent(self, x, idx):
        x = self._to_host(x[idx])
        if isinstance(x, np.ndarray):
            return x[()] if x.ndim == 0 else x.copy()
        else:
            return x

    def _cmd_kind(self, args):
        print(self._kind(self._get_arg(args, 0)))

    def _cmd_type(self, args):
        print(self._type_of(self._get_var(self._get_arg(args, 0))))

    def _cmd_rank(self, args):
        print(self._array_type(self._get_arg(args, 0))[0])

    def _cmd_elemtype(self, args):
        print(self._array_type(self._get_arg(args, 0))[1])

    def _cmd_shape(self, args):
        v = self._get_var(self._get_arg(args, 0))
        rank, _ = self._array_type(self._type_of(v))
        for d in self._shape(v, rank):
            print(d)

    def _cmd_new_array(self, args):
        dst = self._get_arg(args, 0)
        tname = self._get_arg(args, 1)
        self._check_new_var(dst)
        rank, et = self._array_type(tname)
        dims = tuple(self._parse_ints(args[2 : 2 + rank]))
        if len(dims) != rank or any(d < 0 for d in dims):
            raise self.Failure("Expected %d valid dimensions" % rank)
        vnames = args[2 + rank :]
        if len(vnames) != np.prod(dims, dtype=np.int64):
            raise self.Failure(
                "Expected %d values, but got %d"
                % (np.prod(dims, dtype=np.int64), len(vnames))
            )
        elems = [self._payload(self._get_typed_var(v, et)) for v in vnames]
        if tname in self._ctx.opaque_types:
            ts = self._ctx.opaque_types[tname]["payload"]
        else:
            ts = [tname]
        payload = []
        for i, t in enumerate(ts):
            r, pt = self._split_type(t)
            dtype = FUTHARK_PRIMTYPES[pt]["numpy_type"]
            if elems == []:
                payload += [np.zeros(dims + (0,) * (r - rank), dtype=dtype)]
            else:
                try:
                    x = np.array(
                        [self._to_host(e[i]) for e in elems], dtype=dtype
                    )
                except ValueError:
                    raise self.Failure("Array elements have irregular shapes")
                payload += [x.reshape(dims + x.shape[1:])]
        self._vars[dst] = self._from_payload(tname, payload)

    def _cmd_set(self, args):
        arr = self._get_arg(args, 0)
        v = self._get_var(arr)
        rank, et = self._array_type(self._type_of(v))
        val = self._payload(self._get_typed_var(self._get_arg(args, 1), et))
        idx = tuple(self._parse_ints(args[2:]))
        if len(idx) != rank:
            raise self.Failure(
                "%d indices expected but %d values provided."
                % (rank, len(idx))
            )
        shape = self._shape(v, rank)
        if not all(0 <= i < d for i, d in zip(idx, shape)):
            raise self.Failure(
                "Index %s out of bounds for shape %s" % (idx, shape)
            )
        payload = [self._to_host(x) for x in self._payload(v)]
        for x, y in zip(payload, val):
            x[idx] = self._to_host(y)
        self._vars[arr] = self._from_payload(self._type_of(v), payload)

    def _cmd_index(self, args):
        dst = self._get_arg(args, 0)
        v = self._get_var(self._get_arg(args, 1))
        self._check_new_var(dst)
        rank, et = self._array_type(self._type_of(v))
        idx = tuple(self._parse_ints(args[2:]))
        if len(idx) != rank:
            raise self.Failure(
                "%d indices expected but %d values provided."
                % (rank, len(idx))
            )
        shape = self._shape(v, rank)
        if not all(0 <= i < d for i, d in zip(idx, shape)):
            raise self.Failure(
                "Index %s out of bounds for shape %s" % (idx, shape)
            )
        payload = [self._index_transparent(x, idx) for x in self._payload(v)]
        self._vars[dst] = self._from_payload(et, payload)

    def _cmd_zip(self, args):
        dst = self._get_arg(args, 0)
        tname = self._get_arg(args, 1)
        self._check_new_var(dst)
        rank, _ = self._array_type(tname)
        fields = self._opaque_type(tname).get("fields")
        if fields is None:
            raise self.Failure("Cannot zip to this array type")
        vnames = args[2:]
        if len(vnames) != len(fields):
            raise self.Failure(
                "%d arrays expected but %d values provided."
                % (len(fields), len(vnames))
            )
        vs = [self._get_typed_var(v, t) for v, (_, t) in zip(vnames, fields)]
        if len(set(self._shape(v, rank) for v in vs)) > 1:
            raise self.Failure("Arrays have different shapes")
        payload = []
        for v in vs:
            payload += self._payload(v)
        self._vars[dst] = self._opaque(tname, payload)

    def _cmd_unzip(self, args):
        v = self._get_var(self._get_arg(args, 0))
        tname = self._type_of(v)
        self._array_type(tname)
        fields = self._opaque_type(tname).get("fields")
        if fields is None:
            raise self.Failure("Cannot unzip this array type")
        dsts = args[1:]
        if len(dsts) != len(fields):
            raise self.Failure(
                "%d arrays expected but %d values provided."
                % (len(fields), len(dsts))
            )
        for dst in dsts:
            self._check_new_var(dst)
        ts = [t for _, t in fields]
        for dst, t, p in zip(
            dsts, ts, self._split_payload(ts, self._payload(v))
        ):
            self._vars[dst] = self._from_payload(t, p)

    def _cmd_fields(self, args):
        for f, t in self._record_fields(self._get_arg(args, 0)):
            print(f, t)

    def _cmd_new(self, args):
        dst = self._get_arg(args, 0)
        tname = self._get_arg(args, 1)
        self._check_new_var(dst)
        fields = self._record_fields(tname)
        vnames = args[2:]
        if len(vnames) != len(fields):
            raise self.Failure(
                "%d fields expected but %d values provided."
                % (len(fields), len(vnames))
            )
        payload = []
        for v, (_, t) in zip(vnames, fields):
            payload += self._payload(self._get_typed_var(v, t))
        self._vars[dst] = self._opaque(tname, payload)

    def _cmd_project(self, args):
        dst = self._get_arg(args, 0)
        v = self._get_var(self._get_arg(args, 1))
        field = self._get_arg(args, 2)
        self._check_new_var(dst)
        fields = self._record_fields(self._type_of(v))
        ts = [t for _, t in fields]
        for (f, t), p in zip(
            fields, self._split_payload(ts, self._payload(v))
        ):
            if f == field:
                self._vars[dst] = self._from_payload(t, p)
                return
        raise self.Failure("No such field: %s" % field)

    def _cmd_variants(self, args):
        for name, payload in self._sum_variants(self._get_arg(args, 0)):
            print(name)
            for t, _ in payload:
                print("- %s" % t)

    # The variant of a sum-typed value. When there is more than one variant,
    # the first element of the payload is the index of the variant.
    def _variant_of(self, v):
        variants = self._sum_variants(self._type_of(v))
        if len(variants) == 1:
            return variants[0]
        else:
            return variants[int(v.data[0])]

    def _cmd_construct(self, args):
        dst = self._get_arg(args, 0)
        tname = self._get_arg(args, 1)
        vname = self._get_arg(args, 2)
        self._check_new_var(dst)
        variants = self._sum_variants(tname)
        for i, (name, vpayload) in enumerate(variants):
            if name == vname:
                break
        else:
            raise self.Failure("No such variant: %s" % vname)
        vnames = args[3:]
        if len(vnames) != len(vpayload):
            raise self.Failure(
                "%d values expected but %d provided."
                % (len(vpayload), len(vnames))
            )
        ts = self._opaque_type(tname)["payload"]
        payload = [self._blank(t) for t in ts]
        if len(variants) > 1:
            payload[0] = FUTHARK_PRIMTYPES[ts[0]]["numpy_type"](i)
        for v, (t, js) in zip(vnames, vpayload):
            for j, x in zip(js, self._payload(self._get_typed_var(v, t))):
                payload[j] = x
        self._vars[dst] = self._opaque(tname, payload)

    def _cmd_destruct(self, args):
        v = self._get_var(self._get_arg(args, 0))
        _, vpayload = self._variant_of(v)
        dsts = args[1:]
        if len(dsts) != len(vpayload):
            raise self.Failure(
                "%d variables expected but %d provided."
                % (len(vpayload), len(dsts))
            )
        for dst in dsts:
            self._check_new_var(dst)
        for dst, (t, js) in zip(dsts, vpayload):
            self._vars[dst] = self._from_payload(t, [v.data[j] for j in js])

    def _cmd_variant(self, args):
        print(self._variant_of(self._get_var(self._get_arg(args, 0)))[0])

    def _cmd_attributes(self, args):
        return self._get_entry_point(self._get_arg(args, 0))["attributes"]

    def _cmd_entry_points(self, args):
        for k in self._ctx.entry_points.keys():
            print(k)

    _commands = {
        "inputs": _cmd_inputs,
        "output": _cmd_output,
        "call": _cmd_call,
        "restore": _cmd_restore,
        "store": _cmd_store,
        "free": _cmd_free,
        "rename": _cmd_rename,
        "clear": _cmd_dummy,
        "pause_profiling": _cmd_dummy,
        "unpause_profiling": _cmd_dummy,
        "report": _cmd_dummy,
        "types": _cmd_types,
        "entry_points": _cmd_entry_points,
        "kind": _cmd_kind,
        "type": _cmd_type,
        "rank": _cmd_rank,
        "elemtype": _cmd_elemtype,
        "shape": _cmd_shape,
        "new_array": _cmd_new_array,
        "set": _cmd_set,
        "index": _cmd_index,
        "zip": _cmd_zip,
        "unzip": _cmd_unzip,
        "fields": _cmd_fields,
        "new": _cmd_new,
        "project": _cmd_project,
        "variants": _cmd_variants,
        "construct": _cmd_construct,
        "destruct": _cmd_destruct,
        "variant": _cmd_variant,
        "attributes": _cmd_attributes,
    }

    def _process_line(self, line):
        lex = shlex.shlex(line)
        lex.quotes = '"'
        lex.whitespace_split = True
        lex.commenters = ""
        words = list(lex)
        if words == []:
            raise self.Failure("Empty line")
        else:
            cmd = words[0]
            args = [w.strip('"') for w in words[1:]]
            if cmd in self._commands:
                self._commands[cmd](self, args)
            else:
                raise self.Failure("Unknown command: %s" % cmd)

    def run(self):
        while True:
            print("%%% OK", flush=True)
            line = sys.stdin.readline()
            if line == "":
                return
            try:
                self._process_line(line)
            except self.Failure as e:
                print("%%% FAILURE")
                print(e.msg)


# End of server.py