diff --git a/CMakeLists.txt b/CMakeLists.txt index 127ba50b3720..7a137d434d42 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -45,6 +45,7 @@ tvm_option(USE_MICRO "Build with Micro TVM support" OFF) tvm_option(INSTALL_DEV "Install compiler infrastructure" OFF) tvm_option(HIDE_PRIVATE_SYMBOLS "Compile with -fvisibility=hidden." OFF) tvm_option(USE_TF_TVMDSOOP "Build with TensorFlow TVMDSOOp" OFF) +tvm_option(USE_PT_TVMDSOOP "Build with PyTorch TVMDSOOp" OFF) tvm_option(USE_FALLBACK_STL_MAP "Use TVM's POD compatible Map" OFF) tvm_option(USE_ETHOSN "Build with Arm Ethos-N" OFF) tvm_option(INDEX_DEFAULT_I64 "Defaults the index datatype to int64" ON) @@ -412,6 +413,7 @@ include(cmake/modules/contrib/NNPack.cmake) include(cmake/modules/contrib/HybridDump.cmake) include(cmake/modules/contrib/TFLite.cmake) include(cmake/modules/contrib/TF_TVMDSOOP.cmake) +include(cmake/modules/contrib/PT_TVMDSOOP.cmake) include(cmake/modules/contrib/CoreML.cmake) include(cmake/modules/contrib/BNNS.cmake) include(cmake/modules/contrib/ONNX.cmake) diff --git a/apps/pt_tvmdsoop/CMakeLists.txt b/apps/pt_tvmdsoop/CMakeLists.txt new file mode 100644 index 000000000000..05b3b0babc01 --- /dev/null +++ b/apps/pt_tvmdsoop/CMakeLists.txt @@ -0,0 +1,34 @@ +# 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. +cmake_minimum_required(VERSION 3.2) +project(pt_tvmdsoop C CXX) + +set(BUILD_PT_TVMDSOOP_ONLY ON) +set(CMAKE_CURRENT_SOURCE_DIR ${TVM_ROOT}) +set(CMAKE_CURRENT_BINARY_DIR ${TVM_ROOT}/build) + +include_directories(SYSTEM ${TVM_ROOT}/3rdparty/dlpack/include/) +include_directories(SYSTEM ${TVM_ROOT}/3rdparty/dmlc-core/include/) +include_directories(${TVM_ROOT}/include) + +link_directories(${TVM_ROOT}/build) + +include(${TVM_ROOT}/cmake/utils/Utils.cmake) +include(${TVM_ROOT}/cmake/utils/FindCUDA.cmake) +include(${TVM_ROOT}/cmake/modules/CUDA.cmake) + +include(${TVM_ROOT}/cmake/modules/contrib/PT_TVMDSOOP.cmake) diff --git a/apps/pt_tvmdsoop/prepare_and_test_pt_tvm_class.sh b/apps/pt_tvmdsoop/prepare_and_test_pt_tvm_class.sh new file mode 100755 index 000000000000..666f774017c8 --- /dev/null +++ b/apps/pt_tvmdsoop/prepare_and_test_pt_tvm_class.sh @@ -0,0 +1,46 @@ +#!/bin/bash +# 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. + +TVM_ROOT=$(cd $(dirname $0)/../..; pwd) +echo "TVM_ROOT=${TVM_ROOT}" + +export PYTHONPATH=${TVM_ROOT}/python + +if [ ! -f $TVM_ROOT/build/libtvm.so ]; then + echo "$TVM_ROOT/build/libtvm.so missing" + exit 1 +fi + +if [ ! -f $TVM_ROOT/build/libtvm_runtime.so ]; then + echo "$TVM_ROOT/build/libtvm_runtime.so missing" + exit 1 +fi + +python3 -c "import tvm; print(tvm.runtime.enabled('gpu'))" | grep -e 1 + +if [ "$?" -eq 0 ]; then + echo "Build PT_TVMDSOOP with gpu support and execute tests" + CMAKE_OPTIONS="-DUSE_CUDA=ON -DUSE_CUDNN=ON -DPython3_EXECUTABLE=python3 -DTVM_ROOT=${TVM_ROOT}" + mkdir -p build + cd build; cmake .. ${CMAKE_OPTIONS} && make + cp *.so $TVM_ROOT/build/ + cd .. + + LD_LIBRARY_PATH=${TVM_ROOT}/build:./build:$LD_LIBRARY_PATH python3 -m pytest -v ./tests +fi + diff --git a/apps/pt_tvmdsoop/tests/test_torch_compile_cpu.py b/apps/pt_tvmdsoop/tests/test_torch_compile_cpu.py new file mode 100644 index 000000000000..5ad88b45dc80 --- /dev/null +++ b/apps/pt_tvmdsoop/tests/test_torch_compile_cpu.py @@ -0,0 +1,68 @@ +#!/usr/bin/env 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. +"""Test script for torch module""" +import torch +import time +import tvm +from tvm.contrib.torch import compile + + +class Model(torch.nn.Module): + def __init__(self): + super().__init__() + + def forward(self, x: torch.Tensor): + return x * x + + +model = Model() +x = torch.rand([1, 3, 224, 224]) +model_jit = torch.jit.trace(model, x) +print(model_jit.graph) + +print("run torchscript...") +for i in range(20): + t = time.time() + model_jit(x) + print(time.time() - t) + + +option = { + "input_infos": [ + ("x", (1, 3, 224, 224)), + ], + "default_dtype": "float16", + "export_dir": "pytorch_compiled", + "num_outputs": 1, + "tuning_n_trials": 1, # set zero to skip tuning + "tuning_log_file": "tuning.log", + "target": "llvm", + "device": tvm.cpu(), +} + +pytorch_tvm_module = compile(model_jit, option) +torch.jit.script(pytorch_tvm_module).save("model_tvm.pt") + + +print("Run PyTorch...") +for i in range(20): + t = time.time() + outputs = pytorch_tvm_module.forward([x.cpu()]) + print(1000 * (time.time() - t)) +print(outputs[0].shape) diff --git a/apps/pt_tvmdsoop/tests/test_torch_compile_gpu.py b/apps/pt_tvmdsoop/tests/test_torch_compile_gpu.py new file mode 100644 index 000000000000..b2ceb7f5cd6b --- /dev/null +++ b/apps/pt_tvmdsoop/tests/test_torch_compile_gpu.py @@ -0,0 +1,63 @@ +#!/usr/bin/env 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. +"""Test script for torch module""" +import torch +import time +from torchvision.models import resnet50 +import tvm +from tvm.contrib.torch import compile + + +model = resnet50().half().cuda() +x = torch.rand([1, 3, 224, 224]).half().cuda() +model_jit = torch.jit.trace(model, x) +print(model_jit.graph) + +print("run torchscript...") +for i in range(20): + t = time.time() + model_jit(x) + torch.cuda.synchronize() + print(time.time() - t) + + +option = { + "input_infos": [ + ("x", (1, 3, 224, 224)), + ], + "default_dtype": "float16", + "export_dir": "pytorch_compiled", + "num_outputs": 1, + "tuning_n_trials": 1, # set zero to skip tuning + "tuning_log_file": "tuning.log", + "target": "cuda", + "device": tvm.cuda(0), +} + +pytorch_tvm_module = compile(model_jit, option) +torch.jit.script(pytorch_tvm_module).save("model_tvm.pt") + + +print("Run PyTorch...") +for i in range(20): + t = time.time() + outputs = pytorch_tvm_module.forward([x]) + torch.cuda.synchronize() + print(1000 * (time.time() - t)) +print(outputs[0].shape) diff --git a/apps/pt_tvmdsoop/tests/test_torch_graph_module.py b/apps/pt_tvmdsoop/tests/test_torch_graph_module.py new file mode 100644 index 000000000000..4e3b51227cbe --- /dev/null +++ b/apps/pt_tvmdsoop/tests/test_torch_graph_module.py @@ -0,0 +1,129 @@ +#!/usr/bin/env 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. +"""Test script for torch module""" +import tempfile +import os +import logging +import torch +import numpy as np +import tvm +import tvm.testing +from tvm import te, relay +import tvm.contrib.torch +from tvm.contrib import graph_runtime + +TVM_ASSETS = ["mod.so", "graph.json", "params"] + + +def test_use_pt_graph_module(): + """main test function""" + + def build_export_graph(device): + """relay build & export graph""" + x = relay.var("x", shape=(10, 5)) + y = relay.var("y", shape=(1, 5)) + z = relay.add(x, y) + z = relay.exp(z) + func = relay.Function([x, y], z) + x_data = np.random.rand(10, 5).astype("float32") + y_data = np.random.rand(1, 5).astype("float32") + params = {"y": y_data} + + pt_device = torch.device(device) + if pt_device.type == "cuda": + target = "cuda" + ctx = tvm.cuda(pt_device.index) + else: + target = "llvm" + ctx = tvm.cpu(0) + + graph, lib, params = relay.build(tvm.IRModule.from_expr(func), target=target, params=params) + mod = graph_runtime.create(graph, lib, device=ctx) + mod.set_input(**params) + mod.set_input(x=x_data) + mod.run() + res = mod.get_output(0).asnumpy() + ref_res = np.exp(y_data + x_data) + tvm.testing.assert_allclose(res, ref_res, atol=1e-5, rtol=1e-5) + + # export to tempdir + export_dir = tempfile.mkdtemp("tvm_export") + lib.export_library(os.path.join(export_dir, TVM_ASSETS[0])) + with open(os.path.join(export_dir, TVM_ASSETS[1]), "w") as fout: + fout.write(graph) + with open(os.path.join(export_dir, TVM_ASSETS[2]), "wb") as fout: + fout.write(relay.save_param_dict(params)) + + return export_dir + + def test_pt_run(device, trace=True, to_device=None): + """test add lib with Pytorch wrapper""" + print("\n############## Test on device:", device, "#################") + export_dir = build_export_graph(device) + engine = tvm.contrib.torch.GraphModule(num_inputs=2, num_outputs=1).to(device) + + x = np.random.rand(10, 5).astype("float32") + y = np.random.rand(1, 5).astype("float32") + + expect = np.exp(y + x) + + def get_inputs_by_device(device): + inps = [torch.Tensor(x), torch.Tensor(y)] + if device == "cpu": + return inps + else: + device_type, device_id = device.split(":") + assert device_type == "cuda" + return [inp.cuda(int(device_id)) for inp in inps] + + assets = [os.path.join(export_dir, i) for i in TVM_ASSETS] + engine.init((x.shape, y.shape), *assets) + + outputs = engine.forward(get_inputs_by_device(device)) + tvm.testing.assert_allclose(outputs[0].cpu(), expect, atol=1e-5, rtol=1e-5) + + if trace: + print("\n################ Test trace and load #################") + scripted = torch.jit.script(engine) + scripted_dir = tempfile.mkdtemp("scripted") + scripted_path = os.path.join(scripted_dir, "model.pt") + scripted.save(scripted_path) + loaded = torch.jit.load(scripted_path) + outputs = loaded.forward(get_inputs_by_device(device)) + tvm.testing.assert_allclose(outputs[0].cpu(), expect, atol=1e-5, rtol=1e-5) + del scripted + del loaded + + if to_device: + print( + "\n################ Test move from [{}] to [{}] #################".format( + device, to_device + ) + ) + engine = engine.to(to_device) + outputs = engine.forward(get_inputs_by_device(to_device)) + tvm.testing.assert_allclose(outputs[0].cpu(), expect, atol=1e-5, rtol=1e-5) + del engine + + test_pt_run(device="cuda:0", trace=True, to_device="cuda:1") + test_pt_run(device="cpu", trace=True) + + +if __name__ == "__main__": + test_use_pt_graph_module() diff --git a/apps/pt_tvmdsoop/tests/test_torch_script.py b/apps/pt_tvmdsoop/tests/test_torch_script.py new file mode 100644 index 000000000000..34b959714a18 --- /dev/null +++ b/apps/pt_tvmdsoop/tests/test_torch_script.py @@ -0,0 +1,116 @@ +#!/usr/bin/env 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. +"""Test script for torch module""" +import os +import torch +import time +import numpy as np +import tvm +import tvm.testing +import tempfile +from tvm.contrib.torch import PyTorchTVMModule, compile + + +class Model(torch.nn.Module): + def forward(self, x, y): + return torch.matmul(x, y.softmax(1)) + + +model = Model() +model.cuda().half() +x = torch.rand([1280, 2464, 4]).cuda().half() +y = torch.rand([1280, 4, 1]).cuda().half() +for i in range(20): + t = time.time() + o = model(x, y) + torch.cuda.synchronize() + print(1000 * (time.time() - t)) +print(o.shape) + + +model_jit = torch.jit.script(model) +print(model_jit.graph) +input_shapes = [("x", list(x.shape)), ("y", list(y.shape))] +dtype = "float16" +export_dir = tempfile.mkdtemp("pytorch_compiled") +print("tmp export_dir:", export_dir) + + +mod = PyTorchTVMModule() +print("Converting...") +mod.from_pytorch(model_jit, input_shapes, dtype) + +log_file = os.path.join(export_dir, "tuning.log") +if not os.path.exists(log_file): + print("Tuning...") + mod.tune_tvm(log_file=log_file, n_trial=20) + +print("Building...") +tvm_mod = mod.build_tvm(export_dir) +pytorch_mod = mod.build_pytorch_module(num_inputs=2, num_outputs=1) + + +## Or you can load from a prebuilt tvm module +# mod = PyTorchTVMModule() +# tvm_mod = mod.load_tvm(export_dir) +# pytorch_mod = mod.build_pytorch_module(num_inputs=2, num_outputs=1, input_infos=input_shapes) + + +print("Run TVM...") +tvm_x = tvm.nd.array(x.cpu().numpy().astype(dtype), device=tvm.gpu(0)) +tvm_y = tvm.nd.array(y.cpu().numpy().astype(dtype), device=tvm.gpu(0)) +for i in range(20): + t = time.time() + tvm_mod.run(x=tvm_x, y=tvm_y) + print(1000 * (time.time() - t)) +tvm_output = tvm_mod.get_output(0) +print(tvm_output.shape) + + +print("Run PyTorch...") +for i in range(20): + t = time.time() + outputs = pytorch_mod.forward([x, y]) + torch.cuda.synchronize() + print(1000 * (time.time() - t)) +print(outputs[0].shape) + + +class EnsembleModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.layer = torch.jit.script(pytorch_mod) + + def forward(self, x, y, z) -> torch.Tensor: + if x > 1: + out = self.layer(y, z)[0] + else: + out = torch.ones([1280, 2464, 1]) + return out + + +print("Exporting...") +scripted = torch.jit.script(EnsembleModel()) +print(scripted.graph) +scripted_path = os.path.join(export_dir, "model_tvm.pt") +scripted.save(scripted_path) + + +# print(o == outputs[0]) +# print(o - outputs[0]) diff --git a/apps/pt_tvmdsoop/tests/test_torch_vm_module.py b/apps/pt_tvmdsoop/tests/test_torch_vm_module.py new file mode 100644 index 000000000000..81d9dadb02c1 --- /dev/null +++ b/apps/pt_tvmdsoop/tests/test_torch_vm_module.py @@ -0,0 +1,122 @@ +#!/usr/bin/env 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. +"""Test script for torch vm module""" +import tempfile +import os +import logging +import torch +import numpy as np +import tvm +from tvm.contrib.torch.pytorch_tvm import TVM_ASSETS +import tvm.testing +from tvm import te, relay +import tvm.contrib.torch +from tvm.contrib import graph_runtime + +TVM_ASSETS = ["mod.so", "code.ro"] + + +def test_use_pt_vm_module(): + """main test function""" + + def build_export_vm(device): + """relay build & export graph""" + x = relay.var("x", shape=(10, 5)) + y = relay.var("y", shape=(1, 5)) + z = relay.add(x, y) + z = relay.exp(z) + func = relay.Function([x, y], z) + x_data = np.random.rand(10, 5).astype("float32") + y_data = np.random.rand(1, 5).astype("float32") + + pt_device = torch.device(device) + if pt_device.type == "cuda": + target = "cuda" + ctx = tvm.cuda(pt_device.index) + else: + target = "llvm" + ctx = tvm.cpu(0) + exe = relay.vm.compile(tvm.IRModule.from_expr(func), target=target, params={}) + code, lib = exe.save() + export_dir = tempfile.mkdtemp("tvm_export") + # export to tempdir + lib.export_library(os.path.join(export_dir, TVM_ASSETS[0])) + with open(os.path.join(export_dir, TVM_ASSETS[1]), "wb") as fout: + fout.write(code) + vm = tvm.runtime.vm.VirtualMachine(exe, ctx) + res = vm.run(x_data, y_data) + ref_res = np.exp(y_data + x_data) + tvm.testing.assert_allclose(res.numpy(), ref_res, atol=1e-5, rtol=1e-5) + return export_dir + + def test_pt_run(device, trace=True, to_device=None, inp_on_cuda=False): + """test add lib with Pytorch wrapper""" + print("\n############## Test on device:", device, "#################") + export_dir = build_export_vm(device) + engine = tvm.contrib.torch.VMModule(num_inputs=2, num_outputs=1).to(device) + + x = np.random.rand(10, 5).astype("float32") + y = np.random.rand(1, 5).astype("float32") + + expect = np.exp(y + x) + + def get_inputs_by_device(device): + inps = [torch.Tensor(x), torch.Tensor(y)] + if device == "cpu": + return inps + else: + device_type, device_id = device.split(":") + assert device_type == "cuda" + return [inp.cuda(int(device_id)) for inp in inps] + + assets = [os.path.join(export_dir, i) for i in TVM_ASSETS] + engine.init((x.shape, y.shape), *assets) + + outputs = engine.forward(get_inputs_by_device(device)) + tvm.testing.assert_allclose(outputs[0].cpu(), expect, atol=1e-5, rtol=1e-5) + + if trace: + print("\n################ Test trace and load #################") + scripted = torch.jit.script(engine) + scripted_dir = tempfile.mkdtemp("scripted") + scripted_path = os.path.join(scripted_dir, "model.pt") + scripted.save(scripted_path) + loaded = torch.jit.load(scripted_path) + outputs = loaded.forward(get_inputs_by_device(device)) + tvm.testing.assert_allclose(outputs[0].cpu(), expect, atol=1e-5, rtol=1e-5) + del scripted + del loaded + + if to_device: + print( + "\n################ Test move from [{}] to [{}] #################".format( + device, to_device + ) + ) + engine = engine.to(to_device) + outputs = engine.forward(get_inputs_by_device(to_device)) + tvm.testing.assert_allclose(outputs[0].cpu(), expect, atol=1e-5, rtol=1e-5) + del engine + + test_pt_run(device="cuda:0", trace=True, to_device="cuda:1", inp_on_cuda=True) + test_pt_run(device="cpu", trace=True, inp_on_cuda=False) + + +if __name__ == "__main__": + test_use_pt_vm_module() diff --git a/apps/pt_tvmdsoop/tests/test_trace_tvm_module.py b/apps/pt_tvmdsoop/tests/test_trace_tvm_module.py new file mode 100644 index 000000000000..0a12ec529fa0 --- /dev/null +++ b/apps/pt_tvmdsoop/tests/test_trace_tvm_module.py @@ -0,0 +1,58 @@ +#!/usr/bin/env 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. +"""Test script for torch module""" +import torch +import time +import tvm +from tvm.contrib.torch import compile, TraceTvmModule, pytorch_tvm + + +class Model(torch.nn.Module): + def __init__(self): + super().__init__() + + def forward(self, x: torch.Tensor, y: torch.Tensor): + return x * y + + +model = Model() +x = torch.rand([1, 2, 3]) +y = torch.rand([1, 2, 3]) +model_jit = torch.jit.script(model) + +option = { + "input_infos": [("x", (1, 2, 3)), ("y", (1, 2, 3))], + "default_dtype": "float32", + "export_dir": "pytorch_compiled", + "num_outputs": 1, + "tuning_n_trials": 0, # set zero to skip tuning + "tuning_log_file": "tuning.log", + "target": "llvm", + "device": tvm.cpu(), +} + +# use TraceTvmModule to convert List[Tensor] input/output +# to tuple of Tensors +pytorch_tvm_module = compile(model_jit, option) +scripted = torch.jit.script(pytorch_tvm_module) +traced = torch.jit.trace(TraceTvmModule(scripted), (x, y)) + +res_traced = traced.forward(x, y) +res_expected = pytorch_tvm_module.forward([x, y])[0] +tvm.testing.assert_allclose(res_traced, res_expected) diff --git a/cmake/config.cmake b/cmake/config.cmake index e55f1197d90e..59779e034c4c 100644 --- a/cmake/config.cmake +++ b/cmake/config.cmake @@ -266,6 +266,9 @@ set(USE_THRUST OFF) # Whether to build the TensorFlow TVMDSOOp module set(USE_TF_TVMDSOOP OFF) +# Whether to build the PyTorch custom class module +set(USE_PT_TVMDSOOP OFF) + # Whether to use STL's std::unordered_map or TVM's POD compatible Map set(USE_FALLBACK_STL_MAP OFF) diff --git a/cmake/modules/LibInfo.cmake b/cmake/modules/LibInfo.cmake index 163a56dbd1d4..bf548b232512 100644 --- a/cmake/modules/LibInfo.cmake +++ b/cmake/modules/LibInfo.cmake @@ -60,6 +60,7 @@ function(add_lib_info src_file) TVM_INFO_INSTALL_DEV="${INSTALL_DEV}" TVM_INFO_HIDE_PRIVATE_SYMBOLS="${HIDE_PRIVATE_SYMBOLS}" TVM_INFO_USE_TF_TVMDSOOP="${USE_TF_TVMDSOOP}" + TVM_INFO_USE_PT_TVMDSOOP="${USE_PT_TVMDSOOP}" TVM_INFO_USE_FALLBACK_STL_MAP="${USE_FALLBACK_STL_MAP}" TVM_INFO_USE_BYODT_POSIT="${USE_BYODT_POSIT}" TVM_INFO_USE_BLAS="${USE_BLAS}" diff --git a/cmake/modules/contrib/PT_TVMDSOOP.cmake b/cmake/modules/contrib/PT_TVMDSOOP.cmake new file mode 100644 index 000000000000..7ff88693fe4e --- /dev/null +++ b/cmake/modules/contrib/PT_TVMDSOOP.cmake @@ -0,0 +1,59 @@ +# 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. + +if(NOT USE_PT_TVMDSOOP STREQUAL "OFF") + find_package(Python3 COMPONENTS Interpreter Development) + include_directories(${Python3_INCLUDE_DIRS}) + + message(STATUS "Python3_INCLUDE_DIRS: ${Python3_INCLUDE_DIRS}") + + execute_process(COMMAND ${Python3_EXECUTABLE} -c "import torch; print(torch.__path__[0].strip())" + OUTPUT_VARIABLE PT_PATH + RESULT_VARIABLE PT_STATUS) + if (NOT ${PT_STATUS} EQUAL 0) + message(FATAL_ERROR "Fail to get pytorch path") + endif() + + string(REGEX REPLACE "\n" "" PT_PATH "${PT_PATH}") + + set(PT_COMPILE_FLAGS_STR "-I${PT_PATH}/include -D_GLIBCXX_USE_CXX11_ABI=0") + set(PT_LINK_FLAGS_STR "-L${PT_PATH}/lib -l:libtorch.so -l:libtorch_python.so") + + if(NOT USE_CUDA STREQUAL "OFF") + add_definitions(-DPT_TVMDSOOP_ENABLE_GPU) + endif() + + + string(REGEX REPLACE "\n" " " PT_FLAGS "${PT_COMPILE_FLAGS} ${PT_LINK_FLAGS}") + separate_arguments(PT_COMPILE_FLAGS UNIX_COMMAND ${PT_COMPILE_FLAGS_STR}) + separate_arguments(PT_LINK_FLAGS UNIX_COMMAND ${PT_LINK_FLAGS_STR}) + + + set(LIBRARY_NAME pt_tvmdsoop) + file(GLOB_RECURSE PTTVM_SRCS ${CMAKE_CURRENT_SOURCE_DIR}/src/contrib/torch/**/*.cc) + add_library(${LIBRARY_NAME} SHARED ${PTTVM_SRCS}) + set(PTTVM_LINK_FLAGS -ltvm -L${CMAKE_CURRENT_BINARY_DIR}) + + if (NOT BUILD_PT_TVMDSOOP_ONLY STREQUAL "ON") + add_dependencies(${LIBRARY_NAME} tvm) + endif() + + target_compile_options(${LIBRARY_NAME} PUBLIC ${PTTVM_COMPILE_FLAGS} ${PT_COMPILE_FLAGS}) + target_link_libraries(${LIBRARY_NAME} PUBLIC ${PTTVM_LINK_FLAGS} ${PT_LINK_FLAGS}) + +endif() + diff --git a/python/tvm/contrib/torch/__init__.py b/python/tvm/contrib/torch/__init__.py new file mode 100644 index 000000000000..720ac29cc6e2 --- /dev/null +++ b/python/tvm/contrib/torch/__init__.py @@ -0,0 +1,51 @@ +# 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. +# pylint: disable=wrong-import-position,redefined-builtin,invalid-name +"""Module container of Pytorch custom class""" +import os +import platform +import torch +from tvm._ffi import libinfo +from tvm.relay.frontend import pytorch + + +def _load_platform_specific_library(lib_name="libpt_tvmdsoop"): + system = platform.system() + if system == "Darwin": + lib_file_name = lib_name + ".dylib" + elif system == "Windows": + lib_file_name = lib_name + ".dll" + else: + lib_file_name = lib_name + ".so" + lib_path = libinfo.find_lib_path()[0] + lib_dir = os.path.dirname(lib_path) + lib_file_path = os.path.join(lib_dir, lib_file_name) + torch.classes.load_library(lib_file_path) + + +_load_platform_specific_library() + +from . import module + +GraphModule = module.GraphModule +VMModule = module.VMModule +TraceTvmModule = module.TraceTvmModule + +from . import pytorch_tvm + +PyTorchTVMModule = pytorch_tvm.PyTorchTVMModule +compile = pytorch_tvm.compile diff --git a/python/tvm/contrib/torch/module.py b/python/tvm/contrib/torch/module.py new file mode 100644 index 000000000000..3da9c6f591ce --- /dev/null +++ b/python/tvm/contrib/torch/module.py @@ -0,0 +1,121 @@ +# 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. +# pylint: disable=invalid-name +"""Module container of PyTorch custom class""" +from typing import List +import torch + + +class GraphModule(torch.nn.Module): + r"""Module container of Pytorch class which wraps exported + TVM op implementation library to be called on Pytorch side""" + + @classmethod + def shape_repr(cls, input_shapes): + return torch.ops.tvm_dsoop.tvm_shape_repr(input_shapes) + + def __init__(self, num_inputs, num_outputs, device=None): + super().__init__() + self.dummy_param = torch.nn.Parameter(torch.empty(0)) + self.engine = None + + if device is not None: + self.to(device) + self.engine = torch.classes.tvm_dsoop.TvmGraphModule(num_inputs, num_outputs, self.device) + + def init(self, input_shapes, lib_path, graph_path, params_path): + r"""Load tvm module""" + self.engine.load_tvm_module(input_shapes, lib_path, graph_path, params_path) + + def forward(self, inputs: List[torch.Tensor]): + r"""Call tvm module to forward""" + return self.engine.forward(inputs) + + @property + def device(self): + r"""Get the device string""" + return str(self.dummy_param.device) + + def _apply(self, fn): + r"""Override to device function, manually move tvm module to desired device""" + super()._apply(fn) + if self.engine is not None: + self.engine.to(self.device) + return self + + +class VMModule(torch.nn.Module): + r"""Module container of Pytorch class which wraps exported + TVM op implementation library to be called on Pytorch side""" + + @classmethod + def shape_repr(cls, input_shapes): + return torch.ops.tvm_dsoop.tvm_shape_repr(input_shapes) + + def __init__(self, num_inputs, num_outputs, device=None): + super().__init__() + self.dummy_param = torch.nn.Parameter(torch.empty(0)) + self.engine = None + + if device is not None: + self.to(device) + self.engine = torch.classes.tvm_dsoop.TvmVMModule(num_inputs, num_outputs, self.device) + + def init(self, input_shapes, lib_path, code_path): + r"""Load tvm module""" + self.engine.load_tvm_module(input_shapes, lib_path, code_path) + + def forward(self, inputs: List[torch.Tensor]): + r"""Call tvm module to forward""" + return self.engine.forward(inputs) + + @property + def device(self): + r"""Get the device string""" + return str(self.dummy_param.device) + + def _apply(self, fn): + r"""Override to device function, manually move tvm module to desired device""" + super()._apply(fn) + if self.engine is not None: + self.engine.to(self.device) + return self + + +class TraceTvmModule(torch.nn.Module): + r"""Wrapper for trace GraphModule + + GraphModule and VMModule only supports List[Tensor] inputs and cannot be traced. + This is a wrapper class for trace GraphModule or VMModule in order to support + arbitrary number of inputs + + Example: + import tvm.contrib.torch + tvm_module = tvm.contrib.torch.GraphModule(1, 1, 'cuda:0') + tvm_module.init(input_shapes, lib_path, graph_path, params_path) + + trace_wrapper = tvm.contrib.torch.TraceGraphModule(torch.jit.script(tvm_module)) + traced = torch.jit.trace(trace_wrapper, example_inputs) + """ + + def __init__(self, tvm_module): + super().__init__() + self.tvm_module = tvm_module + + def forward(self, *inputs): + outputs = self.tvm_module(inputs) + return outputs[0] if len(outputs) == 1 else tuple(outputs) diff --git a/python/tvm/contrib/torch/pytorch_tvm.py b/python/tvm/contrib/torch/pytorch_tvm.py new file mode 100644 index 000000000000..1e50c98ab883 --- /dev/null +++ b/python/tvm/contrib/torch/pytorch_tvm.py @@ -0,0 +1,249 @@ +#!/usr/bin/env 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. +# pylint: disable=redefined-builtin +"""`compile` api that convert torch module to torch tvm module""" +import os +import tvm +import tvm.testing +from tvm import relay, autotvm +from tvm.runtime import load_module +from tvm.autotvm.tuner import XGBTuner, GATuner, RandomTuner, GridSearchTuner +from tvm.contrib import graph_executor +from tvm.contrib.debugger import debug_executor +from . import GraphModule + + +def tune_tasks( + tasks, + measure_option, + tuner="xgb", + n_trial=1000, + early_stopping=None, + log_filename="tuning.log", + use_transfer_learning=True, +): + """Tune tasks and generate tuning log to file""" + # create tmp log file + tmp_log_file = log_filename + ".tmp" + if os.path.exists(tmp_log_file): + os.remove(tmp_log_file) + + for i, tsk in enumerate(reversed(tasks)): + prefix = f"[Task {i + 1:2d}/{len(tasks):2d}] " + + # create tuner + if tuner in ("xgb", "sgb-rank"): + tuner_obj = XGBTuner(tsk, loss_type="rank") + elif tuner == "ga": + tuner_obj = GATuner(tsk, pop_size=100) + elif tuner == "random": + tuner_obj = RandomTuner(tsk) + elif tuner == "gridsearch": + tuner_obj = GridSearchTuner(tsk) + else: + raise ValueError("Invalid tuner: " + tuner) + + if use_transfer_learning: + if os.path.isfile(tmp_log_file): + tuner_obj.load_history(autotvm.record.load_from_file(tmp_log_file)) + + # do tuning + tsk_trial = min(n_trial, len(tsk.config_space)) + tuner_obj.tune( + n_trial=tsk_trial, + early_stopping=early_stopping, + measure_option=measure_option, + callbacks=[ + autotvm.callback.progress_bar(tsk_trial, prefix=prefix), + autotvm.callback.log_to_file(tmp_log_file), + ], + ) + + # pick best records to a cache file + if not os.path.exists(log_filename): + with open(log_filename, "w", encoding="utf-8"): + pass + if os.path.exists(tmp_log_file): + autotvm.record.pick_best(tmp_log_file, log_filename) + os.remove(tmp_log_file) + + +def get_tuning_opt(log_file="tuning.log", n_trial=200): + """Returns tuning options""" + tuning_opt = { + "log_filename": log_file, + "tuner": "random", + "n_trial": n_trial, + "early_stopping": 60, + "measure_option": autotvm.measure_option( + builder=autotvm.LocalBuilder(timeout=10), + runner=autotvm.LocalRunner(number=20, repeat=3, timeout=4, min_repeat_ms=150), + ), + } + return tuning_opt + + +TVM_ASSETS = ["mod.so", "graph.json", "params"] + + +class PyTorchTVMModule: + """Helper class for compiling pytorch module to tvm module""" + + def __init__(self, target="cuda", device=tvm.cuda(0)) -> None: + self.script_module = None + self.input_infos = None + self.default_dtype = "float32" + self.mod = None + self.params = None + self.tasks = None + self.target = target + self.dev = device + self.log_file = None + self.tvm_module = None + self.tvm_graph = None + self.tvm_lib = None + self.tvm_params = None + + def from_pytorch(self, script_module, input_infos, default_dtype="float32"): + self.script_module = script_module + self.input_infos = input_infos + self.default_dtype = default_dtype + self.mod, self.params = relay.frontend.from_pytorch( + script_module, input_infos, default_dtype=default_dtype + ) + + def tune_tvm(self, log_file="tuning.log", n_trial=200): + self.tasks = autotvm.task.extract_from_program( + self.mod["main"], + target=self.target, + params=self.params, + ) + self.log_file = log_file + tuning_opt = get_tuning_opt(log_file, n_trial) + tune_tasks(self.tasks, **tuning_opt) + + def build_tvm(self, export_dir, debug_runtime=False): + tvm_mod = self._build_tvm(debug_runtime) + self._export_tvm(export_dir) + return tvm_mod + + def _build_tvm(self, debug_runtime=False): + # compile kernels with history best records + with autotvm.apply_history_best(self.log_file): + with tvm.transform.PassContext(opt_level=3): + self.tvm_graph, self.tvm_lib, self.tvm_params = relay.build( + self.mod, target=self.target, params=self.params + ) + + if not debug_runtime: + self.tvm_module = graph_executor.create(self.tvm_graph, self.tvm_lib, device=self.dev) + else: + self.tvm_module = debug_executor.create(self.tvm_graph, self.tvm_lib, device=self.dev) + self.tvm_module.set_input(**self.tvm_params) + return self.tvm_module + + def _export_tvm(self, export_dir): + if not os.path.isdir(export_dir): + os.makedirs(export_dir) + self.export_dir = export_dir + self.tvm_lib.export_library(os.path.join(export_dir, TVM_ASSETS[0])) + with open(os.path.join(export_dir, TVM_ASSETS[1]), "w", encoding="utf8") as fout: + fout.write(self.tvm_graph) + with open(os.path.join(export_dir, TVM_ASSETS[2]), "wb") as fout: + fout.write(relay.save_param_dict(self.tvm_params)) + + def load_tvm(self, export_dir): + """Load tvm module from export directory""" + self.export_dir = export_dir + self.tvm_lib = load_module(os.path.join(export_dir, TVM_ASSETS[0])) + with open(os.path.join(export_dir, TVM_ASSETS[1]), "r", encoding="utf8") as f: + self.tvm_graph = f.read() + with open(os.path.join(export_dir, TVM_ASSETS[2]), "rb") as f: + self.tvm_params = relay.load_param_dict(f.read()) + + self.tvm_module = graph_executor.create(self.tvm_graph, self.tvm_lib, device=self.dev) + self.tvm_module.set_input(**self.tvm_params) + return self.tvm_module + + def build_pytorch_module(self, num_inputs, num_outputs, input_infos=None): + """Build pytorch module containing TVM Graph Module""" + assert self.export_dir, "you must build_tvm or load_tvm before" + input_infos = input_infos or self.input_infos + assert input_infos + assert len(input_infos) == num_inputs + assets = [os.path.join(self.export_dir, i) for i in TVM_ASSETS] + input_shapes = [i[1] for i in input_infos] + + def _tvm_dev_to_pt_dev(device): + """convert tvm device to pytorch device string""" + if tvm.runtime.Device.MASK2STR[device.device_type] == "cpu": + return "cpu" + if tvm.runtime.Device.MASK2STR[device.device_type] == "cuda": + return f"cuda:{device.device_id}" + raise ValueError(f"unsupported device for pt graph module: {device}") + + mod = GraphModule(num_inputs=num_inputs, num_outputs=num_outputs).to( + _tvm_dev_to_pt_dev(self.dev) + ) + mod.init(input_shapes, *assets) + return mod + + +def compile(script_module, option): + """ + example: + option = { + "input_infos": [ + ("x", (1, 3, 244, 244)), + ], + "default_dtype": "float16", + "export_dir": "pytorch_compiled", + "num_outputs": 1, + "tuning_n_trials": 20, # set zero to skip tuning + "tuning_log_file": "tuning.log", + "target": "llvm", + "device": tvm.cpu(), + } + script_module = torch.jit.script(model) + pytorch_tvm_module = compile(script_module, option) + pytorch_tvm_module("model_tvm.pt") + """ + input_infos = option["input_infos"] + default_dtype = option.get("default_dtype", "float32") + export_dir = option.get("export_dir", "pytorch_compiled") + tuning_log_file = option.get("tuning_log_file", "tuning.log") + tuning_n_trials = option.get("tuning_n_trials", 20) + num_outputs = option.get("num_outputs", 1) + target = option.get("target", "cuda") + device = option.get("device", tvm.cuda(0)) + + mod = PyTorchTVMModule(target=target, device=device) + print("Converting...") + + mod.log_file = tuning_log_file + mod.from_pytorch(script_module, input_infos, default_dtype) + + if tuning_n_trials > 0: + print("Tuning...") + mod.tune_tvm(log_file=tuning_log_file, n_trial=tuning_n_trials) + + print("Building...") + mod.build_tvm(export_dir) + pytorch_mod = mod.build_pytorch_module(num_inputs=len(input_infos), num_outputs=num_outputs) + return pytorch_mod diff --git a/src/contrib/torch/pt_call_tvm/tvm_class.cc b/src/contrib/torch/pt_call_tvm/tvm_class.cc new file mode 100644 index 000000000000..5e57dc152f11 --- /dev/null +++ b/src/contrib/torch/pt_call_tvm/tvm_class.cc @@ -0,0 +1,686 @@ +/* + * 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. + */ +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include + +#include "../utils.h" + +namespace tvm { +namespace contrib { +namespace pytorch { + +/*! \brief Class holding necessary components to call TVM graph runtime */ +class TvmGraphModulePack { + public: + /*! + * \brief Constructor. + * + * \param path Encoded path of graph runtime assets. + * \param device_type int64_t, kDLCPU or kDLCUDA. + * \param device_id int64_t. + */ + explicit TvmGraphModulePack(std::string path, int64_t device_type, int64_t device_id) + : path_(std::move(path)) { + LOG(INFO) << "[TvmGraphModule] loading module at path: [" << path_ << "] on device [" + << (device_type == kDLCUDA ? "cuda:" : "cpu:") << device_id << "]..."; + std::string lib_path, graph_path, params_path; + DecodePaths(path_, &lib_path, &graph_path, ¶ms_path); + + // load graph + std::ifstream graph_in(graph_path); + std::string graph_data((std::istreambuf_iterator(graph_in)), + std::istreambuf_iterator()); + graph_in.close(); + + // load mod syslib + tvm::runtime::Module lib = tvm::runtime::Module::LoadFromFile(lib_path); + + const auto runtime_create = *tvm::runtime::Registry::Get("tvm.graph_executor.create"); + + // read params data + std::ifstream params_in(params_path, std::ios::binary); + std::string params_data((std::istreambuf_iterator(params_in)), + std::istreambuf_iterator()); + params_in.close(); + TVMByteArray params_arr; + params_arr.data = params_data.c_str(); + params_arr.size = params_data.length(); + + // set devices + module_ = runtime_create(graph_data, lib, device_type, device_id); + const tvm::runtime::PackedFunc load_params = module_.GetFunction("load_params"); + load_params(params_arr); + + set_input = module_.GetFunction("set_input_zero_copy"); + run = module_.GetFunction("run"); + get_output = module_.GetFunction("get_output"); + set_output = module_.GetFunction("set_output_zero_copy"); + num_outputs_ = module_.GetFunction("get_num_outputs")(); + } + + static constexpr char kPathDelimiter = '|'; + + /*! + * \brief Decode lib_path, graph_path, params_path from encoded path. + * + * \param path The encoded path, concated with `kPathDelimiter`. + * \param lib_path The path of .so lib file. + * \param graph_path The path of graph.json. + * \param params_path The path of params data. + */ + static void DecodePaths(const std::string& path, std::string* lib_path, std::string* graph_path, + std::string* params_path) { + std::vector paths; + for (size_t i = 0, pre = 0, lim = path.size(); i <= lim; ++i) { + if (i == lim || path.at(i) == kPathDelimiter) { + paths.push_back(path.substr(pre, i - pre)); + pre = i + 1; + } + } + CHECK_EQ(paths.size(), 3u); + *lib_path = paths.at(0); + *graph_path = paths.at(1); + *params_path = paths.at(2); + } + + /*! + * \brief Encode lib_path, graph_path, params_path by concat then with `kPathDelimiter`. + * + * \param lib_path The path of .so lib file. + * \param graph_path The path of graph.json. + * \param params_path The path of params data. + * + * \return The encoded path, concated with `kPathDelimiter`. + */ + static std::string EncodePaths(const std::string& lib_path, const std::string& graph_path, + const std::string& params_path) { + return lib_path + kPathDelimiter + graph_path + kPathDelimiter + params_path; + } + + const std::string& path() const { return path_; } + + const int64_t num_outputs() const { return num_outputs_; } + + tvm::runtime::PackedFunc set_input; + tvm::runtime::PackedFunc run; + tvm::runtime::PackedFunc get_output; + tvm::runtime::PackedFunc set_output; + + private: + tvm::runtime::Module module_; + int64_t num_outputs_; + std::string path_; +}; + +/*! \brief Class holding necessary components to call TVM VM runtime */ +class TvmVMModulePack { + public: + /*! + * \brief Constructor. + * + * \param path Encoded path of vm runtime assets. + * \param device_type int64_t, kDLCPU or kDLCUDA. + * \param device_id int64_t. + */ + explicit TvmVMModulePack(std::string path, int64_t device_type, int64_t device_id) + : path_(std::move(path)) { + LOG(INFO) << "[TvmVMModule] loading module at path: [" << path_ << "] on device [" + << (device_type == kDLCUDA ? "cuda:" : "cpu:") << device_id << "]..."; + // build tvm graph runtime + std::string lib_path, code_path; + DecodePaths(path_, &lib_path, &code_path); + // load lib + auto loaded_lib = tvm::runtime::Module::LoadFromFile(lib_path, "so"); + // load code + std::ifstream code_in(code_path); + std::string loaded_code((std::istreambuf_iterator(code_in)), + std::istreambuf_iterator()); + code_in.close(); + exe_ = tvm::runtime::vm::Executable::Load(loaded_code, loaded_lib); + const auto runtime_create = *tvm::runtime::Registry::Get("runtime._VirtualMachine"); + vm_ = runtime_create(exe_); + auto init_func = vm_.GetFunction("init", false); + auto alloc_type = static_cast(tvm::runtime::vm::AllocatorType::kPooled); + if (device_type != kDLCPU) { + // CPU is required for executing shape functions + init_func(static_cast(kDLCPU), 0, alloc_type, device_type, device_id, alloc_type); + } else { + init_func(device_type, device_id, alloc_type); + } + set_input = vm_.GetFunction("set_input", false); + invoke = vm_.GetFunction("invoke", false); + } + + static constexpr char kPathDelimiter = '|'; + + /*! + * \brief Decode lib_path, code_path from encoded path. + * + * \param path The encoded path, concated with `kPathDelimiter`. + * \param lib_path The path of lib file. + * \param code_path The path of code file. + */ + static void DecodePaths(const std::string& path, std::string* lib_path, std::string* code_path) { + std::vector paths; + for (size_t i = 0, pre = 0, lim = path.size(); i <= lim; ++i) { + if (i == lim || path.at(i) == kPathDelimiter) { + paths.push_back(path.substr(pre, i - pre)); + pre = i + 1; + } + } + CHECK_EQ(paths.size(), 2u); + *lib_path = paths.at(0); + *code_path = paths.at(1); + } + + /*! + * \brief Encode lib_path, code_path by concat then with `kPathDelimiter`. + * + * \param lib_path The path of vm lib file. + * \param code_path The path of code. + * + * \return The encoded path, concated with `kPathDelimiter`. + */ + static std::string EncodePaths(const std::string& lib_path, const std::string& code_path) { + return lib_path + kPathDelimiter + code_path; + } + + const std::string& path() const { return path_; } + + tvm::runtime::PackedFunc set_input; + tvm::runtime::PackedFunc invoke; + + private: + tvm::runtime::Module exe_; + tvm::runtime::Module vm_; + std::string path_; +}; + +/*! \brief Pytorch custom class to call TVM */ +class BaseTvmClass : public torch::jit::CustomClassHolder { + public: + /*! + * \brief Constructor. + * + * \param num_inputs Number of inputs. + * \param num_outputs Number of outputs. + * \param device std::string, use the pytorch device str format, e.g. `cuda:0`, 'cpu' + */ + BaseTvmClass(const int64_t num_inputs, const int64_t num_outputs, const std::string& device) + : num_inputs_(num_inputs), num_outputs_(num_outputs) { + auto torch_device = torch::Device(device); + device_type_ = torch_device.is_cuda() ? kDLCUDA : kDLCPU; + device_id_ = torch_device.index(); + } + + /*! \brief Virtual destructor. */ + virtual ~BaseTvmClass() {} + + /*! + * \brief Get repr string of pytorch input shapes. + * + * \param shapes Pytorch shapes of type List[List[int]]. + * + * \return std::string, the representation of inputs shapes. + */ + static std::string TvmShapeRepr(const c10::List>& shapes) { + std::stringstream ss; + for (const auto& shape : shapes) { + for (const auto& sz : static_cast>(shape)) { + ss << sz << "_"; + } + ss << "__"; + } + return ss.str(); + } + + /*! + * \brief Get input shapes. + * + * \param inputs Inputs with type List[Tensor]. + * + * \return outputs with type List[List[int]]. + */ + static c10::List> GetShapes(const c10::List& inputs) { + c10::List> shapes; + for (const auto& input : inputs) { + c10::List shape; + for (const auto sz : static_cast(input).sizes()) { + shape.push_back(sz); + } + shapes.push_back(shape); + } + return shapes; + } + + /*! + * \brief Move the TVM modules to given device. + * + * \param device String repr of the device to be moved to. + */ + virtual void to(const std::string& device) = 0; + + // getters + int64_t num_inputs() const { return num_inputs_; } + + int64_t num_outputs() const { return num_outputs_; } + + int64_t device_type() const { return device_type_; } + + int64_t device_id() const { return device_id_; } + + c10::DeviceType torch_device_type() const { + return device_type() == kDLCUDA ? torch::DeviceType::CUDA : torch::DeviceType::CPU; + } + + bool is_on_same_device(const torch::Tensor& tensor) const { + auto tensor_device_type = tensor.device().type(); + if (tensor_device_type == torch::DeviceType::CUDA) { + return tensor_device_type == torch_device_type() && device_id() == tensor.device().index(); + } + CHECK_EQ(tensor_device_type, torch::DeviceType::CPU); + return tensor_device_type == torch_device_type(); + } + + std::string device() const { return torch::Device(torch_device_type(), device_id()).str(); } + + /*! + * \brief Module forward. + * + * \param inputs Inputs with type List[Tensor]. + * + * \return outputs with type List[Tensor]. + */ + virtual c10::List forward(const c10::List& inputs) = 0; + + /*! + * \brief Serialize TVM Modules to Dict + */ + virtual c10::Dict SerializeTvmModules() const = 0; + + /*! + * \brief deserialize TVM Modules from Dict + */ + virtual void DeserializeTvmModules(const c10::Dict& shape_path_map) = 0; + + protected: + const int64_t num_inputs_; + const int64_t num_outputs_; + int64_t device_type_; + int64_t device_id_; +}; + +/*! \brief Pytorch custom class to call TVM graph runtime */ +class TvmGraphRuntimeClass : public BaseTvmClass { + public: + TvmGraphRuntimeClass(const int64_t num_inputs, const int64_t num_outputs, + const std::string& device) + : BaseTvmClass(num_inputs, num_outputs, device) {} + + /*! + * \brief Module forward. + * + * \param inputs Inputs with type List[Tensor]. + * + * \return outputs with type List[Tensor]. + */ + c10::List forward(const c10::List& inputs) override { + CHECK_EQ(inputs.size(), num_inputs_); + auto shape_repr = TvmShapeRepr(GetShapes(inputs)); + std::vector args(num_inputs_ + num_outputs_); + auto iter = tvm_modules_.find(shape_repr); + CHECK(iter != tvm_modules_.end()); + const auto& tvm_pack = iter->second; + std::vector buf_infos; + buf_infos.reserve(num_inputs_ + num_outputs_); + + for (int i = 0; i < num_inputs_; ++i) { + at::Tensor inp = inputs[i]; + CHECK(is_on_same_device(inp)) + << "input #" << i + << " of forward is not on the same device with TvmGraphRuntime, expected " << device() + << " but got " << inp.device().str(); + inp = inp.contiguous(); + buf_infos.emplace_back(inp); + auto& input_buf = buf_infos[i]; + input_buf.CopyFromOrigin(); + input_buf.MakeDLTensor(&args[i]); + tvm_pack.set_input(i, &args[i]); + } + // prepare output buffers + c10::List outputs; + outputs.reserve(num_outputs_); + + for (int i = 0; i < num_outputs_; ++i) { + tvm::runtime::NDArray output_arr = tvm_pack.get_output(i); + std::vector output_shape(output_arr->shape, output_arr->shape + output_arr->ndim); + + torch::ScalarType output_dtype = torch::ScalarType::Undefined; + CHECK(GetTorchDtype(output_arr.DataType(), &output_dtype)); + + CHECK(device_type_ == kDLCPU || device_type_ == kDLCUDA); + const c10::DeviceType pt_device_type = (device_type_ == kDLCUDA ? torch::kCUDA : torch::kCPU); + const auto options = + torch::TensorOptions().dtype(output_dtype).device(pt_device_type, device_id_); + + outputs.emplace_back(torch::empty(output_shape, options)); + buf_infos.emplace_back(outputs[i]); + auto& output_buf = buf_infos[num_inputs_ + i]; + output_buf.MakeDLTensor(&args[num_inputs_ + i]); + tvm_pack.set_output(i, &args[num_inputs_ + i]); + } + tvm_pack.run(); + for (int i = 0; i < num_outputs_; ++i) { + auto& output_buf = buf_infos[num_inputs_ + i]; + output_buf.CopyToOrigin(); + } + return outputs; + } + + /*! + * \brief Load TVM graph runtime module. + * + * \param shapes Input shapes. List[List[int]]. + * \param lib_path Path of .so lib file. + * \param graph_path Path of graph.json file. + * \param params_path Path of params data file. + */ + void LoadTvmModule(const c10::List>& shapes, const std::string& lib_path, + const std::string& graph_path, const std::string& params_path) { + std::string path = TvmGraphModulePack::EncodePaths(lib_path, graph_path, params_path); + auto shape_repr = TvmShapeRepr(shapes); + auto it_find = tvm_modules_.find(shape_repr); + if (it_find != tvm_modules_.end()) { + tvm_modules_.erase(it_find); + } + const auto it = + tvm_modules_.emplace(shape_repr, TvmGraphModulePack(path, device_type_, device_id_)).first; + if (it->second.num_outputs() != num_outputs_) { + LOG(FATAL) << "tvm class num outputs mismatch, expected " << num_outputs_ << ", got " + << it->second.num_outputs(); + } + } + + const std::map& tvm_modules() const { return tvm_modules_; } + + /*! + * \brief Serialize TVM modules to shape map. + * + * \return shape_path_map Dict of shape_repr to path. + */ + c10::Dict SerializeTvmModules() const override { + c10::Dict shape_path_map; + for (const auto& entry : tvm_modules()) { + shape_path_map.insert(entry.first, entry.second.path()); + } + return shape_path_map; + } + + /*! + * \brief Deserialize TVM modules from shape map. + * + * \param shape_path_map Dict of shape_repr to path. + */ + void DeserializeTvmModules(const c10::Dict& shape_path_map) override { + tvm_modules_.clear(); + for (const auto& entry : shape_path_map) { + const auto& shape_repr = entry.key(); + const auto& path = entry.value(); + tvm_modules_.emplace(shape_repr, TvmGraphModulePack(path, device_type_, device_id_)); + } + } + + /*! + * \brief Move the TVM modules to given device. + * + * \param device String repr of the device to be moved to. + */ + void to(const std::string& device) override { + if (device != this->device()) { + auto torch_device = torch::Device(device); + device_type_ = torch_device.is_cuda() ? kDLCUDA : kDLCPU; + device_id_ = torch_device.index(); + DeserializeTvmModules(SerializeTvmModules()); + } + } + + private: + std::map tvm_modules_; +}; + +/*! \brief Pytorch custom class to call TVM graph runtime */ +class TvmVMRuntimeClass : public BaseTvmClass { + public: + TvmVMRuntimeClass(const int64_t num_inputs, const int64_t num_outputs, const std::string& device) + : BaseTvmClass(num_inputs, num_outputs, device) {} + + /*! + * \brief Module forward. + * + * \param inputs Inputs with type List[Tensor]. + * + * \return outputs with type List[Tensor]. + */ + c10::List forward(const c10::List& inputs) override { + // get inputs repr str + auto shape_repr = TvmShapeRepr(GetShapes(inputs)); + // get tvm pack + auto iter = tvm_modules_.find(shape_repr); + CHECK(iter != tvm_modules_.end()) << "tvm module pack not found for shape_repr " << shape_repr; + const auto& tvm_pack = iter->second; + + // input tensors + CHECK_EQ(inputs.size(), num_inputs_); + std::vector args(num_inputs_); + std::vector args_arr(num_inputs_); + + for (int i = 0; i < num_inputs_; ++i) { + TensorAsBuf input_buf(inputs[i]); + input_buf.CopyFromOrigin(); + input_buf.MakeDLTensor(&args[i]); + args_arr[i] = + tvm::runtime::NDArray::FromDLPack(new DLManagedTensor({args[i], nullptr, nullptr})); + } + // set input + std::vector tvm_values(num_inputs_ + 1); + std::vector tvm_type_codes(num_inputs_ + 1); + tvm::runtime::TVMArgsSetter setter(tvm_values.data(), tvm_type_codes.data()); + setter(0, "main"); + for (int k = 0; k < num_inputs_; ++k) { + setter(k + 1, args_arr[k]); + } + tvm_pack.set_input.CallPacked( + tvm::runtime::TVMArgs(tvm_values.data(), tvm_type_codes.data(), num_inputs_ + 1), nullptr); + + // run tvm + tvm::runtime::TVMRetValue ret = tvm_pack.invoke("main"); + + // get outputs + std::vector output_arrs(num_outputs_); + auto output_mismatch_msg = [](int actual, int expected) { + std::stringstream ss; + ss << "num_outputs not equal, actual:[" << actual << "] != expected:[" << expected << "]"; + return ss.str(); + }; + if (ret.type_code() == kTVMNDArrayHandle) { + CHECK_EQ(num_outputs_, 1) << output_mismatch_msg(1, num_outputs_); + output_arrs.at(0) = ret.AsObjectRef(); + } else if (ret.type_code() == kTVMObjectHandle) { + const auto& adt = ret.AsObjectRef(); + CHECK_EQ(adt.size(), num_outputs_) << output_mismatch_msg(adt.size(), num_outputs_); + for (size_t i = 0; i < adt.size(); ++i) { + CHECK(adt[i]->IsInstance()) + << "adt elements not tvm::runtime::NDArray"; + output_arrs.at(i) = tvm::runtime::Downcast(adt[i]); + } + } else { + LOG(FATAL) << "unsupported return type with type_code = " << ret.type_code(); + } + + std::vector output_args(num_outputs_); + c10::List outputs; + outputs.reserve(num_outputs_); + + for (int i = 0; i < num_outputs_; ++i) { + const auto& output_arr = output_arrs[i]; + std::vector output_shape(output_arr->shape, output_arr->shape + output_arr->ndim); + + torch::ScalarType output_dtype = torch::ScalarType::Undefined; + CHECK(GetTorchDtype(output_arr.DataType(), &output_dtype)); + + CHECK(device_type_ == kDLCPU || device_type_ == kDLCUDA); + const c10::DeviceType pt_device_type = (device_type_ == kDLCUDA ? torch::kCUDA : torch::kCPU); + const auto options = + torch::TensorOptions().dtype(output_dtype).device(pt_device_type, device_id_); + + outputs.emplace_back(torch::empty(output_shape, options)); + TensorAsBuf output_buf(outputs[i]); + output_buf.MakeDLTensor(&output_args[i]); + output_arr.CopyTo(&output_args[i]); + output_buf.CopyToOrigin(); + } + return outputs; + } + + /*! + * \brief Load TVM vm runtime module. + * + * \param shapes Input shapes. List[List[int]]. + * \param lib_path Path of .so lib file. + * \param code_path Path of code file. Typically named code.ro + */ + void LoadTvmModule(const c10::List>& shapes, const std::string& lib_path, + const std::string& code_path) { + std::string path = TvmVMModulePack::EncodePaths(lib_path, code_path); + auto shape_repr = TvmShapeRepr(shapes); + auto it_find = tvm_modules_.find(shape_repr); + if (it_find != tvm_modules_.end()) { + tvm_modules_.erase(it_find); + } + tvm_modules_.emplace(shape_repr, TvmVMModulePack(path, device_type_, device_id_)); + } + + const std::map& tvm_modules() const { return tvm_modules_; } + + /*! + * \brief Serialize TVM modules to shape map. + * + * \return shape_path_map Dict of shape_repr to path. + */ + c10::Dict SerializeTvmModules() const override { + c10::Dict shape_path_map; + for (const auto& entry : tvm_modules()) { + shape_path_map.insert(entry.first, entry.second.path()); + } + return shape_path_map; + } + + /*! + * \brief Deserialize TVM modules from shape map. + * + * \param shape_path_map Dict of shape_repr to path. + */ + void DeserializeTvmModules(const c10::Dict& shape_path_map) override { + tvm_modules_.clear(); + for (const auto& entry : shape_path_map) { + const auto& shape_repr = entry.key(); + const auto& path = entry.value(); + tvm_modules_.emplace(shape_repr, TvmVMModulePack(path, device_type_, device_id_)); + } + } + + /*! + * \brief Move the TVM modules to given device. + * + * \param device String repr of the device to be moved to. + */ + void to(const std::string& device) override { + if (device != this->device()) { + auto torch_device = torch::Device(device); + device_type_ = torch_device.is_cuda() ? kDLCUDA : kDLCPU; + device_id_ = torch_device.index(); + DeserializeTvmModules(SerializeTvmModules()); + } + } + + private: + std::map tvm_modules_; +}; + +// +using SerializeTuple = + std::tuple>; + +/***** registries *****/ +static auto __tvm_dsoop_graph_runtime_registry = + torch::jit::class_("tvm_dsoop", "TvmGraphModule") + .def(torch::init()) + .def("load_tvm_module", &TvmGraphRuntimeClass::LoadTvmModule) + .def("forward", &TvmGraphRuntimeClass::forward) + .def("to", &TvmGraphRuntimeClass::to) + .def_pickle( + [](const c10::intrusive_ptr& self) -> SerializeTuple { + return std::make_tuple(self->num_inputs(), self->num_outputs(), self->device(), + self->SerializeTvmModules()); + }, + [](SerializeTuple tuple) -> c10::intrusive_ptr { + auto ptr = c10::make_intrusive( + /*num_inputs=*/std::get<0>(tuple), + /*num_outputs=*/std::get<1>(tuple), + /*device=*/std::get<2>(tuple)); + ptr->DeserializeTvmModules(std::get<3>(tuple)); + return ptr; + }); + +static auto __tvm_dsoop_vm_runtime_registry = + torch::jit::class_("tvm_dsoop", "TvmVMModule") + .def(torch::init()) + .def("load_tvm_module", &TvmVMRuntimeClass::LoadTvmModule) + .def("forward", &TvmVMRuntimeClass::forward) + .def("to", &TvmVMRuntimeClass::to) + .def_pickle( + [](const c10::intrusive_ptr& self) -> SerializeTuple { + return std::make_tuple(self->num_inputs(), self->num_outputs(), self->device(), + self->SerializeTvmModules()); + }, + [](SerializeTuple tuple) -> c10::intrusive_ptr { + auto ptr = c10::make_intrusive( + /*num_inputs=*/std::get<0>(tuple), + /*num_outputs=*/std::get<1>(tuple), + /*device=*/std::get<2>(tuple)); + ptr->DeserializeTvmModules(std::get<3>(tuple)); + return ptr; + }); + +static auto __tvm_shape_repr_fn_registry = + torch::RegisterOperators("tvm_dsoop::tvm_shape_repr", &BaseTvmClass::TvmShapeRepr); +} // namespace pytorch +} // namespace contrib +} // namespace tvm diff --git a/src/contrib/torch/utils.h b/src/contrib/torch/utils.h new file mode 100644 index 000000000000..a98e058ca346 --- /dev/null +++ b/src/contrib/torch/utils.h @@ -0,0 +1,264 @@ +/* + * 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. + */ + +/*! + * \file utils.h + * \brief Util functions for pytorch tvm interaction. + */ + +#ifndef TVM_CONTRIB_TORCH_UTILS_H_ +#define TVM_CONTRIB_TORCH_UTILS_H_ + +#include +#include +#include +#include +#ifdef PT_TVMDSOOP_ENABLE_GPU +#include +#endif + +#include +#include + +namespace tvm { +namespace contrib { +namespace pytorch { + +inline bool GetTvmDtype(const caffe2::TypeMeta& dtype, DLDataType* res) noexcept { + if (dtype == torch::kFloat16) { + *res = {kDLFloat, 16, 1}; + } else if (dtype == torch::kFloat32) { + *res = {kDLFloat, 32, 1}; + } else if (dtype == torch::kFloat64) { + *res = {kDLFloat, 64, 1}; + } else if (dtype == torch::kInt8) { + *res = {kDLInt, 8, 1}; + } else if (dtype == torch::kInt16) { + *res = {kDLInt, 16, 1}; + } else if (dtype == torch::kInt32) { + *res = {kDLInt, 32, 1}; + } else if (dtype == torch::kInt64) { + *res = {kDLInt, 64, 1}; + } else if (dtype == torch::kUInt8) { + *res = {kDLUInt, 8, 1}; + } else if (dtype == torch::kBool) { + *res = {kDLInt, 1, 1}; + } else { + return false; + } + return true; +} + +inline bool GetTvmDtype(const caffe2::TypeMeta& dtype, tvm::runtime::DataType* res) noexcept { + DLDataType dlpack_dtype; + + if (!GetTvmDtype(dtype, &dlpack_dtype)) { + return false; + } + *res = tvm::runtime::DataType(dlpack_dtype); + return true; +} + +inline bool GetTorchDtype(const DLDataType& dtype, c10::ScalarType* res) noexcept { + if (dtype.lanes != 1) { + // only scalar type + return false; + } + if (dtype.code == kDLFloat) { + if (dtype.bits == 16) { + *res = torch::kFloat16; + } else if (dtype.bits == 32) { + *res = torch::kFloat32; + } else if (dtype.bits == 64) { + *res = torch::kFloat64; + } else { + return false; + } + } else if (dtype.code == kDLInt) { + if (dtype.bits == 16) { + *res = torch::kInt16; + } else if (dtype.bits == 32) { + *res = torch::kInt32; + } else if (dtype.bits == 64) { + *res = torch::kInt64; + } else if (dtype.bits == 1) { + *res = torch::kBool; + } else { + return false; + } + } else if (dtype.code == kDLUInt) { + if (dtype.bits == 8) { + *res = torch::kUInt8; + } else if (dtype.bits == 1) { + *res = torch::kBool; + } else { + return false; + } + } else { + return false; + } + return true; +} + +inline bool GetTorchDtype(const tvm::runtime::DataType& dtype, c10::ScalarType* res) noexcept { + using tvm::runtime::DataType; + if (dtype == DataType::Float(16)) { + *res = torch::kFloat16; + } else if (dtype == DataType::Float(32)) { + *res = torch::kFloat32; + } else if (dtype == DataType::Float(64)) { + *res = torch::kFloat64; + } else if (dtype == DataType::Int(32)) { + *res = torch::kInt32; + } else if (dtype == DataType::Int(64)) { + *res = torch::kInt64; + } else if (dtype == DataType::Int(1)) { + *res = torch::kBool; + } else if (dtype == DataType::Int(8)) { + *res = torch::kInt8; + } else if (dtype == DataType::Int(16)) { + *res = torch::kInt16; + } else if (dtype == DataType::UInt(8)) { + *res = torch::kUInt8; + } else if (dtype == DataType::Bool()) { + *res = torch::kBool; + } else { + return false; + } + return true; +} + +// Buffer information used for actual computation. +// Each buffer is associated with one PyTorch tensor +// whose underlying buffer is record into "origin_buf". +// For input tensor, we copy data from origin_buf to buf +// and for output tensor, copy data from buf to origin_buf +class TensorAsBuf { + public: + explicit TensorAsBuf(const at::Tensor& tensor) + : pt_device_type_(tensor.device().type()), + device_id_(tensor.device().index()), + origin_shape_(tensor.sizes().begin(), tensor.sizes().end()) { + CHECK(pt_device_type_ == torch::kCUDA || pt_device_type_ == torch::kCPU); + device_type_ = (pt_device_type_ == torch::kCUDA ? kDLCUDA : kDLCPU); + + char* buf = static_cast(tensor.data_ptr()); + this->origin_buf_ = buf; + this->size_ = tensor.nbytes(); + + // const int alignment = 64; + const int alignment = tvm::runtime::kAllocAlignment; + char* aligned = reinterpret_cast(((uint64_t)buf + alignment - 1) & (~(alignment - 1))); + if (buf == aligned) { + this->tensor_ = tensor; + this->buf_ = buf; + this->offset_ = 0; + } else { + const auto options = + torch::TensorOptions().dtype(tensor.dtype()).device(pt_device_type_, device_id_); + this->inline_tensor_ = + torch::empty({static_cast(tensor.nbytes() + alignment)}, options); + this->tensor_ = this->inline_tensor_; + + buf = static_cast(this->tensor_.data_ptr()); + char* buf_aligned = reinterpret_cast(((uint64_t)buf + alignment) & (~(alignment - 1))); + this->buf_ = buf; + this->offset_ = buf_aligned - buf; + } + } + + void CopyToOrigin() { + if (buf_ == origin_buf_) { + return; + } + if (device_type_ == kDLCPU) { + memcpy(origin_buf_, buf_ + offset_, size_); +#ifdef PT_TVMDSOOP_ENABLE_GPU + } else if (device_type_ == kDLCUDA) { + cudaMemcpy(origin_buf_, buf_ + offset_, size_, cudaMemcpyDeviceToDevice); +#endif + } else { + LOG(FATAL) << "Only support CPU and CUDA now. Device " << device_type_ + << " is not implemented currently"; + } + } + + void CopyFromOrigin() { + if (buf_ == origin_buf_) { + return; + } + if (device_type_ == kDLCPU) { + memcpy(buf_ + offset_, origin_buf_, size_); +#ifdef PT_TVMDSOOP_ENABLE_GPU + } else if (device_type_ == kDLCUDA) { + cudaMemcpy(buf_ + offset_, origin_buf_, size_, cudaMemcpyDeviceToDevice); +#endif + } else { + LOG(FATAL) << "Only support CPU and CUDA now. Device " << device_type_ + << " is not implemented currently"; + } + } + + // Create DLPack tensor from PyTorch tensor + void MakeDLTensor(DLTensor* out) { + const DLDevice dl_ctx{DLDeviceType(device_type_), device_id_}; + DLDataType dlpack_type; + const auto& tensor = this->tensor_; + CHECK(GetTvmDtype(tensor.dtype(), &dlpack_type)); + + out->device = dl_ctx; + out->ndim = origin_shape_.size(); + out->shape = origin_shape_.data(); + out->strides = nullptr; + out->byte_offset = 0; + out->dtype = dlpack_type; + out->data = buf_ + offset_; + } + + std::string DebugString() { + std::stringstream ss; + ss << "dl device: " << device_type_ << "\npt device: " << static_cast(pt_device_type_) + << "\ndevice_id: " << device_id_ << "\nsize: " << size_ << "\noffset: " << offset_ + << "\nshape:"; + for (auto dim : origin_shape_) { + ss << ' ' << dim; + } + ss << std::endl; + return ss.str(); + } + + private: + DLDeviceType device_type_; + c10::DeviceType pt_device_type_; + int device_id_; + + at::Tensor inline_tensor_; + at::Tensor tensor_; + size_t size_; + size_t offset_; + + std::vector origin_shape_; + + char* origin_buf_; + char* buf_; +}; +} // namespace pytorch +} // namespace contrib +} // namespace tvm +#endif // TVM_CONTRIB_TORCH_UTILS_H_