Files
apache--tvm/python/tvm/micro/project_api/client.py
T
Andrew Reusch a7297870c0 [microTVM] Project API infrastructure (#8380)
* Initial commit of API server impl.

* initial commit of api client

* Add TVM-side glue code to use Project API

* Change tvm.micro.Session to use Project API

* Rework how crt_config.h is used on the host.

 * use template crt_config.h for host test runtime; delete
   src/runtime/crt/host/crt_config.h so that it doesn't diverge from
   the template
 * bring template crt_config.h inline with the one actually in use
  * rename to MAX_STRLEN_DLTYPE
 * Create a dedicated TVM-side host crt_config.h in src/runtime/micro

* Modify Transport infrastructure to work with Project API

* Add host microTVM API server

* Zephyr implementation of microTVM API server

 * move all zephyr projects to apps/microtvm/zephyr/template_project

* consolidate CcompilerAnnotator

* Allow model library format with c backend, add test.

* Update unit tests

* fix incorrect doc

* Delete old Zephyr build infrastructure

* Delete old build abstractions

* Delete old Transport implementations and simplify module

* lint

* ASF header

* address gromero comments

* final fixes?

* fix is_shutdown

* fix user-facing API

* fix TempDirectory / operator

* Update micro_tflite tutorial

* lint

* fix test_crt and test_link_params

* undo global micro import, hopefully fix fixture

* lint

* fix more tests

* Address tmoreau89 comments and mehrdadh comments

 * fix random number generator prj.conf for physical hw
 * uncomment proper aot option
2021-08-07 11:51:32 -07:00

236 lines
7.7 KiB
Python

# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
import base64
import io
import json
import logging
import os
import pathlib
import subprocess
import sys
import typing
from . import server
_LOG = logging.getLogger(__name__)
class ProjectAPIErrorBase(Exception):
"""Base class for all Project API errors."""
class ConnectionShutdownError(ProjectAPIErrorBase):
"""Raised when a request is made but the connection has been closed."""
class MalformedReplyError(ProjectAPIErrorBase):
"""Raised when the server responds with an invalid reply."""
class MismatchedIdError(ProjectAPIErrorBase):
"""Raised when the reply ID does not match the request."""
class ProjectAPIServerNotFoundError(ProjectAPIErrorBase):
"""Raised when the Project API server can't be found in the repo."""
class UnsupportedProtocolVersionError(ProjectAPIErrorBase):
"""Raised when the protocol version returned by the API server is unsupported."""
class RPCError(ProjectAPIErrorBase):
def __init__(self, request, error):
self.request = request
self.error = error
def __str__(self):
return f"Calling project API method {self.request['method']}:" "\n" f"{self.error}"
class ProjectAPIClient:
"""A client for the Project API."""
def __init__(
self,
read_file: typing.BinaryIO,
write_file: typing.BinaryIO,
testonly_did_write_request: typing.Optional[typing.Callable] = None,
):
self.read_file = io.TextIOWrapper(read_file, encoding="UTF-8", errors="strict")
self.write_file = io.TextIOWrapper(
write_file, encoding="UTF-8", errors="strict", write_through=True
)
self.testonly_did_write_request = testonly_did_write_request
self.next_request_id = 1
@property
def is_shutdown(self):
return self.read_file is None
def shutdown(self):
if self.is_shutdown:
return
self.read_file.close()
self.write_file.close()
def _request_reply(self, method, params):
if self.is_shutdown:
raise ConnectionShutdownError("connection already closed")
request = {
"jsonrpc": "2.0",
"method": method,
"params": params,
"id": self.next_request_id,
}
self.next_request_id += 1
request_str = json.dumps(request)
self.write_file.write(request_str)
_LOG.debug("send -> %s", request_str)
self.write_file.write("\n")
if self.testonly_did_write_request:
self.testonly_did_write_request() # Allow test to assert on server processing.
reply_line = self.read_file.readline()
_LOG.debug("recv <- %s", reply_line)
if not reply_line:
self.shutdown()
raise ConnectionShutdownError("got EOF reading reply from API server")
reply = json.loads(reply_line)
if reply.get("jsonrpc") != "2.0":
raise MalformedReplyError(
f"Server reply should include 'jsonrpc': '2.0'; "
f"saw jsonrpc={reply.get('jsonrpc')!r}"
)
if reply["id"] != request["id"]:
raise MismatchedIdError(
f"Reply id ({reply['id']}) does not equal request id ({request['id']}"
)
if "error" in reply:
raise server.JSONRPCError.from_json(f"calling method {method}", reply["error"])
elif "result" not in reply:
raise MalformedReplyError(f"Expected 'result' key in server reply, got {reply!r}")
return reply["result"]
def server_info_query(self, tvm_version: str):
reply = self._request_reply("server_info_query", {"tvm_version": tvm_version})
if reply["protocol_version"] != server.ProjectAPIServer._PROTOCOL_VERSION:
raise UnsupportedProtocolVersionError(
f'microTVM API Server supports protocol version {reply["protocol_version"]}; '
f"want {server.ProjectAPIServer._PROTOCOL_VERSION}"
)
return reply
def generate_project(
self,
model_library_format_path: str,
standalone_crt_dir: str,
project_dir: str,
options: dict = None,
):
return self._request_reply(
"generate_project",
{
"model_library_format_path": model_library_format_path,
"standalone_crt_dir": standalone_crt_dir,
"project_dir": project_dir,
"options": (options if options is not None else {}),
},
)
def build(self, options: dict = None):
return self._request_reply("build", {"options": (options if options is not None else {})})
def flash(self, options: dict = None):
return self._request_reply("flash", {"options": (options if options is not None else {})})
def open_transport(self, options: dict = None):
return self._request_reply(
"open_transport", {"options": (options if options is not None else {})}
)
def close_transport(self):
return self._request_reply("close_transport", {})
def read_transport(self, n, timeout_sec):
reply = self._request_reply("read_transport", {"n": n, "timeout_sec": timeout_sec})
reply["data"] = base64.b85decode(reply["data"])
return reply
def write_transport(self, data, timeout_sec):
return self._request_reply(
"write_transport",
{"data": str(base64.b85encode(data), "utf-8"), "timeout_sec": timeout_sec},
)
# NOTE: windows support untested
SERVER_LAUNCH_SCRIPT_FILENAME = (
f"launch_microtvm_api_server.{'sh' if os.system != 'win32' else '.bat'}"
)
SERVER_PYTHON_FILENAME = "microtvm_api_server.py"
def instantiate_from_dir(project_dir: typing.Union[pathlib.Path, str], debug: bool = False):
"""Launch server located in project_dir, and instantiate a Project API Client connected to it."""
args = None
project_dir = pathlib.Path(project_dir)
python_script = project_dir / SERVER_PYTHON_FILENAME
if python_script.is_file():
args = [sys.executable, str(python_script)]
launch_script = project_dir / SERVER_LAUNCH_SCRIPT_FILENAME
if launch_script.is_file():
args = [str(launch_script)]
if args is None:
raise ProjectAPIServerNotFoundError(
f"No Project API server found in project directory: {project_dir}"
"\n"
f"Tried: {SERVER_LAUNCH_SCRIPT_FILENAME}, {SERVER_PYTHON_FILENAME}"
)
api_server_read_fd, tvm_write_fd = os.pipe()
tvm_read_fd, api_server_write_fd = os.pipe()
args.extend(["--read-fd", str(api_server_read_fd), "--write-fd", str(api_server_write_fd)])
if debug:
args.append("--debug")
api_server_proc = subprocess.Popen(
args, bufsize=0, pass_fds=(api_server_read_fd, api_server_write_fd), cwd=project_dir
)
os.close(api_server_read_fd)
os.close(api_server_write_fd)
return ProjectAPIClient(
os.fdopen(tvm_read_fd, "rb", buffering=0), os.fdopen(tvm_write_fd, "wb", buffering=0)
)